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:
@@ -39,7 +39,6 @@ Roll out open-loop trajectories from initial states, get resulting states and se
|
||||
input arguments (required):
|
||||
model instance of MjModel
|
||||
data associated instance of MjData
|
||||
nroll integer, number of initial states from which to roll out trajectories
|
||||
nstep integer, number of steps to be taken for each trajectory
|
||||
control_spec specification of controls, ncontrol = mj_stateSize(m, control_spec)
|
||||
state0 (nroll x nstate) nroll initial state vectors,
|
||||
@@ -190,7 +189,7 @@ PYBIND11_MODULE(_rollout, pymodule) {
|
||||
pymodule.def(
|
||||
"rollout",
|
||||
[](const MjModelWrapper& m, MjDataWrapper& d,
|
||||
int nroll, int nstep, unsigned int control_spec,
|
||||
int nstep, unsigned int control_spec,
|
||||
const PyCArray state0,
|
||||
std::optional<const PyCArray> warmstart0,
|
||||
std::optional<const PyCArray> control,
|
||||
@@ -201,13 +200,14 @@ PYBIND11_MODULE(_rollout, pymodule) {
|
||||
raw::MjData* data = d.get();
|
||||
|
||||
// check that some steps need to be taken, return if not
|
||||
if (nroll < 1 || nstep < 1) {
|
||||
if (nstep < 1) {
|
||||
return;
|
||||
}
|
||||
|
||||
// get sizes
|
||||
int nstate = mj_stateSize(model, mjSTATE_FULLPHYSICS);
|
||||
int ncontrol = mj_stateSize(model, control_spec);
|
||||
int nroll = state0.shape(0);
|
||||
|
||||
// get raw pointers
|
||||
mjtNum* state0_ptr = get_array_ptr(state0, "state0", nroll, 1, nstate);
|
||||
@@ -232,7 +232,6 @@ PYBIND11_MODULE(_rollout, pymodule) {
|
||||
},
|
||||
py::arg("model"),
|
||||
py::arg("data"),
|
||||
py::arg("nroll"),
|
||||
py::arg("nstep"),
|
||||
py::arg("control_spec"),
|
||||
py::arg("state0"),
|
||||
|
||||
@@ -29,7 +29,6 @@ def rollout(model: mujoco.MjModel,
|
||||
*, # require subsequent arguments to be named
|
||||
control_spec: int = mujoco.mjtState.mjSTATE_CTRL.value,
|
||||
skip_checks: bool = False,
|
||||
nroll: Optional[int] = None,
|
||||
nstep: Optional[int] = None,
|
||||
initial_warmstart: Optional[npt.ArrayLike] = None,
|
||||
state: Optional[npt.ArrayLike] = None,
|
||||
@@ -50,7 +49,6 @@ def rollout(model: mujoco.MjModel,
|
||||
([nroll or 1] x [nstep or 1] x ncontrol)
|
||||
control_spec: mjtState specification of control vectors.
|
||||
skip_checks: Whether to skip internal shape and type checks.
|
||||
nroll: Number of rollouts (inferred if unspecified).
|
||||
nstep: Number of steps in rollouts (inferred if unspecified).
|
||||
initial_warmstart: Initial qfrc_warmstart array (optional).
|
||||
([nroll or 1] x nv)
|
||||
@@ -74,7 +72,7 @@ def rollout(model: mujoco.MjModel,
|
||||
# don't allocate output arrays
|
||||
# just call rollout and return
|
||||
if skip_checks:
|
||||
_rollout.rollout(model, data, nroll, nstep, control_spec, initial_state,
|
||||
_rollout.rollout(model, data, nstep, control_spec, initial_state,
|
||||
initial_warmstart, control, state, sensordata)
|
||||
return state, sensordata
|
||||
|
||||
@@ -83,8 +81,6 @@ def rollout(model: mujoco.MjModel,
|
||||
raise ValueError('control_spec can only contain bits in mjSTATE_USER')
|
||||
|
||||
# check types
|
||||
if nroll and not isinstance(nroll, int):
|
||||
raise ValueError('nroll must be an integer')
|
||||
if nstep and not isinstance(nstep, int):
|
||||
raise ValueError('nstep must be an integer')
|
||||
_check_must_be_numeric(
|
||||
@@ -121,7 +117,7 @@ def rollout(model: mujoco.MjModel,
|
||||
_check_trailing_dimension(model.nsensordata, sensordata=sensordata)
|
||||
|
||||
# infer nroll, check for incompatibilities
|
||||
nroll = _infer_dimension(0, nroll or 1,
|
||||
nroll = _infer_dimension(0, 1,
|
||||
initial_state=initial_state,
|
||||
initial_warmstart=initial_warmstart,
|
||||
control=control,
|
||||
@@ -146,7 +142,7 @@ def rollout(model: mujoco.MjModel,
|
||||
sensordata = np.empty((nroll, nstep, model.nsensordata))
|
||||
|
||||
# call rollout
|
||||
_rollout.rollout(model, data, nroll, nstep, control_spec, initial_state,
|
||||
_rollout.rollout(model, data, nstep, control_spec, initial_state,
|
||||
initial_warmstart, control, state, sensordata)
|
||||
|
||||
# return outputs
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user