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()