Make unsupported feature checking optional in put_model.

PiperOrigin-RevId: 654141976
Change-Id: I35b5f1e4c5b07e32075e847d4a757698a0b603c1
This commit is contained in:
Saran Tunyasuvunakool
2024-07-19 15:22:40 -07:00
committed by Copybara-Service
parent db1d51a1f5
commit d49d4bf64c
+8 -6
View File
@@ -69,17 +69,19 @@ def _make_statistic(s: mujoco.MjStatistic) -> types.Statistic:
) )
def put_model(m: mujoco.MjModel, device=None) -> types.Model: def put_model(
m: mujoco.MjModel, device=None, _check_unsupported=True
) -> types.Model:
"""Puts mujoco.MjModel onto a device, resulting in mjx.Model.""" """Puts mujoco.MjModel onto a device, resulting in mjx.Model."""
if m.ntendon: if _check_unsupported and m.ntendon:
raise NotImplementedError('tendons are not supported') raise NotImplementedError('tendons are not supported')
mesh_geomid = set() mesh_geomid = set()
for g1, g2, ip in collision_driver.geom_pairs(m): for g1, g2, ip in collision_driver.geom_pairs(m):
t1, t2 = m.geom_type[[g1, g2]] t1, t2 = m.geom_type[[g1, g2]]
# check collision function exists for type pair # check collision function exists for type pair
if not collision_driver.has_collision_fn(t1, t2): if _check_unsupported and not collision_driver.has_collision_fn(t1, t2):
t1, t2 = mujoco.mjtGeom(t1), mujoco.mjtGeom(t2) t1, t2 = mujoco.mjtGeom(t1), mujoco.mjtGeom(t2)
raise NotImplementedError(f'({t1}, {t2}) collisions not implemented.') raise NotImplementedError(f'({t1}, {t2}) collisions not implemented.')
# margin/gap not supported for meshes and height fields # margin/gap not supported for meshes and height fields
@@ -89,7 +91,7 @@ def put_model(m: mujoco.MjModel, device=None) -> types.Model:
margin = m.pair_margin[ip] margin = m.pair_margin[ip]
else: else:
margin = m.geom_margin[g1] + m.geom_margin[g2] margin = m.geom_margin[g1] + m.geom_margin[g2]
if margin.any(): if _check_unsupported and margin.any():
t1, t2 = mujoco.mjtGeom(t1), mujoco.mjtGeom(t2) t1, t2 = mujoco.mjtGeom(t1), mujoco.mjtGeom(t2)
raise NotImplementedError(f'({t1}, {t2}) margin/gap not implemented.') raise NotImplementedError(f'({t1}, {t2}) margin/gap not implemented.')
for t, g in [(t1, g1), (t2, g2)]: for t, g in [(t1, g1), (t2, g2)]:
@@ -104,12 +106,12 @@ def put_model(m: mujoco.MjModel, device=None) -> types.Model:
(m.eq_type, types.EqType, mujoco.mjtEq), (m.eq_type, types.EqType, mujoco.mjtEq),
): ):
missing = set(enum_field) - set(enum_type) missing = set(enum_field) - set(enum_type)
if missing: if _check_unsupported and missing:
raise NotImplementedError( raise NotImplementedError(
f'{[mj_type(m) for m in missing]} not supported' f'{[mj_type(m) for m in missing]} not supported'
) )
if not np.allclose(m.dof_frictionloss, 0): if _check_unsupported and not np.allclose(m.dof_frictionloss, 0):
raise NotImplementedError('dof_frictionloss is not implemented.') raise NotImplementedError('dof_frictionloss is not implemented.')
mjx_only = {'mesh_convex', 'geom_rbound_hfield'} mjx_only = {'mesh_convex', 'geom_rbound_hfield'}