Internal change.

PiperOrigin-RevId: 885152209
Change-Id: I56941ba443fe3daad02eff39d6dbcaa9615df41b
This commit is contained in:
Yuval Tassa
2026-03-17 12:28:29 -07:00
committed by Copybara-Service
parent c0ab1cc130
commit 206e245353
5 changed files with 39 additions and 528 deletions
+2 -4
View File
@@ -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]},',
+1
View File
@@ -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
View File
@@ -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)
+27 -49
View File
@@ -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)),
]
+4 -213
View File
@@ -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,