diff --git a/doc/changelog.rst b/doc/changelog.rst
index 471003ff..c97f97e6 100644
--- a/doc/changelog.rst
+++ b/doc/changelog.rst
@@ -16,6 +16,7 @@ General
Bug fixes
^^^^^^^^^
- Fixed a bug where ``actuator_force`` was not set in MJX (:github:issue:`2068`).
+- Fixed bug where MJX data tendon fields were incorrect after calling ``mjx.put_data``.
Version 3.2.3 (Sep 16, 2024)
diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py
index a436d4ac..c848ba9c 100644
--- a/mjx/mujoco/mjx/_src/io_test.py
+++ b/mjx/mujoco/mjx/_src/io_test.py
@@ -68,8 +68,8 @@ _MULTIPLE_CONSTRAINTS = """
-
-
+
+
@@ -78,6 +78,12 @@ _MULTIPLE_CONSTRAINTS = """
+
+
+
+
+
+
"""
@@ -85,8 +91,11 @@ _MULTIPLE_CONSTRAINTS = """
class ModelIOTest(parameterized.TestCase):
"""IO tests for mjx.Model."""
- def test_put_model(self):
- m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONVEX_OBJECTS)
+ @parameterized.parameters(
+ _MULTIPLE_CONVEX_OBJECTS, _MULTIPLE_CONSTRAINTS
+ )
+ def test_put_model(self, xml):
+ m = mujoco.MjModel.from_xml_string(xml)
mx = mjx.put_model(m)
def assert_not_weak_type(x):
if isinstance(x, jax.Array):
@@ -126,6 +135,10 @@ class ModelIOTest(parameterized.TestCase):
np.testing.assert_allclose(mx.actuator_biastype, m.actuator_biastype)
np.testing.assert_allclose(mx.actuator_trnid, m.actuator_trnid)
+ np.testing.assert_equal(mx.wrap_type, m.wrap_type)
+ np.testing.assert_equal(mx.wrap_objid, m.wrap_objid)
+ np.testing.assert_equal(mx.wrap_prm, m.wrap_prm)
+
def test_fluid_params(self):
"""Test that has_fluid_params is set when fluid params are present."""
m = mjx.put_model(
@@ -338,6 +351,13 @@ class DataIOTest(parameterized.TestCase):
np.testing.assert_allclose(dx.geom_xmat.reshape((3, 9)), d.geom_xmat)
np.testing.assert_allclose(dx.site_xmat.reshape((1, 9)), d.site_xmat)
+ # tendon data is correct
+ np.testing.assert_allclose(dx.ten_length, d.ten_length)
+ np.testing.assert_equal(dx.ten_wrapadr, np.zeros((1,)))
+ np.testing.assert_equal(dx.ten_wrapnum, np.zeros((1,)))
+ np.testing.assert_equal(dx.wrap_obj, np.zeros((2, 2)))
+ np.testing.assert_equal(dx.wrap_xpos, np.zeros((2, 6)))
+
# efc_ are also shape transformed and padded
self.assertEqual(dx.efc_J.shape, (45, 8)) # nefc, nv
d_efc_j = d.efc_J.reshape((-1, 8))
diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py
index 17cd499b..dceed8f5 100644
--- a/mjx/mujoco/mjx/_src/types.py
+++ b/mjx/mujoco/mjx/_src/types.py
@@ -1029,9 +1029,9 @@ class Model(PyTreeNode):
tendon_length0: jax.Array
tendon_invweight0: jax.Array
tendon_hasfrictionloss: np.ndarray = _restricted_to('mjx')
- wrap_type: np.ndarray = _restricted_to('mujoco')
- wrap_objid: np.ndarray = _restricted_to('mujoco')
- wrap_prm: np.ndarray = _restricted_to('mujoco')
+ wrap_type: np.ndarray
+ wrap_objid: np.ndarray
+ wrap_prm: np.ndarray
actuator_trntype: np.ndarray
actuator_dyntype: np.ndarray
actuator_gaintype: np.ndarray
@@ -1297,15 +1297,15 @@ class Data(PyTreeNode):
flexedge_J_colind: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name
flexedge_J: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name
flexedge_length: jax.Array = _restricted_to('mujoco')
- ten_wrapadr: jax.Array = _restricted_to('mujoco')
- ten_wrapnum: jax.Array = _restricted_to('mujoco')
+ ten_wrapadr: jax.Array
+ ten_wrapnum: jax.Array
ten_J_rownnz: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name
ten_J_rowadr: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name
ten_J_colind: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name
ten_J: jax.Array # pylint:disable=invalid-name
ten_length: jax.Array
- wrap_obj: jax.Array = _restricted_to('mujoco')
- wrap_xpos: jax.Array = _restricted_to('mujoco')
+ wrap_obj: jax.Array
+ wrap_xpos: jax.Array
actuator_length: jax.Array
actuator_moment: jax.Array
crb: jax.Array