diff --git a/mjx/mujoco/mjx/__init__.py b/mjx/mujoco/mjx/__init__.py index d71d6a24..7b77a2ec 100644 --- a/mjx/mujoco/mjx/__init__.py +++ b/mjx/mujoco/mjx/__init__.py @@ -30,9 +30,12 @@ from mujoco.mjx._src.forward import step from mujoco.mjx._src.inverse import inverse from mujoco.mjx._src.io import get_data from mujoco.mjx._src.io import get_data_into +from mujoco.mjx._src.io import get_state from mujoco.mjx._src.io import make_data from mujoco.mjx._src.io import put_data from mujoco.mjx._src.io import put_model +from mujoco.mjx._src.io import set_state +from mujoco.mjx._src.io import state_size from mujoco.mjx._src.passive import passive from mujoco.mjx._src.ray import ray from mujoco.mjx._src.sensor import sensor_acc diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 0a169b04..9805e878 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -1496,3 +1496,171 @@ def get_data( get_data_into(result, m, d) return result + + +_STATE_MAP = { + mujoco.mjtState.mjSTATE_TIME: 'time', + mujoco.mjtState.mjSTATE_QPOS: 'qpos', + mujoco.mjtState.mjSTATE_QVEL: 'qvel', + mujoco.mjtState.mjSTATE_ACT: 'act', + mujoco.mjtState.mjSTATE_WARMSTART: 'qacc_warmstart', + mujoco.mjtState.mjSTATE_CTRL: 'ctrl', + mujoco.mjtState.mjSTATE_QFRC_APPLIED: 'qfrc_applied', + mujoco.mjtState.mjSTATE_XFRC_APPLIED: 'xfrc_applied', + mujoco.mjtState.mjSTATE_EQ_ACTIVE: 'eq_active', + mujoco.mjtState.mjSTATE_MOCAP_POS: 'mocap_pos', + mujoco.mjtState.mjSTATE_MOCAP_QUAT: 'mocap_quat', + mujoco.mjtState.mjSTATE_USERDATA: 'userdata', + mujoco.mjtState.mjSTATE_PLUGIN: 'plugin_state', +} + + +def _state_elem_size(m: types.Model, state_enum: mujoco.mjtState) -> int: + """Returns the size of a state component.""" + if state_enum not in _STATE_MAP: + raise ValueError(f'Invalid state element {state_enum}') + name = _STATE_MAP[state_enum] + if name == 'time': + return 1 + if name in ( + 'qpos', + 'qvel', + 'act', + 'qacc_warmstart', + 'ctrl', + 'qfrc_applied', + 'eq_active', + 'mocap_pos', + 'mocap_quat', + 'userdata', + 'plugin_state', + ): + val = getattr( + m, + { + 'qpos': 'nq', + 'qvel': 'nv', + 'act': 'na', + 'qacc_warmstart': 'nv', + 'ctrl': 'nu', + 'qfrc_applied': 'nv', + 'eq_active': 'neq', + 'mocap_pos': 'nmocap', + 'mocap_quat': 'nmocap', + 'userdata': 'nuserdata', + 'plugin_state': 'npluginstate', + }[name], + ) + if name == 'mocap_pos': + val *= 3 + if name == 'mocap_quat': + val *= 4 + return val + if name == 'xfrc_applied': + return 6 * m.nbody + + raise NotImplementedError(f'state component {name} not implemented') + + +def state_size(m: types.Model, spec: Union[int, mujoco.mjtState]) -> int: + """Returns the size of a state vector for a given spec. + + Args: + m: model describing the simulation + spec: int bitmask or mjtState enum specifying which state components to + include + + Returns: + size of the state vector + """ + size = 0 + spec_int = int(spec) + for i in range(mujoco.mjtState.mjNSTATE.value): + element = mujoco.mjtState(1 << i) + if element & spec_int: + size += _state_elem_size(m, element) + return size + + +def get_state( + m: types.Model, d: types.Data, spec: Union[int, mujoco.mjtState] +) -> jax.Array: + """Gets state from mjx.Data. This is equivalent to `mujoco.mj_getState`. + + Args: + m: model describing the simulation + d: data for the simulation + spec: int bitmask or mjtState enum specifying which state components to + include + + Returns: + a flat array of state values + """ + spec_int = int(spec) + if spec_int >= (1 << mujoco.mjtState.mjNSTATE.value): + raise ValueError(f'Invalid state spec {spec}') + + state = [] + for i in range(mujoco.mjtState.mjNSTATE.value): + element = mujoco.mjtState(1 << i) + if element & spec_int: + if element not in _STATE_MAP: + raise ValueError(f'Invalid state element {element}') + name = _STATE_MAP[element] + value = getattr(d, name) + if element == mujoco.mjtState.mjSTATE_EQ_ACTIVE: + value = value.astype(jp.float32) + state.append(value.flatten()) + + return jp.concatenate(state) if state else jp.array([]) + + +def set_state( + m: types.Model, + d: types.Data, + state: jax.Array, + spec: Union[int, mujoco.mjtState], +) -> types.Data: + """Sets state in mjx.Data. This is equivalent to `mujoco.mj_setState`. + + Args: + m: model describing the simulation + d: data for the simulation + state: a flat array of state values + spec: int bitmask or mjtState enum specifying which state components to + include + + Returns: + data with state set to provided values + """ + spec_int = int(spec) + if spec_int >= (1 << mujoco.mjtState.mjNSTATE.value): + raise ValueError(f'Invalid state spec {spec}') + + expected_size = state_size(m, spec) + if state.size != expected_size: + raise ValueError( + f'state has size {state.size} but expected {expected_size}' + ) + + updates = {} + offset = 0 + for i in range(mujoco.mjtState.mjNSTATE.value): + element = mujoco.mjtState(1 << i) + if element & spec_int: + if element not in _STATE_MAP: + raise ValueError(f'Invalid state element {element}') + name = _STATE_MAP[element] + size = _state_elem_size(m, element) + value = state[offset : offset + size] + if name == 'time': + value = value[0] + else: + orig_shape = getattr(d, name).shape + value = value.reshape(orig_shape) + if element == mujoco.mjtState.mjSTATE_EQ_ACTIVE: + value = value.astype(bool) + updates[name] = value + offset += size + + return d.replace(**updates) diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index aaa2500a..4c52cdcd 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -1071,5 +1071,107 @@ class ResolveImplAndDeviceTest(parameterized.TestCase): self.assertEqual(device.platform, 'gpu') +class StateIOTest(parameterized.TestCase): + + @parameterized.parameters( + mujoco.mjtState.mjSTATE_TIME, + mujoco.mjtState.mjSTATE_QPOS, + mujoco.mjtState.mjSTATE_QVEL, + mujoco.mjtState.mjSTATE_ACT, + mujoco.mjtState.mjSTATE_WARMSTART, + mujoco.mjtState.mjSTATE_CTRL, + mujoco.mjtState.mjSTATE_QFRC_APPLIED, + mujoco.mjtState.mjSTATE_XFRC_APPLIED, + mujoco.mjtState.mjSTATE_EQ_ACTIVE, + mujoco.mjtState.mjSTATE_INTEGRATION, + ) + def test_state_size(self, spec): + m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) + mx = mjx.put_model(m) + self.assertEqual(mjx.state_size(mx, spec), mujoco.mj_stateSize(m, spec)) + + def test_get_set_state(self): + m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) + d = mujoco.MjData(m) + # give the model a little kick to get some non-zero state + d.qvel = np.random.random(m.nv) + mujoco.mj_step(m, d) + mx = mjx.put_model(m) + dx = mjx.put_data(m, d) + + # test full state + spec_full = mujoco.mjtState.mjSTATE_INTEGRATION + state = mjx.get_state(mx, dx, spec_full) + state_mj = np.empty(state.shape, dtype=np.float64) + mujoco.mj_getState(m, d, state_mj, int(spec_full)) + np.testing.assert_allclose(state, state_mj, atol=1e-6) + dx2 = mjx.set_state(mx, mjx.make_data(m), state, spec_full) + np.testing.assert_allclose(dx.qpos, dx2.qpos) + np.testing.assert_allclose(dx.qvel, dx2.qvel) + np.testing.assert_allclose(dx.act, dx2.act) + np.testing.assert_allclose(dx.qacc_warmstart, dx2.qacc_warmstart) + np.testing.assert_allclose(dx.ctrl, dx2.ctrl) + np.testing.assert_allclose(dx.qfrc_applied, dx2.qfrc_applied) + np.testing.assert_allclose(dx.xfrc_applied, dx2.xfrc_applied) + np.testing.assert_allclose(dx.eq_active, dx2.eq_active) + + # test single state + for spec in [ + mujoco.mjtState.mjSTATE_TIME, + mujoco.mjtState.mjSTATE_QPOS, + mujoco.mjtState.mjSTATE_QVEL, + mujoco.mjtState.mjSTATE_ACT, + mujoco.mjtState.mjSTATE_WARMSTART, + mujoco.mjtState.mjSTATE_CTRL, + mujoco.mjtState.mjSTATE_QFRC_APPLIED, + mujoco.mjtState.mjSTATE_XFRC_APPLIED, + mujoco.mjtState.mjSTATE_EQ_ACTIVE, + ]: + state = mjx.get_state(mx, dx, spec) + state_mj = np.empty(state.shape, dtype=np.float64) + mujoco.mj_getState(m, d, state_mj, int(spec)) + np.testing.assert_allclose(state, state_mj) + dx2 = mjx.set_state(mx, mjx.make_data(m), state, spec) + np.testing.assert_allclose( + getattr(dx, mjx_io._STATE_MAP[spec]), + getattr(dx2, mjx_io._STATE_MAP[spec]), + ) + + # test partial state + spec = ( + mujoco.mjtState.mjSTATE_QPOS + | mujoco.mjtState.mjSTATE_QVEL + ) + state = mjx.get_state(mx, dx, spec) + state_mj = np.empty(state.shape, dtype=np.float64) + mujoco.mj_getState(m, d, state_mj, int(spec)) + np.testing.assert_allclose(state, state_mj) + + # check that we only set qpos/qvel and other values are at init + dx_init = mjx.make_data(m) + dx2 = mjx.set_state(mx, dx_init, state, spec) + np.testing.assert_allclose(dx.qpos, dx2.qpos) + np.testing.assert_allclose(dx.qvel, dx2.qvel) + np.testing.assert_allclose(dx_init.time, dx2.time) + np.testing.assert_allclose(dx_init.act, dx2.act) + + def test_jit(self): + m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) + d = mujoco.MjData(m) + mx = mjx.put_model(m) + dx = mjx.put_data(m, d) + spec = mujoco.mjtState.mjSTATE_INTEGRATION + + get_state_jit = jax.jit(mjx.get_state, static_argnames='spec') + state = get_state_jit(mx, dx, spec) + state_nojit = mjx.get_state(mx, dx, spec) + np.testing.assert_allclose(state, state_nojit) + + set_state_jit = jax.jit(mjx.set_state, static_argnames='spec') + dx2 = set_state_jit(mx, dx, state, spec) + dx2_nojit = mjx.set_state(mx, dx, state, spec) + np.testing.assert_allclose(dx2.qpos, dx2_nojit.qpos) + + if __name__ == '__main__': absltest.main()