Remove nroll argument from rollout

nroll can always be implied from the other arguments of rollout.

Add tests for nroll inference.
This commit is contained in:
Levi Burner
2024-11-24 15:34:35 -05:00
parent 7e46e21ef3
commit 96a8f002bc
4 changed files with 112 additions and 12 deletions
+105 -1
View File
@@ -192,6 +192,110 @@ 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_infer_nroll_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
nstep = 1 # number of steps
initial_state = np.random.randn(nroll, 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))
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):
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
nstep = 1 # number of steps
initial_state = np.random.randn(nstate)
control = np.random.randn(nroll, 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))
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):
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
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))
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))
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):
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
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, 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))
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):
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
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))
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))
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_one_rollout_fixed_ctrl(self, model_name):
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
@@ -328,7 +432,7 @@ class MuJoCoRolloutTest(parameterized.TestCase):
def call_rollout(initial_state, control, state, sensordata):
rollout.rollout(model, thread_local.data, initial_state, control,
skip_checks=True, nroll=initial_state.shape[0],
skip_checks=True,
nstep=nstep, state=state, sensordata=sensordata)
n = nroll // num_workers # integer division