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:
committed by
Copybara-Service
parent
e310c23267
commit
b716837395
@@ -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,),
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user