diff --git a/doc/changelog.rst b/doc/changelog.rst index c9f62f1a..e212eb61 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -14,7 +14,7 @@ MJX Bug fixes ^^^^^^^^^ 1. Fix typos and supported fields in the docs (fixes :github:issue:`1105` and :github:issue:`1106`). - +2. Fix bug where mixed `jnt_limited` joints are not being constrained correctly. Version 3.0.0 (October 18, 2023) -------------------------------- diff --git a/mjx/mujoco/mjx/_src/constraint.py b/mjx/mujoco/mjx/_src/constraint.py index bd0bbaa6..3f020198 100644 --- a/mjx/mujoco/mjx/_src/constraint.py +++ b/mjx/mujoco/mjx/_src/constraint.py @@ -228,13 +228,19 @@ def _instantiate_limit(m: Model, d: Data) -> _Efc: return jp.stack(js), jp.stack(rs), jp.stack(arefs) + jnt_range = jp.where( + m.jnt_limited[:, None], + m.jnt_range, + jp.array([-jp.inf, jp.inf]), + ) + j, r, aref = scan.flat( m, fn, 'jjjjjqvv', 'jjj', m.jnt_type, - m.jnt_range, + jnt_range, m.jnt_solref, m.jnt_solimp, m.jnt_margin, diff --git a/mjx/mujoco/mjx/_src/constraint_test.py b/mjx/mujoco/mjx/_src/constraint_test.py index ce289df7..8bc77042 100644 --- a/mjx/mujoco/mjx/_src/constraint_test.py +++ b/mjx/mujoco/mjx/_src/constraint_test.py @@ -24,6 +24,7 @@ from mujoco.mjx._src import constraint from mujoco.mjx._src import test_util # pylint: disable=g-importing-member from mujoco.mjx._src.types import DisableBit +from mujoco.mjx._src.types import SolverType # pylint: enable=g-importing-member import numpy as np @@ -79,6 +80,39 @@ class ConstraintTest(parameterized.TestCase): fname, ) + _JNT_RANGE = """ + + + + + + + + + + + + + """ + + def test_jnt_range(self): + """Tests that mixed joint ranges are respected.""" + m = mujoco.MjModel.from_xml_string(self._JNT_RANGE) + m.opt.solver = SolverType.CG.value + d = mujoco.MjData(m) + d.qpos = np.array([2.0, 15.0]) + + mx = mjx.device_put(m) + dx = mjx.device_put(d) + efc = jax.jit(constraint._instantiate_limit)(mx, dx) + + # first joint is outside the joint range + np.testing.assert_array_almost_equal(efc.J[0, 0], -1.0) + + # second joint does not hit joint range + np.testing.assert_array_almost_equal(efc.aref[1], 0.0) + def test_disable_refsafe(self): m = test_util.load_test_file('ant.xml') diff --git a/mjx/mujoco/mjx/test_data/cherry_pendulum.xml b/mjx/mujoco/mjx/test_data/cherry_pendulum.xml index bfe2db9b..ea9082de 100644 --- a/mjx/mujoco/mjx/test_data/cherry_pendulum.xml +++ b/mjx/mujoco/mjx/test_data/cherry_pendulum.xml @@ -5,7 +5,7 @@ - +