diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 1eada20e..bf64260c 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -393,6 +393,7 @@ def _make_data_public_fields(m: types.Model) -> Dict[str, Any]: 'time': (float_,), 'qvel': (m.nv, float_), 'act': (m.na, float_), + 'plugin_state': (m.npluginstate, float_), 'qacc_warmstart': (m.nv, float_), 'ctrl': (m.nu, float_), 'qfrc_applied': (m.nv, float_), diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index a5843cfd..42af00a7 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -579,6 +579,7 @@ class ModelC(PyTreeNode): sensor_plugin: jax.Array sensor_intprm: jax.Array plugin: jax.Array + plugin_stateadr: jax.Array class ModelJAX(PyTreeNode): @@ -636,6 +637,7 @@ class Model(PyTreeNode): ngravcomp: int nuserdata: int nsensordata: int + npluginstate: int opt: Option stat: Statistic qpos0: jax.Array @@ -1084,6 +1086,7 @@ class Data(PyTreeNode): qvel: jax.Array act: jax.Array qacc_warmstart: jax.Array + plugin_state: jax.Array # control: ctrl: jax.Array qfrc_applied: jax.Array