Implement equivalents to mj_getState, mj_setState and mj_stateSize in mjx
PiperOrigin-RevId: 823158825 Change-Id: I06bbf603f61ef2de8f83f32a211e8d468e0b7c5c
This commit is contained in:
committed by
Copybara-Service
parent
99c18f07d3
commit
7bf065c75b
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user