diff --git a/mjx/mujoco/mjx/_src/collision_driver.py b/mjx/mujoco/mjx/_src/collision_driver.py index 3371a674..70d3a59b 100644 --- a/mjx/mujoco/mjx/_src/collision_driver.py +++ b/mjx/mujoco/mjx/_src/collision_driver.py @@ -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]},', diff --git a/mjx/mujoco/mjx/_src/forward_test.py b/mjx/mujoco/mjx/_src/forward_test.py index f90fae1a..ba200b72 100644 --- a/mjx/mujoco/mjx/_src/forward_test.py +++ b/mjx/mujoco/mjx/_src/forward_test.py @@ -170,6 +170,7 @@ class ForwardTest(absltest.TestCase): np.testing.assert_allclose(dx.qvel, 1 + m.opt.timestep) + class ActuatorTest(parameterized.TestCase): @parameterized.parameters( diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 253a62a5..54e40926 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -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) diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 9910e9b5..dfcb541d 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -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(""" @@ -374,11 +369,8 @@ class ModelIOTest(parameterized.TestCase): """) - 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)), ] diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 898c844a..5b8ad301 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -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,