From ab3954d892cef084299b5535c64fc346d1b52760 Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Fri, 20 Sep 2024 09:15:45 -0700 Subject: [PATCH] Set `actuator_force` in MJX. Fixes #2068. PiperOrigin-RevId: 676867005 Change-Id: I80514a050ff40cb6d43a95b74c4a36c4991db588 --- doc/changelog.rst | 5 +++++ mjx/mujoco/mjx/_src/forward.py | 4 +++- mjx/mujoco/mjx/_src/forward_test.py | 8 +++++++- 3 files changed, 15 insertions(+), 2 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index f4d1828f..471003ff 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -13,6 +13,11 @@ General `__ for examples of flex objects that previously required these plugins. +Bug fixes +^^^^^^^^^ +- Fixed a bug where ``actuator_force`` was not set in MJX (:github:issue:`2068`). + + Version 3.2.3 (Sep 16, 2024) ---------------------------- diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py index 456e748c..c187d685 100644 --- a/mjx/mujoco/mjx/_src/forward.py +++ b/mjx/mujoco/mjx/_src/forward.py @@ -192,7 +192,9 @@ def fwd_actuation(m: Model, d: Data) -> Data: actfrcrange = actfrcrange[m.dof_jntid] qfrc_actuator = jp.clip(qfrc_actuator, actfrcrange[:, 0], actfrcrange[:, 1]) - d = d.replace(act_dot=act_dot, qfrc_actuator=qfrc_actuator) + d = d.replace( + act_dot=act_dot, qfrc_actuator=qfrc_actuator, actuator_force=force + ) return d diff --git a/mjx/mujoco/mjx/_src/forward_test.py b/mjx/mujoco/mjx/_src/forward_test.py index a93f04e0..7c669062 100644 --- a/mjx/mujoco/mjx/_src/forward_test.py +++ b/mjx/mujoco/mjx/_src/forward_test.py @@ -52,9 +52,15 @@ class ForwardTest(absltest.TestCase): mx = mjx.put_model(m) # fwd_actuation - dx = jax.jit(mjx.fwd_actuation)(mx, mjx.put_data(m, d)) + dx = mjx.put_data(m, d).replace( + act_dot=np.zeros_like(d.act_dot), + qfrc_actuator=np.zeros_like(d.qfrc_actuator), + actuator_force=np.zeros_like(d.actuator_force), + ) + dx = jax.jit(mjx.fwd_actuation)(mx, dx) _assert_attr_eq(d, dx, 'act_dot') _assert_attr_eq(d, dx, 'qfrc_actuator') + _assert_attr_eq(d, dx, 'actuator_force') # fwd_accleration (fwd_position and fwd_velocity already tested elsewhere) dx = jax.jit(mjx.fwd_acceleration)(mx, mjx.put_data(m, d))