Rename rollout variable nroll to nbatch.
PiperOrigin-RevId: 718196100 Change-Id: Ic36b1a96ea1af2de351539115d4eee051d195aaf
This commit is contained in:
committed by
Copybara-Service
parent
ac76540eb8
commit
0090e1e908
@@ -185,11 +185,11 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 5 # number of rollouts
|
||||
nbatch = 5 # number of rollouts
|
||||
nstep = 1 # number of steps
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
initial_state = np.random.randn(nbatch, nstate)
|
||||
control = np.random.randn(nbatch, nstep, model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
@@ -198,108 +198,108 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_infer_nroll_initial_state(self, model_name):
|
||||
def test_infer_nbatch_initial_state(self, model_name):
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 5 # number of rollouts
|
||||
nbatch = 5 # number of rollouts
|
||||
nstep = 1 # number of steps
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
initial_state = np.random.randn(nbatch, nstate)
|
||||
control = np.random.randn(nstep, model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
control = np.tile(control, (nroll, 1, 1))
|
||||
control = np.tile(control, (nbatch, 1, 1))
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_infer_nroll_control(self, model_name):
|
||||
def test_infer_nbatch_control(self, model_name):
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 5 # number of rollouts
|
||||
nbatch = 5 # number of rollouts
|
||||
nstep = 1 # number of steps
|
||||
|
||||
initial_state = np.random.randn(nstate)
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
control = np.random.randn(nbatch, nstep, model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
initial_state = np.tile(initial_state, (nroll, 1))
|
||||
initial_state = np.tile(initial_state, (nbatch, 1))
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_infer_nroll_warmstart(self, model_name):
|
||||
def test_infer_nbatch_warmstart(self, model_name):
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 5 # number of rollouts
|
||||
nbatch = 5 # number of rollouts
|
||||
nstep = 1 # number of steps
|
||||
|
||||
initial_state = np.random.randn(nstate)
|
||||
control = np.random.randn(nstep, model.nu)
|
||||
initial_warmstart = np.tile(data.qacc_warmstart.copy(), (nroll, 1))
|
||||
initial_warmstart = np.tile(data.qacc_warmstart.copy(), (nbatch, 1))
|
||||
state, sensordata = rollout.rollout(
|
||||
model, data, initial_state, control, initial_warmstart=initial_warmstart
|
||||
)
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
initial_state = np.tile(initial_state, (nroll, 1))
|
||||
control = np.tile(control, (nroll, 1, 1))
|
||||
initial_state = np.tile(initial_state, (nbatch, 1))
|
||||
control = np.tile(control, (nbatch, 1, 1))
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_infer_nroll_state(self, model_name):
|
||||
def test_infer_nbatch_state(self, model_name):
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 5 # number of rollouts
|
||||
nbatch = 5 # number of rollouts
|
||||
nstep = 1 # number of steps
|
||||
|
||||
initial_state = np.random.randn(nstate)
|
||||
control = np.random.randn(nstep, model.nu)
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
state = np.empty((nbatch, nstep, nstate))
|
||||
state, sensordata = rollout.rollout(
|
||||
model, data, initial_state, control, state=state
|
||||
)
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
initial_state = np.tile(initial_state, (nroll, 1))
|
||||
control = np.tile(control, (nroll, 1, 1))
|
||||
initial_state = np.tile(initial_state, (nbatch, 1))
|
||||
control = np.tile(control, (nbatch, 1, 1))
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_infer_nroll_sensordata(self, model_name):
|
||||
def test_infer_nbatch_sensordata(self, model_name):
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 5 # number of rollouts
|
||||
nbatch = 5 # number of rollouts
|
||||
nstep = 1 # number of steps
|
||||
|
||||
initial_state = np.random.randn(nstate)
|
||||
control = np.random.randn(nstep, model.nu)
|
||||
sensordata = np.empty((nroll, nstep, model.nsensordata))
|
||||
sensordata = np.empty((nbatch, nstep, model.nsensordata))
|
||||
state, sensordata = rollout.rollout(
|
||||
model, data, initial_state, control, sensordata=sensordata
|
||||
)
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
initial_state = np.tile(initial_state, (nroll, 1))
|
||||
control = np.tile(control, (nroll, 1, 1))
|
||||
initial_state = np.tile(initial_state, (nbatch, 1))
|
||||
control = np.tile(control, (nbatch, 1, 1))
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
@@ -310,13 +310,13 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 1 # number of rollouts
|
||||
nbatch = 1 # number of rollouts
|
||||
nstep = 3 # number of steps
|
||||
|
||||
initial_state = np.random.randn(nstate)
|
||||
control = np.random.randn(model.nu)
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
sensordata = np.empty((nroll, nstep, model.nsensordata))
|
||||
state = np.empty((nbatch, nstep, nstate))
|
||||
sensordata = np.empty((nbatch, nstep, model.nsensordata))
|
||||
rollout.rollout(
|
||||
model, data, initial_state, control, state=state, sensordata=sensordata
|
||||
)
|
||||
@@ -332,11 +332,11 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 2 # number of initial states
|
||||
nbatch = 2 # number of initial states
|
||||
nstep = 3 # number of timesteps
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
initial_state = np.random.randn(nbatch, nstate)
|
||||
control = np.random.randn(nbatch, nstep, model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
@@ -345,26 +345,26 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_multi_model(self, model_name):
|
||||
nroll = 3 # number of initial states and models
|
||||
nbatch = 3 # number of initial states and models
|
||||
nstep = 3 # number of timesteps
|
||||
|
||||
spec = mujoco.MjSpec.from_string(ALL_MODELS[model_name])
|
||||
|
||||
if len(spec.bodies) > 1:
|
||||
model = []
|
||||
for i in range(nroll):
|
||||
for i in range(nbatch):
|
||||
body = spec.bodies[1]
|
||||
assert body.name != 'world'
|
||||
body.pos = body.pos + i
|
||||
model.append(spec.compile())
|
||||
else:
|
||||
model = [spec.compile() for _ in range(nroll)]
|
||||
model = [spec.compile() for _ in range(nbatch)]
|
||||
|
||||
nstate = mujoco.mj_stateSize(model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model[0])
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
control = np.random.randn(nroll, nstep, model[0].nu)
|
||||
initial_state = np.random.randn(nbatch, nstate)
|
||||
control = np.random.randn(nbatch, nstep, model[0].nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
@@ -377,12 +377,12 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 2 # number of rollouts
|
||||
nbatch = 2 # number of rollouts
|
||||
nstep = 3 # number of timesteps
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
control = np.random.randn(nroll, 1, model.nu)
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
initial_state = np.random.randn(nbatch, nstate)
|
||||
control = np.random.randn(nbatch, 1, model.nu)
|
||||
state = np.empty((nbatch, nstep, nstate))
|
||||
state, sensordata = rollout.rollout(
|
||||
model, data, initial_state, control, state=state
|
||||
)
|
||||
@@ -398,10 +398,10 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 4 # number of rollouts
|
||||
nbatch = 4 # number of rollouts
|
||||
nstep = 3 # number of timesteps
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
initial_state = np.random.randn(nbatch, nstate)
|
||||
|
||||
control_spec = (
|
||||
mujoco.mjtState.mjSTATE_CTRL
|
||||
@@ -409,7 +409,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
| mujoco.mjtState.mjSTATE_XFRC_APPLIED
|
||||
)
|
||||
ncontrol = mujoco.mj_stateSize(model, control_spec)
|
||||
control = np.random.randn(nroll, nstep, ncontrol)
|
||||
control = np.random.randn(nbatch, nstep, ncontrol)
|
||||
|
||||
state, sensordata = rollout.rollout(
|
||||
model, data, initial_state, control, control_spec=control_spec
|
||||
@@ -426,8 +426,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 4 # number of rollouts
|
||||
initial_state = np.empty((nroll, nstate))
|
||||
nbatch = 4 # number of rollouts
|
||||
initial_state = np.empty((nbatch, nstate))
|
||||
|
||||
# get diverging (0, 2) and non-diverging (1, 3) states
|
||||
mujoco.mj_getState(
|
||||
@@ -446,7 +446,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
|
||||
nstep = 10000 # divergence after ~15s, timestep = 2e-3
|
||||
|
||||
state = np.random.randn(nroll, nstep, nstate)
|
||||
state = np.random.randn(nbatch, nstep, nstate)
|
||||
|
||||
rollout.rollout(model, data, initial_state, state=state)
|
||||
|
||||
@@ -464,19 +464,19 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML)
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
num_workers = 32
|
||||
nroll = 100
|
||||
nbatch = 100
|
||||
nstep = 5
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
sensordata = np.empty((nroll, nstep, model.nsensordata))
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
initial_state = np.random.randn(nbatch, nstate)
|
||||
state = np.empty((nbatch, nstep, nstate))
|
||||
sensordata = np.empty((nbatch, nstep, model.nsensordata))
|
||||
control = np.random.randn(nbatch, nstep, model.nu)
|
||||
|
||||
thread_local = threading.local()
|
||||
|
||||
def thread_initializer():
|
||||
thread_local.data = mujoco.MjData(model)
|
||||
|
||||
model_list = [copy.copy(model) for _ in range(nroll)]
|
||||
model_list = [copy.copy(model) for _ in range(nbatch)]
|
||||
|
||||
def call_rollout(initial_state, control, state, sensordata):
|
||||
rollout.rollout(
|
||||
@@ -490,7 +490,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
sensordata=sensordata,
|
||||
)
|
||||
|
||||
n = nroll // num_workers # integer division
|
||||
n = nbatch // num_workers # integer division
|
||||
chunks = [] # a list of tuples, one per worker
|
||||
for i in range(num_workers - 1):
|
||||
chunks.append((
|
||||
@@ -526,14 +526,14 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML)
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
num_workers = 32
|
||||
nroll = 100
|
||||
nbatch = 100
|
||||
nstep = 5
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
sensordata = np.empty((nroll, nstep, model.nsensordata))
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
initial_state = np.random.randn(nbatch, nstate)
|
||||
state = np.empty((nbatch, nstep, nstate))
|
||||
sensordata = np.empty((nbatch, nstep, model.nsensordata))
|
||||
control = np.random.randn(nbatch, nstep, model.nu)
|
||||
|
||||
model_list = [copy.copy(model) for _ in range(nroll)]
|
||||
model_list = [copy.copy(model) for _ in range(nbatch)]
|
||||
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
|
||||
|
||||
rollout.rollout(
|
||||
@@ -555,14 +555,14 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML)
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
num_workers = 32
|
||||
nroll = 100
|
||||
nbatch = 100
|
||||
nstep = 5
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
sensordata = np.empty((nroll, nstep, model.nsensordata))
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
initial_state = np.random.randn(nbatch, nstate)
|
||||
state = np.empty((nbatch, nstep, nstate))
|
||||
sensordata = np.empty((nbatch, nstep, model.nsensordata))
|
||||
control = np.random.randn(nbatch, nstep, model.nu)
|
||||
|
||||
model_list = [copy.copy(model) for _ in range(nroll)]
|
||||
model_list = [copy.copy(model) for _ in range(nbatch)]
|
||||
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
|
||||
|
||||
with rollout.Rollout(nthread=num_workers) as rollout_:
|
||||
@@ -604,14 +604,14 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML)
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
num_workers = 32
|
||||
nroll = 100
|
||||
nbatch = 100
|
||||
nstep = 5
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
sensordata = np.empty((nroll, nstep, model.nsensordata))
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
initial_state = np.random.randn(nbatch, nstate)
|
||||
state = np.empty((nbatch, nstep, nstate))
|
||||
sensordata = np.empty((nbatch, nstep, model.nsensordata))
|
||||
control = np.random.randn(nbatch, nstep, model.nu)
|
||||
|
||||
model_list = [copy.copy(model) for _ in range(nroll)]
|
||||
model_list = [copy.copy(model) for _ in range(nbatch)]
|
||||
data_list = [mujoco.MjData(model) for _ in range(num_workers)]
|
||||
|
||||
for _ in range(2):
|
||||
@@ -699,11 +699,11 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 1
|
||||
nbatch = 1
|
||||
nstep = 3
|
||||
|
||||
initial_state = np.zeros((nroll, nstate))
|
||||
ctrl = np.zeros((nroll, nstep, model.nu))
|
||||
initial_state = np.zeros((nbatch, nstate))
|
||||
ctrl = np.zeros((nbatch, nstep, model.nu))
|
||||
|
||||
model.opt.solver = 10 # invalid solver type
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
@@ -716,9 +716,9 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 1
|
||||
nbatch = 1
|
||||
|
||||
initial_state = np.zeros((nroll, nstate))
|
||||
initial_state = np.zeros((nbatch, nstate))
|
||||
|
||||
control = 'string'
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
@@ -737,31 +737,31 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 1
|
||||
nbatch = 1
|
||||
nstep = 3
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate + 1)
|
||||
initial_state = np.random.randn(nbatch, nstate + 1)
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'trailing dimension of initial_state must be 6, got 7'
|
||||
):
|
||||
rollout.rollout(model, data, initial_state)
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
initial_state = np.random.randn(nbatch, nstate)
|
||||
control = np.random.randn(1, nstep, model.nu + 1)
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'trailing dimension of control must be 2, got 3'
|
||||
):
|
||||
rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
state = np.random.randn(nroll, nstep + 1, nstate) # incompatible nstep
|
||||
control = np.random.randn(nbatch, nstep, model.nu)
|
||||
state = np.random.randn(nbatch, nstep + 1, nstate) # incompatible nstep
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'dimension 1 inferred as 3 but state has 4'
|
||||
):
|
||||
rollout.rollout(model, data, initial_state, control, state=state)
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
initial_state = np.random.randn(nbatch, nstate)
|
||||
control = np.random.randn(nbatch, nstep, model.nu)
|
||||
bad_spec = mujoco.mjtState.mjSTATE_ACT
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'control_spec can only contain bits in mjSTATE_USER'
|
||||
@@ -946,17 +946,17 @@ def py_rollout(
|
||||
):
|
||||
initial_state = ensure_2d(initial_state)
|
||||
control = ensure_3d(control)
|
||||
nroll = initial_state.shape[0]
|
||||
nbatch = initial_state.shape[0]
|
||||
nstep = control.shape[1]
|
||||
|
||||
if isinstance(model, mujoco.MjModel):
|
||||
model = [copy.copy(model) for _ in range(nroll)]
|
||||
model = [copy.copy(model) for _ in range(nbatch)]
|
||||
|
||||
nstate = mujoco.mj_stateSize(model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
sensordata = np.empty((nroll, nstep, model[0].nsensordata))
|
||||
for r in range(nroll):
|
||||
state = np.empty((nbatch, nstep, nstate))
|
||||
sensordata = np.empty((nbatch, nstep, model[0].nsensordata))
|
||||
for r in range(nbatch):
|
||||
state_r, sensordata_r = one_rollout(
|
||||
model[r], data, initial_state[r], control[r], control_spec
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user