Implement equivalents to mj_getState, mj_setState and mj_stateSize in mjx

PiperOrigin-RevId: 823158825
Change-Id: I06bbf603f61ef2de8f83f32a211e8d468e0b7c5c
This commit is contained in:
Google DeepMind
2025-10-23 13:08:37 -07:00
committed by Copybara-Service
parent 99c18f07d3
commit 7bf065c75b
3 changed files with 273 additions and 0 deletions
+3
View File
@@ -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
+168
View File
@@ -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)
+102
View File
@@ -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()