diff --git a/doc/changelog.rst b/doc/changelog.rst index deaba59e..f2590378 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -23,6 +23,10 @@ General - Added :ref:`tendon actuator force limits` and :ref:`tendon actuator force sensor`. +MJX +^^^ +- Added tendon actuator force limits. + Bug fixes ^^^^^^^^^ - :ref:`mj_jacDot` was missing a term that accounts for the motion of the point with respect to diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py index 71a6e3f5..6ad1ae71 100644 --- a/mjx/mujoco/mjx/_src/forward.py +++ b/mjx/mujoco/mjx/_src/forward.py @@ -38,6 +38,7 @@ from mujoco.mjx._src.types import GainType from mujoco.mjx._src.types import IntegratorType from mujoco.mjx._src.types import JointType from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import TrnType # pylint: enable=g-importing-member import numpy as np @@ -178,6 +179,34 @@ def fwd_actuation(m: Model, d: Data) -> Data: jp.array(m.actuator_acc0), group_by='u', ) + + # tendon total force clamping + if np.any(m.tendon_actfrclimited): + (tendon_actfrclimited_id,) = np.nonzero(m.tendon_actfrclimited) + actuator_tendon = m.actuator_trntype == TrnType.TENDON + + force_mask = [ + actuator_tendon & (m.actuator_trnid[:, 0] == tendon_id) + for tendon_id in tendon_actfrclimited_id + ] + force_ids = np.concatenate([np.nonzero(mask)[0] for mask in force_mask]) + force_mat = np.array(force_mask)[:, force_ids] + tendon_total_force = force_mat @ force[force_ids] + + force_scaling = jp.where( + tendon_total_force < m.tendon_actfrcrange[tendon_actfrclimited_id, 0], + m.tendon_actfrcrange[tendon_actfrclimited_id, 0] / tendon_total_force, + 1, + ) + force_scaling = jp.where( + tendon_total_force > m.tendon_actfrcrange[tendon_actfrclimited_id, 1], + m.tendon_actfrcrange[tendon_actfrclimited_id, 1] / tendon_total_force, + force_scaling, + ) + + tendon_forces = force[force_ids] * (force_mat.T @ force_scaling) + force = force.at[force_ids].set(tendon_forces) + forcerange = jp.where( m.actuator_forcelimited[:, None], m.actuator_forcerange, diff --git a/mjx/mujoco/mjx/_src/forward_test.py b/mjx/mujoco/mjx/_src/forward_test.py index c6466449..efdbcd65 100644 --- a/mjx/mujoco/mjx/_src/forward_test.py +++ b/mjx/mujoco/mjx/_src/forward_test.py @@ -17,6 +17,7 @@ from absl.testing import absltest from absl.testing import parameterized import jax +from jax import numpy as jp import mujoco from mujoco import mjx from mujoco.mjx._src import test_util @@ -196,6 +197,21 @@ class ActuatorTest(parameterized.TestCase): dx = jax.jit(mjx.euler)(mx, dx) _assert_attr_eq(d, dx, 'act') + def test_tendon_force_clamp(self): + m = test_util.load_test_file('actuator/tendon_force_clamp.xml') + d = mujoco.MjData(m) + mx = mjx.put_model(m) + dx = mjx.put_data(m, d) + + dx = dx.replace(ctrl=jp.array([1.0, 1.0, 1.0, -1.0, 1.0, -20.0, 5.0, -5.0])) + dx = mjx.forward(mx, dx) + + _assert_eq( + dx.actuator_force, + jp.array([1.0, 1.0, 1.0, -1.0, 1.0, -10.0, 5.0, -5.0]), + 'actuator_force', + ) + if __name__ == '__main__': absltest.main() diff --git a/mjx/mujoco/mjx/test_data/actuator/tendon_force_clamp.xml b/mjx/mujoco/mjx/test_data/actuator/tendon_force_clamp.xml new file mode 100644 index 00000000..43d4c69b --- /dev/null +++ b/mjx/mujoco/mjx/test_data/actuator/tendon_force_clamp.xml @@ -0,0 +1,49 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +