Modify MJX solver to support elliptic friction cones with condim=3 and condim=4.

PiperOrigin-RevId: 700306231
Change-Id: I50fd20cd367ecab11b23c03c0e0080266cc404e6
This commit is contained in:
Taylor Howell
2024-11-26 05:16:33 -08:00
committed by Copybara-Service
parent e310c23267
commit b716837395
4 changed files with 55 additions and 2 deletions
+5
View File
@@ -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,),
+19
View File
@@ -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("""
<mujoco>
<worldbody>
<geom size="0 0 1e-5" type="plane" condim="1"/>
<body>
<freejoint/>
<geom size="0.1" condim="1"/>
</body>
</worldbody>
</mujoco>
""")
m.opt.cone = ConeType.ELLIPTIC
with self.assertRaises(NotImplementedError):
mjx.make_data(m)
if __name__ == '__main__':
absltest.main()
+12 -2
View File
@@ -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)
+19
View File
@@ -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"""
<mujoco>
<worldbody>
<geom size="0 0 1e-5" type="plane" condim="1"/>
<body pos="0 0 0.09">
<freejoint/>
<geom size="0.1" condim="{condim}"/>
</body>
</worldbody>
</mujoco>
""")
m.opt.cone = cone
solver.solve(mjx.put_model(m), mjx.put_data(m, mujoco.MjData(m)))
if __name__ == '__main__':
absltest.main()