diff --git a/mjx/mujoco/mjx/_src/inverse.py b/mjx/mujoco/mjx/_src/inverse.py
index 5df72de9..086dd101 100644
--- a/mjx/mujoco/mjx/_src/inverse.py
+++ b/mjx/mujoco/mjx/_src/inverse.py
@@ -93,11 +93,15 @@ def inverse(m: Model, d: Data) -> Data:
d = discrete_acc(m, d)
d = inv_constraint(m, d)
- d = smooth.rne(m, d, flg_acc=True)
+ d = smooth.rne(m, d)
+ d = smooth.tendon_bias(m, d)
d = sensor.sensor_acc(m, d)
qfrc_inverse = (
- d.qfrc_bias + m.dof_armature * d.qacc - d.qfrc_passive - d.qfrc_constraint
+ d.qfrc_bias
+ + support.mul_m(m, d, d.qacc)
+ - d.qfrc_passive
+ - d.qfrc_constraint
)
if m.opt.enableflags & EnableBit.INVDISCRETE:
diff --git a/mjx/mujoco/mjx/_src/inverse_test.py b/mjx/mujoco/mjx/_src/inverse_test.py
index 227406e3..63720a15 100644
--- a/mjx/mujoco/mjx/_src/inverse_test.py
+++ b/mjx/mujoco/mjx/_src/inverse_test.py
@@ -19,6 +19,7 @@ from jax import numpy as jp
import mujoco
from mujoco import mjx
from mujoco.mjx._src import support
+from mujoco.mjx._src import test_util
import numpy as np
# tolerance for difference between MuJoCo and MJX calculations - mostly
@@ -109,6 +110,41 @@ class InverseTest(parameterized.TestCase):
self.assertLess(fwdinv1, 1.0e-3)
_assert_eq(dxinv.qacc, dx.qacc, 'qacc')
+ def test_inverse_tendon_armature(self):
+ m = test_util.load_test_file('tendon/armature.xml')
+
+ d = mujoco.MjData(m)
+ d.qvel = np.random.uniform(low=-0.01, high=0.01, size=d.qvel.shape)
+ d.ctrl = np.random.uniform(low=-0.01, high=0.01, size=d.ctrl.shape)
+ d.qfrc_applied = np.random.uniform(
+ low=-0.01, high=0.01, size=d.qfrc_applied.shape
+ )
+ d.xfrc_applied = np.random.uniform(
+ low=-0.01, high=0.01, size=d.xfrc_applied.shape
+ )
+ mujoco.mj_step(m, d, 10)
+
+ mx = mjx.put_model(m)
+ dx = mjx.put_data(m, d)
+
+ dx = mjx.forward(mx, dx)
+ dxinv = mjx.inverse(mx, dx)
+
+ fwdinv0 = jp.linalg.norm(
+ dxinv.qfrc_constraint - dx.qfrc_constraint, ord=np.inf
+ )
+ fwdinv1 = jp.linalg.norm(
+ dxinv.qfrc_inverse
+ - (
+ dx.qfrc_applied + dx.qfrc_actuator + support.xfrc_accumulate(mx, dx)
+ ),
+ ord=np.inf,
+ )
+
+ self.assertLess(fwdinv0, 1.0e-3)
+ self.assertLess(fwdinv1, 1.0e-3)
+ _assert_eq(dxinv.qacc, dx.qacc, 'qacc')
+
if __name__ == '__main__':
absltest.main()
diff --git a/mjx/mujoco/mjx/_src/smooth_test.py b/mjx/mujoco/mjx/_src/smooth_test.py
index 506d8109..64c4d32b 100644
--- a/mjx/mujoco/mjx/_src/smooth_test.py
+++ b/mjx/mujoco/mjx/_src/smooth_test.py
@@ -321,32 +321,7 @@ class TendonTest(parameterized.TestCase):
@parameterized.parameters(JacobianType.DENSE, JacobianType.SPARSE)
def test_tendon_armature(self, jacobian):
"""Tests MJX tendon armature matches MuJoCo."""
- m = mujoco.MjModel.from_xml_string("""
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
- """)
+ m = test_util.load_test_file('tendon/armature.xml')
m.opt.jacobian = jacobian
d = mujoco.MjData(m)
mujoco.mj_resetDataKeyframe(m, d, 0)
diff --git a/mjx/mujoco/mjx/test_data/tendon/armature.xml b/mjx/mujoco/mjx/test_data/tendon/armature.xml
new file mode 100644
index 00000000..0b416feb
--- /dev/null
+++ b/mjx/mujoco/mjx/test_data/tendon/armature.xml
@@ -0,0 +1,24 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+