Fix more misleading line for creating list of models in rollout_test.py

PiperOrigin-RevId: 717628943
Change-Id: I83a314565ce4ec683afeea18c8c8ce9e23d5452e
This commit is contained in:
Yuval Tassa
2025-01-20 14:26:33 -08:00
committed by Copybara-Service
parent 68cd263337
commit baf1c43b3c
+4 -4
View File
@@ -532,7 +532,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
sensordata = np.empty((nroll, nstep, model.nsensordata))
control = np.random.randn(nroll, nstep, model.nu)
model_list = [model] * nroll
model_list = [copy.copy(model) for _ in range(nroll)]
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
rollout.rollout(
@@ -561,7 +561,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
sensordata = np.empty((nroll, nstep, model.nsensordata))
control = np.random.randn(nroll, nstep, model.nu)
model_list = [model] * nroll
model_list = [copy.copy(model) for _ in range(nroll)]
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
with rollout.Rollout(nthread=num_workers) as rollout_:
@@ -610,7 +610,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
sensordata = np.empty((nroll, nstep, model.nsensordata))
control = np.random.randn(nroll, nstep, model.nu)
model_list = [model] * nroll
model_list = [copy.copy(model) for _ in range(nroll)]
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
for _ in range(2):
@@ -871,7 +871,7 @@ def py_rollout(
nstep = control.shape[1]
if isinstance(model, mujoco.MjModel):
model = [model] * nroll
model = [copy.copy(model) for _ in range(nroll)]
nstate = mujoco.mj_stateSize(model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS)