diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py
index fedb014e..860c8a44 100644
--- a/mjx/mujoco/mjx/_src/io.py
+++ b/mjx/mujoco/mjx/_src/io.py
@@ -224,6 +224,11 @@ def make_data(
efc_address=efc_address,
)
+ if m.opt.cone == types.ConeType.ELLIPTIC and np.any(contact.dim == 1):
+ raise NotImplementedError(
+ 'condim=1 with ConeType.ELLIPTIC not implemented.'
+ )
+
zero_fields = {
'solver_niter': (int,),
'time': (float,),
diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py
index cbb10ebe..289c3890 100644
--- a/mjx/mujoco/mjx/_src/io_test.py
+++ b/mjx/mujoco/mjx/_src/io_test.py
@@ -20,6 +20,7 @@ import jax
from jax import numpy as jp
import mujoco
from mujoco import mjx
+from mujoco.mjx._src.types import ConeType
import numpy as np
@@ -495,5 +496,23 @@ class DataIOTest(parameterized.TestCase):
# calling make_data. they should be interchangeable for jax functions:
step_fn_jit(mjx.make_data(m))
+ def test_contact_elliptic_condim1(self):
+ """Test that condim=1 with ConeType.ELLIPTIC is not implemented."""
+ m = mujoco.MjModel.from_xml_string("""
+
+
+
+
+
+
+
+
+
+ """)
+ m.opt.cone = ConeType.ELLIPTIC
+ with self.assertRaises(NotImplementedError):
+ mjx.make_data(m)
+
+
if __name__ == '__main__':
absltest.main()
diff --git a/mjx/mujoco/mjx/_src/solver.py b/mjx/mujoco/mjx/_src/solver.py
index 4c1bcbf9..52db56f7 100644
--- a/mjx/mujoco/mjx/_src/solver.py
+++ b/mjx/mujoco/mjx/_src/solver.py
@@ -279,7 +279,12 @@ def _update_constraint(m: Model, d: Data, ctx: _Context) -> _Context:
friction = d.contact.friction[d.contact.dim > 1]
efc_address = d.contact.efc_address[d.contact.dim > 1]
dim = d.contact.dim[d.contact.dim > 1]
- slice_fn = jax.vmap(lambda x: jax.lax.dynamic_slice(ctx.Jaref, (x,), (6,)))
+ # to prevent out of range append zeros to ctx.Jaref
+ slice_fn = jax.vmap(
+ lambda x: jax.lax.dynamic_slice(
+ jp.concatenate((ctx.Jaref, jp.zeros((3)))), (x,), (6,)
+ )
+ )
u = slice_fn(efc_address) * ctx.fri
mu, n, t = ctx.fri[:, 0], u[:, 0], jax.vmap(math.norm)(u[:, 1:])
@@ -433,7 +438,12 @@ def _linesearch(m: Model, d: Data, ctx: _Context) -> _Context:
quad = quad.at[jp.array(efc_con)].add(quad[jp.array(efc_fri)])
# rescale to make primal cone circular
- jv_fn = jax.vmap(lambda x: jax.lax.dynamic_slice(jv, (x,), (6,)))
+ # to prevent out of range append zeros to jv
+ jv_fn = jax.vmap(
+ lambda x: jax.lax.dynamic_slice(
+ jp.concatenate((jv, jp.zeros(3))), (x,), (6,)
+ )
+ )
efc_elliptic = d.contact.efc_address[mask]
v = jv_fn(efc_elliptic) * ctx.fri
uu = jp.sum(ctx.u[:, 1:] * ctx.u[:, 1:], axis=1)
diff --git a/mjx/mujoco/mjx/_src/solver_test.py b/mjx/mujoco/mjx/_src/solver_test.py
index 989e4015..57d4a116 100644
--- a/mjx/mujoco/mjx/_src/solver_test.py
+++ b/mjx/mujoco/mjx/_src/solver_test.py
@@ -21,6 +21,7 @@ import mujoco
from mujoco import mjx
from mujoco.mjx._src import solver
from mujoco.mjx._src import test_util
+from mujoco.mjx._src.types import ConeType
import numpy as np
@@ -146,6 +147,24 @@ class SolverTest(parameterized.TestCase):
nnz = dx.efc_J.any(axis=1)
_assert_eq(d.efc_force, dx.efc_force[nnz], 'efc_force')
+ # TODO(taylorhowell): condim=1 with ConeType.ELLIPTIC
+ @parameterized.product(condim=(3, 4, 6), cone=tuple(ConeType))
+ def test_condim(self, condim, cone):
+ """Test contact dimension."""
+ m = mujoco.MjModel.from_xml_string(f"""
+
+
+
+
+
+
+
+
+
+ """)
+ m.opt.cone = cone
+ solver.solve(mjx.put_model(m), mjx.put_data(m, mujoco.MjData(m)))
+
if __name__ == '__main__':
absltest.main()