Internal change.
PiperOrigin-RevId: 885152209 Change-Id: I56941ba443fe3daad02eff39d6dbcaa9615df41b
This commit is contained in:
committed by
Copybara-Service
parent
c0ab1cc130
commit
206e245353
@@ -349,8 +349,8 @@ def make_condim(
|
||||
m: Union[Model, mujoco.MjModel], impl: Impl = Impl.JAX
|
||||
) -> np.ndarray:
|
||||
"""Returns the dims of the contacts for a Model."""
|
||||
if impl not in (Impl.JAX, Impl.C):
|
||||
raise ValueError('make_condim only supports JAX and C backends.')
|
||||
if impl != Impl.JAX:
|
||||
raise ValueError('make_condim only supports JAX backend.')
|
||||
|
||||
if isinstance(m, mujoco.MjModel):
|
||||
sdf_initpoints = m.opt.sdf_initpoints
|
||||
@@ -389,8 +389,6 @@ def make_condim(
|
||||
func = _COLLISION_FUNC.get(k.types, None)
|
||||
if func is not None:
|
||||
ncon = func.ncon # pytype: disable=attribute-error
|
||||
elif impl == Impl.C:
|
||||
ncon = _MAX_NCON
|
||||
else:
|
||||
raise ValueError(
|
||||
f'Collision function not found for geom types {k.types[0]},',
|
||||
|
||||
@@ -170,6 +170,7 @@ class ForwardTest(absltest.TestCase):
|
||||
np.testing.assert_allclose(dx.qvel, 1 + m.opt.timestep)
|
||||
|
||||
|
||||
|
||||
class ActuatorTest(parameterized.TestCase):
|
||||
|
||||
@parameterized.parameters(
|
||||
|
||||
+5
-262
@@ -90,7 +90,7 @@ def _resolve_device(
|
||||
logging.debug('Picking default device: %s.', device_0)
|
||||
return device_0
|
||||
|
||||
if impl == types.Impl.C or impl == types.Impl.CPP:
|
||||
if impl == types.Impl.CPP:
|
||||
cpu_0 = jax.devices('cpu')[0]
|
||||
logging.debug('Picking default device: %s', cpu_0)
|
||||
return cpu_0
|
||||
@@ -128,7 +128,7 @@ def _check_impl_device_compatibility(
|
||||
_check_warp_installed()
|
||||
|
||||
is_cpu_device = device.platform == 'cpu'
|
||||
if impl == types.Impl.C or impl == types.Impl.CPP:
|
||||
if impl == types.Impl.CPP:
|
||||
if not is_cpu_device:
|
||||
raise AssertionError(
|
||||
f'C implementation requires a CPU device, got {device}.'
|
||||
@@ -248,7 +248,6 @@ def _put_option(
|
||||
fields['jacobian'] = types.JacobianType(o.jacobian)
|
||||
|
||||
option_obj = {
|
||||
types.Impl.C: types.OptionC,
|
||||
types.Impl.JAX: types.OptionJAX,
|
||||
types.Impl.WARP: mjxw.types.OptionWarp,
|
||||
}[impl]
|
||||
@@ -266,8 +265,6 @@ def _put_option(
|
||||
impl_fields['has_fluid_params'] = has_fluid_params
|
||||
return types.Option(**fields, _impl=types.OptionJAX(**impl_fields))
|
||||
|
||||
if impl == types.Impl.C:
|
||||
return types.Option(**fields, _impl=types.OptionC(**impl_fields))
|
||||
|
||||
if impl == types.Impl.WARP:
|
||||
impl_fields = {k: _wp_to_np_type(v) for k, v in impl_fields.items()}
|
||||
@@ -425,28 +422,6 @@ def _put_model_jax(
|
||||
return _strip_weak_type(model)
|
||||
|
||||
|
||||
def _put_model_c(
|
||||
m: mujoco.MjModel,
|
||||
device: Optional[jax.Device] = None,
|
||||
) -> types.Model:
|
||||
"""Puts mujoco.MjModel onto a device, resulting in mjx.Model."""
|
||||
mj_field_names = {f.name for f in types.Model.fields() if f.name != '_impl'}
|
||||
fields = {f: getattr(m, f) for f in mj_field_names}
|
||||
fields['cam_mat0'] = fields['cam_mat0'].reshape((-1, 3, 3))
|
||||
fields['opt'] = _put_option(m.opt, impl=types.Impl.C)
|
||||
fields['stat'] = _put_statistic(m.stat, impl=types.Impl.C)
|
||||
|
||||
c_impl_keys = (
|
||||
types.ModelC.__annotations__.keys() - types.Model.__annotations__.keys()
|
||||
)
|
||||
c_impl_dict = {k: getattr(m, k) for k in c_impl_keys}
|
||||
c_impl_obj = types.ModelC(**{k: copy.copy(v) for k, v in c_impl_dict.items()})
|
||||
|
||||
model = types.Model(
|
||||
**{k: copy.copy(v) for k, v in fields.items()}, _impl=c_impl_obj
|
||||
)
|
||||
model = jax.device_put(model, device=device)
|
||||
return _strip_weak_type(model)
|
||||
|
||||
|
||||
def _put_model_warp(
|
||||
@@ -506,8 +481,8 @@ def _put_model_cpp(
|
||||
mj_field_names = {f.name for f in types.Model.fields() if f.name != '_impl'}
|
||||
fields = {f: getattr(m, f) for f in mj_field_names}
|
||||
fields['cam_mat0'] = fields['cam_mat0'].reshape((-1, 3, 3))
|
||||
fields['opt'] = _put_option(m.opt, impl=types.Impl.C)
|
||||
fields['stat'] = _put_statistic(m.stat, impl=types.Impl.C)
|
||||
fields['opt'] = _put_option(m.opt, impl=types.Impl.JAX)
|
||||
fields['stat'] = _put_statistic(m.stat, impl=types.Impl.JAX)
|
||||
|
||||
# get the pointer address
|
||||
# we use a 0-d array
|
||||
@@ -562,8 +537,6 @@ def put_model(
|
||||
impl, device = _resolve_impl_and_device(impl, device)
|
||||
if impl == types.Impl.JAX:
|
||||
return _put_model_jax(m, device)
|
||||
elif impl == types.Impl.C:
|
||||
return _put_model_c(m, device)
|
||||
elif impl == types.Impl.WARP:
|
||||
_check_warp_installed()
|
||||
graph_mode = graph_mode or getattr(mjxw.types.GraphMode, 'WARP')
|
||||
@@ -740,125 +713,6 @@ def _make_data_jax(
|
||||
return d
|
||||
|
||||
|
||||
def _make_data_c(
|
||||
m: Union[types.Model, mujoco.MjModel],
|
||||
device: Optional[jax.Device] = None,
|
||||
) -> types.Data:
|
||||
"""Allocate and initialize Data for the C implementation."""
|
||||
# TODO(stunya): The C implementation should not use static dimensions, and
|
||||
# the backend implementation details should be kept hidden from JAX
|
||||
# altogether.
|
||||
dim = collision_driver.make_condim(m, impl=types.Impl.C)
|
||||
efc_type = constraint.make_efc_type(m, dim)
|
||||
efc_address = constraint.make_efc_address(m, dim, efc_type)
|
||||
ne, nf, nl, nc = constraint.counts(efc_type)
|
||||
ncon, nefc = dim.size, ne + nf + nl + nc
|
||||
|
||||
float_ = jp.zeros(1, float).dtype
|
||||
int_ = jp.zeros(1, int).dtype
|
||||
# TODO(stunya): remove the JAX contact from C data.
|
||||
contact = _make_data_contact_jax(dim, efc_address)
|
||||
|
||||
def get(m, name: str):
|
||||
return getattr(m._impl, name) if hasattr(m, '_impl') else getattr(m, name) # pylint: disable=protected-access
|
||||
|
||||
nflexvert = get(m, 'nflexvert')
|
||||
nflexedge = get(m, 'nflexedge')
|
||||
nflexelem = get(m, 'nflexelem')
|
||||
nbvh = get(m, 'nbvh')
|
||||
nbvhdynamic = get(m, 'nbvhdynamic')
|
||||
zero_impl_fields = {
|
||||
'solver_niter': (int_,),
|
||||
'cinert': (m.nbody, 10, float_),
|
||||
'light_xpos': (m.nlight, 3, float_),
|
||||
'light_xdir': (m.nlight, 3, float_),
|
||||
'flexvert_xpos': (nflexvert, 3, float_),
|
||||
'flexelem_aabb': (nflexelem, 6, float_),
|
||||
'flexedge_J': (m.nJfe, float_),
|
||||
'flexedge_length': (nflexedge, float_),
|
||||
'ten_J_rownnz': (m.ntendon, np.int32),
|
||||
'ten_J_rowadr': (m.ntendon, np.int32),
|
||||
'ten_J_colind': (m.nJten, np.int32),
|
||||
'ten_J': (m.nJten, float_),
|
||||
'ten_wrapadr': (m.ntendon, np.int32),
|
||||
'ten_wrapnum': (m.ntendon, np.int32),
|
||||
'wrap_obj': (m.nwrap, 2, np.int32),
|
||||
'wrap_xpos': (m.nwrap, 6, float_),
|
||||
'moment_rownnz': (m.nu, np.int32),
|
||||
'moment_rowadr': (m.nu, np.int32),
|
||||
'moment_colind': (m.nJmom, np.int32),
|
||||
'actuator_moment': (m.nJmom, float_),
|
||||
'bvh_aabb_dyn': (nbvhdynamic, 6, float_),
|
||||
'bvh_active': (nbvh, np.uint8),
|
||||
'tree_asleep': (m.ntree, int_),
|
||||
'tree_awake': (m.ntree, int_),
|
||||
'body_awake': (m.nbody, int_),
|
||||
'body_awake_ind': (m.nbody, int_),
|
||||
'parent_awake_ind': (m.nbody, int_),
|
||||
'dof_awake_ind': (m.nv, int_),
|
||||
'tree_island': (m.ntree, int_),
|
||||
'map_itree2tree': (m.ntree, int_),
|
||||
'flexedge_velocity': (nflexedge, float_),
|
||||
'crb': (m.nbody, 10, float_),
|
||||
'qM': (m.nM, float_),
|
||||
'M': (m.nC, float_),
|
||||
'qLD': (m.nC, float_),
|
||||
'qH': (m.nC, float_),
|
||||
'qHDiagInv': (m.nv, float_),
|
||||
'qLDiagInv': (m.nv, float_),
|
||||
'ten_velocity': (m.ntendon, float_),
|
||||
'actuator_velocity': (m.nu, float_),
|
||||
'plugin_data': (get(m, 'nplugin'), np.uint64),
|
||||
'qDeriv': (m.nD, float_),
|
||||
'qLU': (m.nD, float_),
|
||||
'qfrc_spring': (m.nv, float_),
|
||||
'qfrc_damper': (m.nv, float_),
|
||||
'cacc': (m.nbody, 6, float_),
|
||||
'cfrc_int': (m.nbody, 6, float_),
|
||||
'cfrc_ext': (m.nbody, 6, float_),
|
||||
'subtree_linvel': (m.nbody, 3, float_),
|
||||
'subtree_angmom': (m.nbody, 3, float_),
|
||||
'efc_J': (nefc, m.nv, float_),
|
||||
'efc_pos': (nefc, float_),
|
||||
'efc_margin': (nefc, float_),
|
||||
'efc_frictionloss': (nefc, float_),
|
||||
'efc_D': (nefc, float_),
|
||||
'efc_aref': (nefc, float_),
|
||||
'efc_force': (nefc, float_),
|
||||
}
|
||||
zero_impl_fields = {
|
||||
k: np.zeros(v[:-1], dtype=v[-1]) for k, v in zero_impl_fields.items()
|
||||
}
|
||||
impl = types.DataC(
|
||||
ne=ne,
|
||||
nf=nf,
|
||||
nl=nl,
|
||||
nefc=nefc,
|
||||
ncon=ncon,
|
||||
contact=contact,
|
||||
efc_type=efc_type,
|
||||
**zero_impl_fields,
|
||||
)
|
||||
|
||||
d = types.Data(
|
||||
qpos=jp.array(m.qpos0, dtype=float_),
|
||||
eq_active=m.eq_active0,
|
||||
_impl=impl,
|
||||
**_make_data_public_fields(m),
|
||||
)
|
||||
|
||||
if m.nmocap:
|
||||
# Set mocap_pos/quat = body_pos/quat for mocap bodies as done in C MuJoCo.
|
||||
body_mask = m.body_mocapid >= 0
|
||||
body_pos = m.body_pos[body_mask]
|
||||
body_quat = m.body_quat[body_mask]
|
||||
d = d.replace(
|
||||
mocap_pos=body_pos[m.body_mocapid[body_mask]],
|
||||
mocap_quat=body_quat[m.body_mocapid[body_mask]],
|
||||
)
|
||||
|
||||
d = jax.device_put(d, device=device)
|
||||
return d
|
||||
|
||||
|
||||
def _get_nested_attr(obj: Any, attr_name: str, split: str) -> Any:
|
||||
@@ -1024,8 +878,6 @@ def make_data(
|
||||
|
||||
if impl == types.Impl.JAX:
|
||||
return _make_data_jax(m, device)
|
||||
elif impl == types.Impl.C:
|
||||
return _make_data_c(m, device)
|
||||
elif impl == types.Impl.CPP:
|
||||
return _make_data_cpp(m, device, keepalive_refs=keepalive_refs)
|
||||
elif impl == types.Impl.WARP:
|
||||
@@ -1226,111 +1078,6 @@ def _put_data_jax(
|
||||
return _strip_weak_type(data)
|
||||
|
||||
|
||||
def _put_data_c(
|
||||
m: mujoco.MjModel, d: mujoco.MjData, device: Optional[jax.Device] = None
|
||||
) -> types.Data:
|
||||
"""Puts mujoco.MjData onto a device, resulting in mjx.Data."""
|
||||
# TODO(stunya): ncon, nefc should potentially be jax.Array, and contact/efc
|
||||
# should not be materialized in JAX.
|
||||
dim = collision_driver.make_condim(m, impl=types.Impl.C)
|
||||
efc_type = constraint.make_efc_type(m, dim)
|
||||
efc_address = constraint.make_efc_address(m, dim, efc_type)
|
||||
ne, nf, nl, nc = constraint.counts(efc_type)
|
||||
ncon, nefc = dim.size, ne + nf + nl + nc
|
||||
|
||||
# TODO(stunya): remove this check.
|
||||
for d_val, val, name in (
|
||||
(d.ncon, ncon, 'ncon'),
|
||||
(d.ne, ne, 'ne'),
|
||||
(d.nf, nf, 'nf'),
|
||||
(d.nl, nl, 'nl'),
|
||||
(d.nefc, nefc, 'nefc'),
|
||||
):
|
||||
if d_val > val:
|
||||
raise ValueError(f'd.{name} too high, d.{name} = {d_val}, model = {val}')
|
||||
|
||||
fields = _put_data_public_fields(d)
|
||||
|
||||
# Implementation specific fields.
|
||||
impl_fields = {
|
||||
f.name: getattr(d, f.name)
|
||||
for f in types.DataC.fields()
|
||||
if hasattr(d, f.name)
|
||||
}
|
||||
for f in types.DataC.fields():
|
||||
if not hasattr(d, f.name) and hasattr(m, f.name):
|
||||
impl_fields[f.name] = getattr(m, f.name)
|
||||
|
||||
# TODO(stunya): support islanding via C impl.
|
||||
impl_fields['solver_niter'] = impl_fields['solver_niter'][0]
|
||||
|
||||
# TODO(stunya): remove reliance on JAX _put_contact.
|
||||
contact, contact_map = _put_contact(d.contact, dim, efc_address)
|
||||
|
||||
# TODO(stunya): remove reliance on dense efc_J.
|
||||
if mujoco.mj_isSparse(m):
|
||||
efc_j = np.zeros((d.efc_J_rownnz.shape[0], m.nv))
|
||||
mujoco.mju_sparse2dense(
|
||||
efc_j,
|
||||
impl_fields['efc_J'],
|
||||
d.efc_J_rownnz,
|
||||
d.efc_J_rowadr,
|
||||
d.efc_J_colind,
|
||||
)
|
||||
impl_fields['efc_J'] = efc_j
|
||||
else:
|
||||
impl_fields['efc_J'] = impl_fields['efc_J'].reshape(
|
||||
(-1 if m.nv else 0, m.nv)
|
||||
)
|
||||
|
||||
# move efc rows to their correct offsets
|
||||
for fname in (
|
||||
'efc_J',
|
||||
'efc_pos',
|
||||
'efc_margin',
|
||||
'efc_frictionloss',
|
||||
'efc_D',
|
||||
'efc_aref',
|
||||
'efc_force',
|
||||
):
|
||||
value = np.zeros((nefc, m.nv)) if fname == 'efc_J' else np.zeros(nefc)
|
||||
for i in range(3):
|
||||
value_beg = sum([ne, nf][:i])
|
||||
d_beg = sum([d.ne, d.nf][:i])
|
||||
size = [d.ne, d.nf, d.nl][i]
|
||||
value[value_beg : value_beg + size] = impl_fields[fname][
|
||||
d_beg : d_beg + size
|
||||
]
|
||||
|
||||
# for nc, we may reorder contacts so they match MJX order: group by dim
|
||||
for id_to, id_from in enumerate(contact_map):
|
||||
if id_from == -1:
|
||||
continue
|
||||
num_rows = dim[id_to]
|
||||
if num_rows > 1 and m.opt.cone == mujoco.mjtCone.mjCONE_PYRAMIDAL:
|
||||
num_rows = (num_rows - 1) * 2
|
||||
efc_i, efc_o = d.contact.efc_address[id_from], efc_address[id_to]
|
||||
if efc_i == -1:
|
||||
continue
|
||||
value[efc_o : efc_o + num_rows] = impl_fields[fname][
|
||||
efc_i : efc_i + num_rows
|
||||
]
|
||||
|
||||
impl_fields[fname] = value
|
||||
|
||||
impl_fields['contact'] = contact
|
||||
impl_fields.update(
|
||||
ne=ne, nf=nf, nl=nl, nefc=nefc, ncon=ncon, efc_type=efc_type
|
||||
)
|
||||
|
||||
# copy because device_put is async:
|
||||
data_jax = types.DataC(**{k: copy.copy(v) for k, v in impl_fields.items()})
|
||||
data = types.Data(
|
||||
**{k: copy.copy(v) for k, v in fields.items()}, _impl=data_jax
|
||||
)
|
||||
|
||||
data = jax.device_put(data, device=device)
|
||||
return _strip_weak_type(data)
|
||||
|
||||
|
||||
# TODO(josechenf): Iterate on the keepalive implementation to make it easier to
|
||||
@@ -1470,8 +1217,6 @@ def put_data(
|
||||
impl, device = _resolve_impl_and_device(impl, device)
|
||||
if impl == types.Impl.JAX:
|
||||
return _put_data_jax(m, d, device)
|
||||
elif impl == types.Impl.C:
|
||||
return _put_data_c(m, d, device)
|
||||
elif impl == types.Impl.CPP:
|
||||
return _put_data_cpp(
|
||||
m,
|
||||
@@ -1613,8 +1358,6 @@ def _get_data_into(
|
||||
|
||||
if d.impl == types.Impl.JAX:
|
||||
all_fields = types.Data.fields() + types.DataJAX.fields()
|
||||
elif d.impl == types.Impl.C:
|
||||
all_fields = types.Data.fields() + types.DataC.fields()
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'get_data_into for implementation "{d.impl}" not implemented yet.'
|
||||
@@ -1814,7 +1557,7 @@ def get_data_into(
|
||||
|
||||
d = jax.device_get(d)
|
||||
|
||||
if d.impl in (types.Impl.JAX, types.Impl.C):
|
||||
if d.impl == types.Impl.JAX:
|
||||
# TODO(stunya): Split out _get_data_into once codepaths diverge enough.
|
||||
return _get_data_into(result, m, d)
|
||||
|
||||
|
||||
@@ -136,7 +136,7 @@ class ModelIOTest(parameterized.TestCase):
|
||||
|
||||
@parameterized.product(
|
||||
xml=(_MULTIPLE_CONVEX_OBJECTS, _MULTIPLE_CONSTRAINTS),
|
||||
impl=('jax', 'c', 'warp', 'cpp'),
|
||||
impl=('jax', 'warp', 'cpp'),
|
||||
)
|
||||
@mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'true'})
|
||||
def test_put_model(self, xml, impl):
|
||||
@@ -171,11 +171,7 @@ class ModelIOTest(parameterized.TestCase):
|
||||
if impl == 'jax':
|
||||
# fields restricted to MuJoCo should not be populated
|
||||
self.assertFalse(hasattr(mx, 'bvh_aabb'))
|
||||
elif impl == 'c':
|
||||
# Options specific to C are populated.
|
||||
self.assertEqual(mx.opt._impl.noslip_iterations, m.opt.noslip_iterations)
|
||||
# Fields private to C backend impl are populated.
|
||||
self.assertTrue(hasattr(mx._impl, 'bvh_aabb'))
|
||||
|
||||
elif impl == 'warp':
|
||||
# Options specific to Warp are populated.
|
||||
self.assertTrue(hasattr(mx.opt._impl, 'ls_parallel'))
|
||||
@@ -353,8 +349,7 @@ class ModelIOTest(parameterized.TestCase):
|
||||
expected = graph_mode or mjxw_types.GraphMode.WARP
|
||||
self.assertEqual(mx.opt._impl.graph_mode, expected)
|
||||
|
||||
@parameterized.parameters('c', 'jax')
|
||||
def test_unsupported_contact_types(self, impl):
|
||||
def test_unsupported_contact_types(self):
|
||||
"""Tests that unsupported contact types raise an error."""
|
||||
m = mujoco.MjModel.from_xml_string("""
|
||||
<mujoco>
|
||||
@@ -374,11 +369,8 @@ class ModelIOTest(parameterized.TestCase):
|
||||
</mujoco>
|
||||
""")
|
||||
|
||||
if impl == 'jax':
|
||||
with self.assertRaises(ValueError):
|
||||
mjx.make_data(m, impl=impl)
|
||||
if impl == 'c':
|
||||
mjx.make_data(m, impl=impl)
|
||||
with self.assertRaises(ValueError):
|
||||
mjx.make_data(m, impl='jax')
|
||||
|
||||
|
||||
class DataIOTest(parameterized.TestCase):
|
||||
@@ -390,7 +382,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
self.tempdir = tempfile.TemporaryDirectory()
|
||||
wp.config.kernel_cache_dir = self.tempdir.name
|
||||
|
||||
@parameterized.parameters('jax', 'c', 'cpp')
|
||||
@parameterized.parameters('jax', 'cpp')
|
||||
def test_make_data(self, impl: str):
|
||||
"""Test that make_data returns the correct shapes."""
|
||||
m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONVEX_OBJECTS)
|
||||
@@ -437,8 +429,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
self.assertEqual(d._impl.crb.shape, (nbody, 10))
|
||||
if impl == 'jax':
|
||||
self.assertEqual(d._impl.actuator_moment.shape, (1, nv))
|
||||
elif impl == 'c':
|
||||
self.assertEqual(d._impl.actuator_moment.shape, (m.nJmom,))
|
||||
|
||||
self.assertEqual(d._impl.contact.dist.shape, (ncon,))
|
||||
self.assertEqual(d._impl.contact.pos.shape, (ncon, 3))
|
||||
self.assertEqual(d._impl.contact.frame.shape, (ncon, 3, 3))
|
||||
@@ -466,10 +457,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
self.assertEqual(d._impl.qM.shape, (nv, nv))
|
||||
self.assertEqual(d._impl.qLD.shape, (nv, nv))
|
||||
self.assertEqual(d._impl.qLDiagInv.shape, (0,))
|
||||
elif impl == 'c':
|
||||
self.assertEqual(d._impl.qM.shape, (nm,))
|
||||
self.assertEqual(d._impl.qLD.shape, (nm,))
|
||||
self.assertEqual(d._impl.qLDiagInv.shape, (nv,))
|
||||
|
||||
|
||||
# test sparse
|
||||
m.opt.jacobian = mujoco.mjtJacobian.mjJAC_SPARSE
|
||||
@@ -478,10 +466,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
self.assertEqual(d._impl.qLD.shape, (nm,))
|
||||
self.assertEqual(d._impl.qLDiagInv.shape, (nv,))
|
||||
|
||||
if impl == 'c':
|
||||
# check C specific fields
|
||||
self.assertEqual(d._impl.light_xpos.shape, (m.nlight, 3))
|
||||
self.assertEqual(d._impl.bvh_active.shape, (m.nbvh,))
|
||||
|
||||
|
||||
@mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'true'})
|
||||
def test_make_data_warp(self):
|
||||
@@ -494,7 +479,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
self.assertEqual(d._impl.contact__dist.shape[0], 9)
|
||||
self.assertEqual(d._impl.efc__pos.shape[0], 23)
|
||||
|
||||
@parameterized.parameters('jax', 'c', 'cpp', 'warp')
|
||||
@parameterized.parameters('jax', 'cpp', 'warp')
|
||||
def test_put_data(self, impl: str):
|
||||
"""Test that put_data puts the correct data for dense and sparse."""
|
||||
if impl == 'warp':
|
||||
@@ -537,10 +522,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
qm = np.zeros((m.nv, m.nv), dtype=np.float64)
|
||||
mujoco.mj_fullM(m, qm, d.qM)
|
||||
np.testing.assert_allclose(qm, mjx.full_m(mjx.put_model(m), dx))
|
||||
elif impl == 'c':
|
||||
np.testing.assert_allclose(dx._impl.qM, d.qM)
|
||||
np.testing.assert_allclose(dx._impl.qLD, d.qLD)
|
||||
np.testing.assert_allclose(dx._impl.qLDiagInv, d.qLDiagInv)
|
||||
|
||||
elif impl == 'cpp':
|
||||
self.assertTrue(hasattr(dx._impl, 'pointer_lo'))
|
||||
self.assertTrue(hasattr(dx._impl, 'pointer_hi'))
|
||||
@@ -616,8 +598,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
qm = np.zeros((m.nv, m.nv))
|
||||
mujoco.mj_fullM(m, qm, d.qM)
|
||||
np.testing.assert_allclose(dx_from_dense._impl.qM, qm, atol=1e-8)
|
||||
elif impl == 'c':
|
||||
np.testing.assert_allclose(dx_from_dense._impl.qM, d.qM, atol=1e-8)
|
||||
|
||||
|
||||
def test_put_data_warp_ndim(self):
|
||||
"""Tests that put_data produces expected dimensions for Warp fields."""
|
||||
@@ -645,7 +626,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
_ = jax.tree.map_with_path(check_ndim, dx)
|
||||
|
||||
@parameterized.parameters(
|
||||
('jax', False), ('jax', True), ('c', False), ('c', True)
|
||||
('jax', False), ('jax', True)
|
||||
)
|
||||
def test_get_data(self, impl: str, sparse: bool):
|
||||
"""Test that get_data makes correct MjData."""
|
||||
@@ -710,9 +691,6 @@ class DataIOTest(parameterized.TestCase):
|
||||
np.testing.assert_allclose(d_2.efc_aref, d.efc_aref)
|
||||
np.testing.assert_allclose(d_2.contact.efc_address, d.contact.efc_address)
|
||||
|
||||
if impl == 'c':
|
||||
# check fields specific to the C implementation
|
||||
np.testing.assert_allclose(d_2.bvh_active, d.bvh_active)
|
||||
|
||||
def test_get_data_simplebody(self):
|
||||
"""Test that get_data works with simple bodies where nC < nM."""
|
||||
@@ -744,7 +722,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
dx = mjx.put_data(m, d)
|
||||
mjx.get_data(m, dx)
|
||||
|
||||
@parameterized.parameters('jax', 'c')
|
||||
@parameterized.parameters(('jax',))
|
||||
def test_get_data_batched(self, impl):
|
||||
"""Test that get_data makes correct List[MjData] for batched Data."""
|
||||
|
||||
@@ -761,7 +739,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
self.assertEqual(ds[0].ncon, 1)
|
||||
self.assertEqual(ds[1].ncon, 0)
|
||||
|
||||
@parameterized.parameters('jax', 'c', 'cpp')
|
||||
@parameterized.parameters('jax', 'cpp')
|
||||
def test_get_data_into(self, impl):
|
||||
"""Test that get_data_into correctly populates an MjData."""
|
||||
|
||||
@@ -801,7 +779,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
dx = mjx.make_data(m, impl='warp')
|
||||
mjx.get_data_into(d, mx, dx)
|
||||
|
||||
@parameterized.parameters('jax', 'c')
|
||||
@parameterized.parameters(('jax',))
|
||||
def test_get_data_into_wrong_shape(self, impl):
|
||||
"""Tests that get_data_into throwsif input and output shapes don't match."""
|
||||
|
||||
@@ -814,7 +792,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
with self.assertRaisesRegex(ValueError, r'Input field.*has shape.*'):
|
||||
mjx.get_data_into(d_2, m, dx)
|
||||
|
||||
@parameterized.parameters('jax', 'c')
|
||||
@parameterized.parameters(('jax',))
|
||||
def test_make_matches_put(self, impl):
|
||||
"""Test that make_data produces a pytree that matches put_data."""
|
||||
m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS)
|
||||
@@ -960,11 +938,11 @@ _DEVICE_TEST_CASES = [
|
||||
('gpu-notnvidia', 'warp', ('cpu', 'error')),
|
||||
('gpu-nvidia', 'warp', ('gpu', Impl.WARP)),
|
||||
('tpu', 'warp', ('tpu', 'error')),
|
||||
# C backend specified.
|
||||
('cpu', 'c', ('cpu', Impl.C)),
|
||||
('gpu-notnvidia', 'c', ('cpu', 'error')),
|
||||
('gpu-nvidia', 'c', ('cpu', 'error')),
|
||||
('tpu', 'c', ('tpu', 'error')),
|
||||
# CPP backend specified.
|
||||
('cpu', 'cpp', ('cpu', Impl.CPP)),
|
||||
('gpu-notnvidia', 'cpp', ('cpu', 'error')),
|
||||
('gpu-nvidia', 'cpp', ('cpu', 'error')),
|
||||
('tpu', 'cpp', ('tpu', 'error')),
|
||||
]
|
||||
|
||||
# Test cases for `_resolve_impl_and_device` where the user does NOT
|
||||
@@ -988,11 +966,11 @@ _DEFAULT_DEVICE_TEST_CASES = [
|
||||
('gpu-notnvidia', 'warp', ('cpu', 'error')),
|
||||
('gpu-nvidia', 'warp', ('gpu', Impl.WARP)),
|
||||
('tpu', 'warp', ('tpu', 'error')),
|
||||
# C backend impl specified, CPU should always be available.
|
||||
('cpu', 'c', ('cpu', Impl.C)),
|
||||
('gpu-notnvidia', 'c', ('cpu', Impl.C)),
|
||||
('gpu-nvidia', 'c', ('cpu', Impl.C)),
|
||||
('tpu', 'c', ('cpu', Impl.C)),
|
||||
# CPP backend impl specified, CPU should always be available.
|
||||
('cpu', 'cpp', ('cpu', Impl.CPP)),
|
||||
('gpu-notnvidia', 'cpp', ('cpu', Impl.CPP)),
|
||||
('gpu-nvidia', 'cpp', ('cpu', Impl.CPP)),
|
||||
('tpu', 'cpp', ('cpu', Impl.CPP)),
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -29,8 +29,8 @@ import numpy as np
|
||||
class Impl(enum.Enum):
|
||||
"""Implementation to use."""
|
||||
|
||||
C = 'c'
|
||||
CPP = 'cpp'
|
||||
C = 'cpp' # alias C -> CPP
|
||||
JAX = 'jax'
|
||||
WARP = 'warp'
|
||||
|
||||
@@ -490,22 +490,6 @@ class OptionJAX(PyTreeNode):
|
||||
has_fluid_params: bool
|
||||
|
||||
|
||||
class OptionC(PyTreeNode):
|
||||
"""C-specific option."""
|
||||
|
||||
o_margin: jax.Array
|
||||
o_solref: jax.Array
|
||||
o_solimp: jax.Array
|
||||
o_friction: jax.Array
|
||||
disableactuator: int
|
||||
sdf_initpoints: int
|
||||
has_fluid_params: bool
|
||||
noslip_tolerance: jax.Array
|
||||
ccd_tolerance: jax.Array
|
||||
sleep_tolerance: jax.Array
|
||||
noslip_iterations: int
|
||||
ccd_iterations: int
|
||||
sdf_iterations: int
|
||||
|
||||
|
||||
class Option(PyTreeNode):
|
||||
@@ -528,7 +512,7 @@ class Option(PyTreeNode):
|
||||
integrator: IntegratorType
|
||||
solver: SolverType
|
||||
timestep: jax.Array
|
||||
_impl: Union[OptionJAX, OptionC, mjxw_types.OptionWarp]
|
||||
_impl: Union[OptionJAX, mjxw_types.OptionWarp]
|
||||
|
||||
|
||||
class ModelCPP(PyTreeNode):
|
||||
@@ -549,121 +533,6 @@ class DataCPP(PyTreeNode):
|
||||
pointer_hi: jax.Array
|
||||
|
||||
|
||||
class ModelC(PyTreeNode):
|
||||
"""CPU-specific model data."""
|
||||
|
||||
nbvh: jax.Array
|
||||
nbvhstatic: jax.Array
|
||||
nbvhdynamic: jax.Array
|
||||
ntree: jax.Array
|
||||
nflex: jax.Array
|
||||
nflexvert: jax.Array
|
||||
nflexedge: jax.Array
|
||||
nflexelem: jax.Array
|
||||
nflexelemdata: jax.Array
|
||||
nflexshelldata: jax.Array
|
||||
nflexevpair: jax.Array
|
||||
nflextexcoord: jax.Array
|
||||
nplugin: jax.Array
|
||||
narena: jax.Array
|
||||
body_bvhadr: jax.Array
|
||||
body_bvhnum: jax.Array
|
||||
bvh_child: jax.Array
|
||||
bvh_nodeid: jax.Array
|
||||
bvh_aabb: jax.Array
|
||||
oct_child: jax.Array
|
||||
oct_aabb: jax.Array
|
||||
oct_coeff: jax.Array
|
||||
dof_length: jax.Array
|
||||
tree_bodyadr: jax.Array
|
||||
tree_bodynum: jax.Array
|
||||
tree_dofadr: jax.Array
|
||||
tree_dofnum: jax.Array
|
||||
tree_sleep_policy: jax.Array
|
||||
geom_plugin: jax.Array
|
||||
light_bodyid: jax.Array
|
||||
light_targetbodyid: jax.Array
|
||||
flex_contype: jax.Array
|
||||
flex_conaffinity: jax.Array
|
||||
flex_condim: jax.Array
|
||||
flex_priority: jax.Array
|
||||
flex_solmix: jax.Array
|
||||
flex_solref: jax.Array
|
||||
flex_solimp: jax.Array
|
||||
flex_friction: jax.Array
|
||||
flex_margin: jax.Array
|
||||
flex_gap: jax.Array
|
||||
flex_internal: jax.Array
|
||||
flex_selfcollide: jax.Array
|
||||
flex_activelayers: jax.Array
|
||||
flex_passive: jax.Array
|
||||
flex_dim: jax.Array
|
||||
flex_vertadr: jax.Array
|
||||
flex_vertnum: jax.Array
|
||||
flex_edgeadr: jax.Array
|
||||
flex_edgenum: jax.Array
|
||||
flex_elemadr: jax.Array
|
||||
flex_elemnum: jax.Array
|
||||
flex_elemdataadr: jax.Array
|
||||
flex_evpairadr: jax.Array
|
||||
flex_evpairnum: jax.Array
|
||||
flex_vertbodyid: jax.Array
|
||||
flex_edge: jax.Array
|
||||
flex_elem: jax.Array
|
||||
flex_elemlayer: jax.Array
|
||||
flex_evpair: jax.Array
|
||||
flex_vert: jax.Array
|
||||
flexedge_length0: jax.Array
|
||||
flexedge_invweight0: jax.Array
|
||||
flex_radius: jax.Array
|
||||
flex_edgestiffness: jax.Array
|
||||
flex_edgedamping: jax.Array
|
||||
flex_edgeequality: jax.Array
|
||||
flex_rigid: jax.Array
|
||||
flexedge_rigid: jax.Array
|
||||
flex_centered: jax.Array
|
||||
flex_bvhadr: jax.Array
|
||||
flex_bvhnum: jax.Array
|
||||
flexedge_J_rownnz: jax.Array
|
||||
flexedge_J_rowadr: jax.Array
|
||||
flexedge_J_colind: jax.Array
|
||||
mesh_polynum: jax.Array
|
||||
mesh_polyadr: jax.Array
|
||||
mesh_polynormal: jax.Array
|
||||
mesh_polyvertadr: jax.Array
|
||||
mesh_polyvertnum: jax.Array
|
||||
mesh_polyvert: jax.Array
|
||||
mesh_polymapadr: jax.Array
|
||||
mesh_polymapnum: jax.Array
|
||||
mesh_polymap: jax.Array
|
||||
tendon_treenum: jax.Array
|
||||
tendon_treeid: jax.Array
|
||||
actuator_plugin: jax.Array
|
||||
actuator_history: jax.Array
|
||||
actuator_historyadr: jax.Array
|
||||
actuator_delay: jax.Array
|
||||
sensor_plugin: jax.Array
|
||||
sensor_history: jax.Array
|
||||
sensor_historyadr: jax.Array
|
||||
sensor_delay: jax.Array
|
||||
sensor_interval: jax.Array
|
||||
plugin: jax.Array
|
||||
plugin_stateadr: jax.Array
|
||||
B_rownnz: jax.Array # pylint:disable=invalid-name
|
||||
B_rowadr: jax.Array # pylint:disable=invalid-name
|
||||
B_colind: jax.Array # pylint:disable=invalid-name
|
||||
M_rownnz: jax.Array # pylint:disable=invalid-name
|
||||
M_rowadr: jax.Array # pylint:disable=invalid-name
|
||||
M_colind: jax.Array # pylint:disable=invalid-name
|
||||
mapM2M: jax.Array # pylint:disable=invalid-name
|
||||
D_rownnz: jax.Array # pylint:disable=invalid-name
|
||||
D_rowadr: jax.Array # pylint:disable=invalid-name
|
||||
D_diag: jax.Array # pylint:disable=invalid-name
|
||||
D_colind: jax.Array # pylint:disable=invalid-name
|
||||
mapM2D: jax.Array # pylint:disable=invalid-name
|
||||
mapD2M: jax.Array # pylint:disable=invalid-name
|
||||
|
||||
|
||||
class ModelJAX(PyTreeNode):
|
||||
"""JAX-specific model data."""
|
||||
|
||||
@@ -1025,12 +894,11 @@ class Model(PyTreeNode):
|
||||
names: bytes
|
||||
signature: np.uint64
|
||||
_sizes: jax.Array
|
||||
_impl: Union[ModelC, ModelJAX, mjxw_types.ModelWarp]
|
||||
_impl: Union[ModelJAX, mjxw_types.ModelWarp]
|
||||
|
||||
@property
|
||||
def impl(self) -> Impl:
|
||||
return {
|
||||
ModelC: Impl.C,
|
||||
ModelCPP: Impl.CPP,
|
||||
ModelJAX: Impl.JAX,
|
||||
mjxw_types.ModelWarp: Impl.WARP,
|
||||
@@ -1094,82 +962,6 @@ class Contact(PyTreeNode):
|
||||
efc_address: np.ndarray
|
||||
|
||||
|
||||
class DataC(PyTreeNode):
|
||||
"""C-specific data."""
|
||||
|
||||
# constant sizes:
|
||||
# TODO(stunya): make these sizes jax.Array?
|
||||
ncon: int
|
||||
ne: int
|
||||
nf: int
|
||||
nl: int
|
||||
nefc: int
|
||||
# TODO(stunya): remove most of these fields
|
||||
solver_niter: jax.Array
|
||||
tree_asleep: jax.Array
|
||||
plugin_data: jax.Array
|
||||
light_xpos: jax.Array
|
||||
light_xdir: jax.Array
|
||||
cinert: jax.Array
|
||||
flexvert_xpos: jax.Array
|
||||
flexelem_aabb: jax.Array
|
||||
flexedge_J: jax.Array # pylint:disable=invalid-name
|
||||
flexedge_length: jax.Array
|
||||
bvh_aabb_dyn: jax.Array
|
||||
ten_wrapadr: jax.Array
|
||||
ten_wrapnum: jax.Array
|
||||
ten_J_rownnz: jax.Array # pylint:disable=invalid-name
|
||||
ten_J_rowadr: jax.Array # pylint:disable=invalid-name
|
||||
ten_J_colind: jax.Array # pylint:disable=invalid-name
|
||||
ten_J: jax.Array # pylint:disable=invalid-name
|
||||
wrap_obj: jax.Array
|
||||
wrap_xpos: jax.Array
|
||||
moment_rownnz: jax.Array # pylint:disable=invalid-name
|
||||
moment_rowadr: jax.Array # pylint:disable=invalid-name
|
||||
moment_colind: jax.Array # pylint:disable=invalid-name
|
||||
actuator_moment: jax.Array
|
||||
crb: jax.Array
|
||||
qM: jax.Array # pylint:disable=invalid-name
|
||||
M: jax.Array # pylint:disable=invalid-name
|
||||
qLD: jax.Array # pylint:disable=invalid-name
|
||||
qLDiagInv: jax.Array # pylint:disable=invalid-name
|
||||
bvh_active: jax.Array
|
||||
tree_awake: jax.Array
|
||||
body_awake: jax.Array
|
||||
body_awake_ind: jax.Array
|
||||
parent_awake_ind: jax.Array
|
||||
dof_awake_ind: jax.Array
|
||||
# position, velocity dependent:
|
||||
flexedge_velocity: jax.Array
|
||||
ten_velocity: jax.Array
|
||||
actuator_velocity: jax.Array
|
||||
|
||||
qfrc_spring: jax.Array
|
||||
qfrc_damper: jax.Array
|
||||
subtree_linvel: jax.Array
|
||||
subtree_angmom: jax.Array
|
||||
qH: jax.Array # pylint:disable=invalid-name
|
||||
qHDiagInv: jax.Array # pylint:disable=invalid-name
|
||||
qDeriv: jax.Array # pylint:disable=invalid-name
|
||||
qLU: jax.Array # pylint:disable=invalid-name
|
||||
cacc: jax.Array
|
||||
cfrc_int: jax.Array
|
||||
cfrc_ext: jax.Array
|
||||
# dynamically sized arrays which are made static for the frontend JAX API
|
||||
# TODO(stunya): remove these dynamic fields entirely
|
||||
contact: Contact
|
||||
efc_type: jax.Array
|
||||
efc_J: jax.Array # pylint:disable=invalid-name
|
||||
efc_pos: jax.Array
|
||||
efc_margin: jax.Array
|
||||
efc_frictionloss: jax.Array
|
||||
efc_D: jax.Array # pylint:disable=invalid-name
|
||||
tree_island: jax.Array
|
||||
map_itree2tree: jax.Array
|
||||
efc_aref: jax.Array
|
||||
efc_force: jax.Array
|
||||
|
||||
|
||||
class DataJAX(PyTreeNode):
|
||||
"""JAX-specific data."""
|
||||
|
||||
@@ -1316,12 +1108,11 @@ class Data(PyTreeNode):
|
||||
qacc_smooth: jax.Array
|
||||
qfrc_constraint: jax.Array
|
||||
qfrc_inverse: jax.Array
|
||||
_impl: Union[DataC, DataCPP, DataJAX, mjxw_types.DataWarp]
|
||||
_impl: Union[DataCPP, DataJAX, mjxw_types.DataWarp]
|
||||
|
||||
@property
|
||||
def impl(self) -> Impl:
|
||||
return {
|
||||
DataC: Impl.C,
|
||||
DataCPP: Impl.CPP,
|
||||
DataJAX: Impl.JAX,
|
||||
mjxw_types.DataWarp: Impl.WARP,
|
||||
|
||||
Reference in New Issue
Block a user