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 @@
-
+