rollout accepts a list of models of length nroll
This commit is contained in:
@@ -334,6 +334,34 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_multi_model(self, model_name):
|
||||
nroll = 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):
|
||||
body = spec.bodies[1]
|
||||
assert body.name != 'world'
|
||||
body.pos = body.pos + i
|
||||
model.append(spec.compile())
|
||||
else:
|
||||
model = [spec.compile() for i in range(nroll)]
|
||||
|
||||
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)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
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_multi_rollout_fixed_ctrl_infer_from_output(self, model_name):
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
@@ -430,8 +458,9 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
def thread_initializer():
|
||||
thread_local.data = mujoco.MjData(model)
|
||||
|
||||
model_list = [model]*nroll
|
||||
def call_rollout(initial_state, control, state, sensordata):
|
||||
rollout.rollout(model, thread_local.data, initial_state, control,
|
||||
rollout.rollout(model_list, thread_local.data, initial_state, control,
|
||||
skip_checks=True,
|
||||
nstep=nstep, state=state, sensordata=sensordata)
|
||||
|
||||
@@ -677,13 +706,17 @@ def py_rollout(model, data, initial_state, control,
|
||||
control = ensure_3d(control)
|
||||
nroll = initial_state.shape[0]
|
||||
nstep = control.shape[1]
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
|
||||
if isinstance(model, mujoco.MjModel):
|
||||
model = [model]*nroll
|
||||
|
||||
nstate = mujoco.mj_stateSize(model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
sensordata = np.empty((nroll, nstep, model.nsensordata))
|
||||
sensordata = np.empty((nroll, nstep, model[0].nsensordata))
|
||||
for r in range(nroll):
|
||||
state_r, sensordata_r = one_rollout(
|
||||
model, data, initial_state[r], control[r], control_spec
|
||||
model[r], data, initial_state[r], control[r], control_spec
|
||||
)
|
||||
state[r] = state_r
|
||||
sensordata[r] = sensordata_r
|
||||
|
||||
Reference in New Issue
Block a user