diff --git a/mjx/mujoco/mjx/_src/collision_convex.py b/mjx/mujoco/mjx/_src/collision_convex.py index 31ce2057..90a709ab 100644 --- a/mjx/mujoco/mjx/_src/collision_convex.py +++ b/mjx/mujoco/mjx/_src/collision_convex.py @@ -28,8 +28,10 @@ from mujoco.mjx._src.collision_types import FunctionKey from mujoco.mjx._src.collision_types import GeomInfo from mujoco.mjx._src.collision_types import HFieldInfo from mujoco.mjx._src.types import Data +from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import GeomType from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import ModelJAX # pylint: enable=g-importing-member _GeomInfo = Union[GeomInfo, ConvexInfo] @@ -42,6 +44,9 @@ def collider(ncon: int): def collide( m: Model, d: Data, key: FunctionKey, geom: jax.Array ) -> Collision: + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('collider requires JAX backend implementation.') + g1, g2 = geom.T infos = [ GeomInfo(d.geom_xpos[g1], d.geom_xmat[g1], m.geom_size[g1]), @@ -56,7 +61,7 @@ def collider(ncon: int): pos=0, mat=0, size=0, face=0, vert=0 ) elif key.types[i] == GeomType.MESH: - c, cm = infos[i], m.mesh_convex[key.data_ids[i]] + c, cm = infos[i], m._impl.mesh_convex[key.data_ids[i]] infos[i] = ConvexInfo(**vars(c), **vars(cm)) in_axes[i] = jax.tree_util.tree_map(lambda x: None, infos[i]).replace( pos=0, mat=0, size=0 diff --git a/mjx/mujoco/mjx/_src/collision_driver.py b/mjx/mujoco/mjx/_src/collision_driver.py index 7b947ce1..e618fe19 100644 --- a/mjx/mujoco/mjx/_src/collision_driver.py +++ b/mjx/mujoco/mjx/_src/collision_driver.py @@ -71,9 +71,11 @@ from mujoco.mjx._src.collision_sdf import sphere_ellipsoid from mujoco.mjx._src.collision_types import FunctionKey from mujoco.mjx._src.types import Contact from mujoco.mjx._src.types import Data +from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import GeomType from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import ModelJAX # pylint: enable=g-importing-member import numpy as np @@ -227,7 +229,7 @@ def _geom_groups( if types[0] == mujoco.mjtGeom.mjGEOM_HFIELD: # add static grid bounds to the grouping key for hfield collisions geom_rbound_hfield = ( - m.geom_rbound_hfield if isinstance(m, Model) else m.geom_rbound + m._impl.geom_rbound_hfield if isinstance(m, Model) else m.geom_rbound # pytype: disable=attribute-error ) nrow, ncol = m.hfield_nrow[data_ids[0]], m.hfield_ncol[data_ids[0]] xsize, ysize = m.hfield_size[data_ids[0]][:2] @@ -323,11 +325,11 @@ def _contact_groups(m: Model, d: Data) -> Dict[FunctionKey, Contact]: solref=solref, solreffriction=solreffriction, solimp=solimp, - dim=d.contact.dim, + dim=d._impl.contact.dim, # pytype: disable=attribute-error geom1=jp.array(geom[:, 0]), geom2=jp.array(geom[:, 1]), geom=jp.array(geom[:, :2]), - efc_address=d.contact.efc_address, + efc_address=d._impl.contact.efc_address, # pytype: disable=attribute-error ) return groups @@ -374,7 +376,10 @@ def make_condim(m: Union[Model, mujoco.MjModel]) -> np.ndarray: def collision(m: Model, d: Data) -> Data: """Collides geometries.""" - if d.ncon == 0: + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('collision requires JAX backend implementation.') + + if d._impl.ncon == 0: # pytype: disable=attribute-error return d max_geom_pairs = _numeric(m, 'max_geom_pairs') @@ -424,4 +429,4 @@ def collision(m: Model, d: Data) -> Data: contacts = sum([condim_groups[k] for k in sorted(condim_groups)], []) contact = jax.tree_util.tree_map(lambda *x: jp.concatenate(x), *contacts) - return d.replace(contact=contact) + return d.replace(_impl=d._impl.replace(contact=contact)) # pytype: disable=attribute-error diff --git a/mjx/mujoco/mjx/_src/collision_driver_test.py b/mjx/mujoco/mjx/_src/collision_driver_test.py index f33cbb07..ab41bc26 100644 --- a/mjx/mujoco/mjx/_src/collision_driver_test.py +++ b/mjx/mujoco/mjx/_src/collision_driver_test.py @@ -119,7 +119,7 @@ class SphereCollisionTest(parameterized.TestCase): def test_sphere(self, name, mjcf): d, dx = _collide(mjcf) for field in dataclasses.fields(Contact): - _assert_attr_eq(dx.contact, d.contact, field.name, name, 1e-5) + _assert_attr_eq(dx._impl.contact, d.contact, field.name, name, 1e-5) _SPHERE_CONVEX = """ @@ -143,7 +143,7 @@ class SphereCollisionTest(parameterized.TestCase): ) d, dx = _collide(xml) self.assertEmpty(d.contact.dist) - self.assertGreater(dx.contact.dist, 0) + self.assertGreater(dx._impl.contact.dist, 0) # face contact xml = self._SPHERE_CONVEX.replace( @@ -151,38 +151,38 @@ class SphereCollisionTest(parameterized.TestCase): ) d, dx = _collide(xml) for field in dataclasses.fields(Contact): - _assert_attr_eq(dx.contact, d.contact, field.name, 'face', 1e-4) + _assert_attr_eq(dx._impl.contact, d.contact, field.name, 'face', 1e-4) # deep face contact xml = self._SPHERE_CONVEX.replace( '', '' ) d, dx = _collide(xml) - self.assertTrue((dx.contact.dist < 0).all()) + self.assertTrue((dx._impl.contact.dist < 0).all()) self.assertTrue((d.contact.dist < 0).all()) - np.testing.assert_allclose(dx.contact.dist, [-0.07], atol=1e-5) - np.testing.assert_array_almost_equal(dx.contact.pos, d.contact.pos) + np.testing.assert_allclose(dx._impl.contact.dist, [-0.07], atol=1e-5) + np.testing.assert_array_almost_equal(dx._impl.contact.pos, d.contact.pos) np.testing.assert_array_almost_equal( - dx.contact.frame, d.contact.frame.reshape((-1, 3, 3)) + dx._impl.contact.frame, d.contact.frame.reshape((-1, 3, 3)) ) def test_sphere_convex_edge(self): # edge contact d, dx = _collide(self._SPHERE_CONVEX) for field in dataclasses.fields(Contact): - _assert_attr_eq(dx.contact, d.contact, field.name, 'edge', 1e-4) + _assert_attr_eq(dx._impl.contact, d.contact, field.name, 'edge', 1e-4) # deep edge penetration xml = self._SPHERE_CONVEX.replace( '', '' ) d, dx = _collide(xml) - self.assertTrue((dx.contact.dist < 0).all()) + self.assertTrue((dx._impl.contact.dist < 0).all()) self.assertTrue((d.contact.dist < 0).all()) - np.testing.assert_allclose(dx.contact.dist, [-0.06], atol=1e-5) - np.testing.assert_array_almost_equal(dx.contact.pos, d.contact.pos) + np.testing.assert_allclose(dx._impl.contact.dist, [-0.06], atol=1e-5) + np.testing.assert_array_almost_equal(dx._impl.contact.pos, d.contact.pos) np.testing.assert_array_almost_equal( - dx.contact.frame, d.contact.frame.reshape((-1, 3, 3)) + dx._impl.contact.frame, d.contact.frame.reshape((-1, 3, 3)) ) # vertex contact @@ -191,7 +191,7 @@ class SphereCollisionTest(parameterized.TestCase): ) d, dx = _collide(xml) for field in dataclasses.fields(Contact): - _assert_attr_eq(dx.contact, d.contact, field.name, 'vertex', 1e-4) + _assert_attr_eq(dx._impl.contact, d.contact, field.name, 'vertex', 1e-4) # sphere center on vertex xml = self._SPHERE_CONVEX.replace( @@ -199,7 +199,9 @@ class SphereCollisionTest(parameterized.TestCase): ) d, dx = _collide(xml) for field in dataclasses.fields(Contact): - _assert_attr_eq(dx.contact, d.contact, field.name, 'vertex_center', 1e-4) + _assert_attr_eq( + dx._impl.contact, d.contact, field.name, 'vertex_center', 1e-4 + ) class EllipsoidCollisionTest(parameterized.TestCase): @@ -219,10 +221,10 @@ class EllipsoidCollisionTest(parameterized.TestCase): def test_plane_ellipsoid(self): """Tests ellipsoid plane contact.""" d, dx = _collide(self._ELLIPSOID_PLANE) - self.assertLess(dx.contact.dist[0], 0) + self.assertLess(dx._impl.contact.dist[0], 0) for field in dataclasses.fields(Contact): _assert_attr_eq( - dx.contact, d.contact, field.name, 'ellipsoid-plane', 1e-5 + dx._impl.contact, d.contact, field.name, 'ellipsoid-plane', 1e-5 ) _ELLIPSOID_ELLIPSOID = """ @@ -242,10 +244,10 @@ class EllipsoidCollisionTest(parameterized.TestCase): def test_ellipsoid_ellipsoid(self): """Tests ellipsoid ellipsoid contact.""" d, dx = _collide(self._ELLIPSOID_ELLIPSOID) - self.assertLess(dx.contact.dist[0], 0) + self.assertLess(dx._impl.contact.dist[0], 0) for field in dataclasses.fields(Contact): _assert_attr_eq( - dx.contact, d.contact, field.name, 'ellipsoid-ellipsoid', 1e-2 + dx._impl.contact, d.contact, field.name, 'ellipsoid-ellipsoid', 1e-2 ) _ELLIPSOID_SPHERE = """ @@ -266,10 +268,10 @@ class EllipsoidCollisionTest(parameterized.TestCase): """Tests ellipsoid capsule contact.""" d, dx = _collide(self._ELLIPSOID_SPHERE) d.contact.pos[0][2] = 0.03 # MJX finds the point on the surface - self.assertLess(dx.contact.dist[0], 0) + self.assertLess(dx._impl.contact.dist[0], 0) for field in dataclasses.fields(Contact): _assert_attr_eq( - dx.contact, d.contact, field.name, 'ellipsoid-sphere', 1e-4 + dx._impl.contact, d.contact, field.name, 'ellipsoid-sphere', 1e-4 ) _ELLIPSOID_CAPSULE = """ @@ -289,10 +291,10 @@ class EllipsoidCollisionTest(parameterized.TestCase): def test_capsule_ellipsoid(self): """Tests ellipsoid capsule contact.""" d, dx = _collide(self._ELLIPSOID_CAPSULE) - self.assertLess(dx.contact.dist[0], 0) + self.assertLess(dx._impl.contact.dist[0], 0) for field in dataclasses.fields(Contact): _assert_attr_eq( - dx.contact, d.contact, field.name, 'ellipsoid-capsule', 1e-5 + dx._impl.contact, d.contact, field.name, 'ellipsoid-capsule', 1e-5 ) _ELLIPSOID_CYLINDER = """ @@ -313,10 +315,10 @@ class EllipsoidCollisionTest(parameterized.TestCase): """Tests ellipsoid cylinder contact.""" d, dx = _collide(self._ELLIPSOID_CYLINDER) d.contact.pos[0][2] = 0.04 # MJX finds the deepest point on the surface - self.assertLess(dx.contact.dist[0], 0) + self.assertLess(dx._impl.contact.dist[0], 0) for field in dataclasses.fields(Contact): _assert_attr_eq( - dx.contact, d.contact, field.name, 'ellipsoid-cylinder', 1e-4 + dx._impl.contact, d.contact, field.name, 'ellipsoid-cylinder', 1e-4 ) @@ -357,7 +359,7 @@ class CapsuleCollisionTest(parameterized.TestCase): def test_capsule(self, name, mjcf): d, dx = _collide(mjcf) for field in dataclasses.fields(Contact): - _assert_attr_eq(dx.contact, d.contact, field.name, name, 1e-4) + _assert_attr_eq(dx._impl.contact, d.contact, field.name, name, 1e-4) _PARALLEL_CAP = """ @@ -378,14 +380,14 @@ class CapsuleCollisionTest(parameterized.TestCase): """Tests that two parallel capsules are colliding at the midpoint.""" _, dx = _collide(self._PARALLEL_CAP) - np.testing.assert_allclose(dx.contact.dist, -0.05) + np.testing.assert_allclose(dx._impl.contact.dist, -0.05) np.testing.assert_allclose( - dx.contact.pos[0], + dx._impl.contact.pos[0], np.array([0.0, 0.1, (0.15 + 0.2) / 2.0]), atol=1e-5, ) np.testing.assert_allclose( - dx.contact.frame[0, 0, :], np.array([0, 0.0, -1.0]), atol=1e-5 + dx._impl.contact.frame[0, 0, :], np.array([0, 0.0, -1.0]), atol=1e-5 ) _CAP_BOX = """ @@ -408,30 +410,32 @@ class CapsuleCollisionTest(parameterized.TestCase): d, dx = _collide(self._CAP_BOX) # sort positions for comparison - idx = np.lexsort((dx.contact.pos[:, 0], dx.contact.pos[:, 1])) - dx = dx.tree_replace({'contact.pos': dx.contact.pos[idx]}) + idx = np.lexsort((dx._impl.contact.pos[:, 0], dx._impl.contact.pos[:, 1])) + dx = dx.tree_replace({'_impl.contact.pos': dx._impl.contact.pos[idx]}) idx = np.lexsort((d.contact.pos[:, 0], d.contact.pos[:, 1])) d.contact.pos[:] = d.contact.pos[idx] d.contact.frame[:] = d.contact.frame[idx] d.contact.dist[:] = d.contact.dist[idx] for field in dataclasses.fields(Contact): - _assert_attr_eq(dx.contact, d.contact, field.name, 'capsule_convex', 1e-4) + _assert_attr_eq( + dx._impl.contact, d.contact, field.name, 'capsule_convex', 1e-4 + ) def test_capsule_convex_face_deep(self): """Tests deep face penetration.""" xml = self._CAP_BOX.replace('', '') _, dx = _collide(xml) - self.assertTrue((dx.contact.dist < 0).all()) + self.assertTrue((dx._impl.contact.dist < 0).all()) np.testing.assert_array_almost_equal( - dx.contact.pos, np.array([[0.5, 0, 0.425], [-0.4, 0, 0.425]]) + dx._impl.contact.pos, np.array([[0.5, 0, 0.425], [-0.4, 0, 0.425]]) ) np.testing.assert_array_almost_equal( - dx.contact.dist, np.array([-0.15, -0.15]) + dx._impl.contact.dist, np.array([-0.15, -0.15]) ) np.testing.assert_array_almost_equal( - dx.contact.frame[:, 0], np.array([[0, 0, -1]] * 2) + dx._impl.contact.frame[:, 0], np.array([[0, 0, -1]] * 2) ) _CAP_EDGE_BOX = """ @@ -453,11 +457,11 @@ class CapsuleCollisionTest(parameterized.TestCase): """Tests edge contact.""" d, dx = _collide(self._CAP_EDGE_BOX) - c = dx.contact + c = dx._impl.contact self.assertEqual(c.pos.shape[0], 2) self.assertGreater(c.dist[1], 0) # extract the contact point with penetration - c = jax.tree_util.tree_map(lambda x: x[:1], dx.contact) + c = jax.tree_util.tree_map(lambda x: x[:1], dx._impl.contact) c = c.replace(dim=c.dim[:1], efc_address=c.efc_address[:1]) for field in dataclasses.fields(Contact): _assert_attr_eq(c, d.contact, field.name, 'capsule_convex_edge', 1e-4) @@ -469,13 +473,17 @@ class CapsuleCollisionTest(parameterized.TestCase): ) _, dx = _collide(xml) - np.testing.assert_array_equal(dx.contact.dist < 0, np.array([True, False])) - np.testing.assert_array_almost_equal(dx.contact.dist[0], np.array([-0.13])) - np.testing.assert_array_almost_equal( - dx.contact.pos[0], np.array([0.5, 0, 0.435]), decimal=3 + np.testing.assert_array_equal( + dx._impl.contact.dist < 0, np.array([True, False]) ) np.testing.assert_array_almost_equal( - dx.contact.frame[0, 0], np.array([0, 0, -1]), decimal=3 + dx._impl.contact.dist[0], np.array([-0.13]) + ) + np.testing.assert_array_almost_equal( + dx._impl.contact.pos[0], np.array([0.5, 0, 0.435]), decimal=3 + ) + np.testing.assert_array_almost_equal( + dx._impl.contact.frame[0, 0], np.array([0, 0, -1]), decimal=3 ) def test_capsule_convex_edge_shallow_tip(self): @@ -489,16 +497,16 @@ class CapsuleCollisionTest(parameterized.TestCase): xml = xml.replace(' @@ -586,10 +596,10 @@ class CylinderTest(absltest.TestCase): """Tests sphere cylinder contact.""" d, dx = _collide(self._SPHERE_CYLINDER) d.contact.pos[0][2] = 0.05 # MJX finds the deepest point on the surface - self.assertLess(dx.contact.dist[0], 0) + self.assertLess(dx._impl.contact.dist[0], 0) for field in dataclasses.fields(Contact): _assert_attr_eq( - dx.contact, d.contact, field.name, 'sphere-cylinder', 1e-4 + dx._impl.contact, d.contact, field.name, 'sphere-cylinder', 1e-4 ) @@ -612,11 +622,11 @@ class ConvexTest(absltest.TestCase): """Tests box collision with a plane.""" d, dx = _collide(self._BOX_PLANE) - np.testing.assert_array_less(dx.contact.dist[:2], 0) - np.testing.assert_array_less(-dx.contact.dist[2:], 0) + np.testing.assert_array_less(dx._impl.contact.dist[:2], 0) + np.testing.assert_array_less(-dx._impl.contact.dist[2:], 0) # extract the contact points with penetration c = jax.tree_util.tree_map( - lambda x: jp.take(x, jp.array([0, 1]), axis=0), dx.contact + lambda x: jp.take(x, jp.array([0, 1]), axis=0), dx._impl.contact ) c = c.replace(dim=c.dim[[0, 1]], efc_address=c.efc_address[[0, 1]]) for field in dataclasses.fields(Contact): @@ -638,16 +648,18 @@ class ConvexTest(absltest.TestCase): """Tests box collision with a plane.""" d, dx = _collide(self._FLAT_BOX_PLANE) - np.testing.assert_array_less(dx.contact.dist, 0) + np.testing.assert_array_less(dx._impl.contact.dist, 0) # sort positions for comparison - idx = np.lexsort((dx.contact.pos[:, 0], dx.contact.pos[:, 1])) - dx = dx.tree_replace({'contact.pos': dx.contact.pos[idx]}) + idx = np.lexsort((dx._impl.contact.pos[:, 0], dx._impl.contact.pos[:, 1])) + dx = dx.tree_replace({'_impl.contact.pos': dx._impl.contact.pos[idx]}) idx = np.lexsort((d.contact.pos[:, 0], d.contact.pos[:, 1])) d.contact.pos[:] = d.contact.pos[idx] for field in dataclasses.fields(Contact): - _assert_attr_eq(dx.contact, d.contact, field.name, 'flat_box_plane', 1e-5) + _assert_attr_eq( + dx._impl.contact, d.contact, field.name, 'flat_box_plane', 1e-5 + ) _BOX_BOX = """ @@ -669,7 +681,7 @@ class ConvexTest(absltest.TestCase): def test_box_box(self): """Tests a face contact for a box-box collision.""" d, dx = _collide(self._BOX_BOX, keyframe=0) - c = dx.contact + c = dx._impl.contact self.assertEqual(c.pos.shape[0], 4) np.testing.assert_array_less(c.dist, 0) @@ -701,10 +713,10 @@ class ConvexTest(absltest.TestCase): d, dx = _collide(self._BOX_BOX_EDGE) # Only one contact point. - np.testing.assert_array_less(dx.contact.dist[:1], 0) - np.testing.assert_array_less(-dx.contact.dist[1:], 0) + np.testing.assert_array_less(dx._impl.contact.dist[:1], 0) + np.testing.assert_array_less(-dx._impl.contact.dist[1:], 0) # extract the contact point with penetration - c = jax.tree_util.tree_map(lambda x: x[:1], dx.contact) + c = jax.tree_util.tree_map(lambda x: x[:1], dx._impl.contact) c = c.replace(dim=c.dim[:1], efc_address=c.efc_address[:1]) for field in dataclasses.fields(Contact): _assert_attr_eq(c, d.contact, field.name, 'box_box_edge', 1e-2) @@ -736,7 +748,7 @@ class ConvexTest(absltest.TestCase): ).read_bytes(), } _, dx = _collide(self._CONVEX_CONVEX, assets=assets) - c = dx.contact + c = dx._impl.contact # Only one contact point for an edge contact. self.assertLess(c.dist[0], 0) @@ -766,7 +778,7 @@ class ConvexTest(absltest.TestCase): def test_convex_convex_edge(self): """Tests convex-convex collisions with edge contact via _sat_gaussmap.""" _, dx = _collide(self._CONVEX_CONVEX_THIN) - c = dx.contact + c = dx._impl.contact # Only one contact point for an edge contact. self.assertLess(c.dist[0], 0) @@ -781,7 +793,7 @@ class ConvexTest(absltest.TestCase): 'pos="0.0 2.0 0.35"', 'pos="0.0 2.0 0"' ) ) - c = dx.contact + c = dx._impl.contact self.assertTrue((c.dist > 0).all()) @@ -837,8 +849,8 @@ class HFieldTest(absltest.TestCase): # check that all geoms are colliding with the hfield for geom_id in [1, 2, 3]: - mask = (dx.contact.geom == np.array([0, geom_id])).all(axis=1) - c = jax.tree_util.tree_map(lambda x, m=mask: x[m], dx.contact) + mask = (dx._impl.contact.geom == np.array([0, geom_id])).all(axis=1) + c = jax.tree_util.tree_map(lambda x, m=mask: x[m], dx._impl.contact) self.assertTrue((c.dist < 0).any()) self.assertTrue((c.dist > -1e-3).any()) # all contact normals are roughly pointing in the right direction @@ -850,7 +862,7 @@ class HFieldTest(absltest.TestCase): for p in positions: xml = self._HFIELD.replace('= 0).all()) + self.assertTrue((dx._impl.contact.dist >= 0).all()) def test_hfield_deep(self): """Tests that objects with deep penetration do not get stuck.""" @@ -868,8 +880,8 @@ class HFieldTest(absltest.TestCase): # check that all geoms are colliding with the hfield for geom_id in [1, 2, 3]: - mask = (dx.contact.geom == np.array([0, geom_id])).all(axis=1) - c = jax.tree_util.tree_map(lambda x, m=mask: x[m], dx.contact) + mask = (dx._impl.contact.geom == np.array([0, geom_id])).all(axis=1) + c = jax.tree_util.tree_map(lambda x, m=mask: x[m], dx._impl.contact) # all contact normals are in the top half-face of the hfield self.assertTrue((c.frame[:, 0].dot(np.array([0, 0, 1])) > 0.7).all()) @@ -892,8 +904,8 @@ class BodyPairFilterTest(absltest.TestCase): def test_filter_self_collision(self): """Tests that self collisions get filtered.""" d, dx = _collide(self._SELF_COLLISION) - self.assertEqual(dx.contact.pos.shape[0], d.contact.pos.shape[0]) - self.assertEqual(dx.contact.pos.shape[0], 0) + self.assertEqual(dx._impl.contact.pos.shape[0], d.contact.pos.shape[0]) + self.assertEqual(dx._impl.contact.pos.shape[0], 0) _PARENT_CHILD = """ @@ -923,8 +935,8 @@ class BodyPairFilterTest(absltest.TestCase): dx = kinematics_jit_fn(mx, dx) dx = collision_jit_fn(mx, dx) - self.assertEqual(dx.contact.pos.shape[0], d.contact.pos.shape[0]) - self.assertEqual(dx.contact.pos.shape[0], 0) + self.assertEqual(dx._impl.contact.pos.shape[0], d.contact.pos.shape[0]) + self.assertEqual(dx._impl.contact.pos.shape[0], 0) def test_disable_filter_parent_child(self): """Tests that filterparent flag disables parent-child filtering.""" @@ -941,8 +953,8 @@ class BodyPairFilterTest(absltest.TestCase): dx = collision_jit_fn(mx, dx) # one collision between parent-child spheres - self.assertEqual(dx.contact.pos.shape[0], d.contact.pos.shape[0]) - self.assertEqual(dx.contact.pos.shape[0], 1) + self.assertEqual(dx._impl.contact.pos.shape[0], d.contact.pos.shape[0]) + self.assertEqual(dx._impl.contact.pos.shape[0], 1) class DimTest(parameterized.TestCase): @@ -1027,8 +1039,8 @@ class TopKContactTest(absltest.TestCase): dx_all = collision_jit_fn(mx_all, dx) dx_top_k = collision_jit_fn(mx_top_k, dx) - self.assertEqual(dx_all.contact.dist.shape, (3,)) - self.assertEqual(dx_top_k.contact.dist.shape, (2,)) + self.assertEqual(dx_all._impl.contact.dist.shape, (3,)) + self.assertEqual(dx_top_k._impl.contact.dist.shape, (2,)) _CAPSULES_MAX_PAIR = """ @@ -1073,9 +1085,9 @@ class TopKContactTest(absltest.TestCase): dx_all = collision_jit_fn(mx_all, dx) dx_top_k = collision_jit_fn(mx_top_k, dx) - self.assertEqual(dx_all.contact.dist.shape, (6,)) - self.assertEqual(dx_top_k.contact.dist.shape, (2,)) - self.assertTrue((dx_top_k.contact.dist < 0).all()) + self.assertEqual(dx_all._impl.contact.dist.shape, (6,)) + self.assertEqual(dx_top_k._impl.contact.dist.shape, (2,)) + self.assertTrue((dx_top_k._impl.contact.dist < 0).all()) if __name__ == '__main__': diff --git a/mjx/mujoco/mjx/_src/collision_sdf.py b/mjx/mujoco/mjx/_src/collision_sdf.py index 1874211d..3bf8b824 100644 --- a/mjx/mujoco/mjx/_src/collision_sdf.py +++ b/mjx/mujoco/mjx/_src/collision_sdf.py @@ -32,7 +32,9 @@ from mujoco.mjx._src.collision_types import Collision from mujoco.mjx._src.collision_types import GeomInfo from mujoco.mjx._src.dataclasses import PyTreeNode from mujoco.mjx._src.types import Data +from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import ModelJAX # pylint: enable=g-importing-member # the SDF function takes position in, and returns a distance or objective diff --git a/mjx/mujoco/mjx/_src/constraint.py b/mjx/mujoco/mjx/_src/constraint.py index e778bcd8..9b41daa7 100644 --- a/mjx/mujoco/mjx/_src/constraint.py +++ b/mjx/mujoco/mjx/_src/constraint.py @@ -28,10 +28,12 @@ from mujoco.mjx._src.types import ConeType from mujoco.mjx._src.types import ConstraintType from mujoco.mjx._src.types import Contact from mujoco.mjx._src.types import Data +from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import EqType from mujoco.mjx._src.types import JointType from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import ModelJAX from mujoco.mjx._src.types import ObjType # pylint: enable=g-importing-member import numpy as np @@ -281,6 +283,10 @@ def _efc_equality_joint(m: Model, d: Data) -> Optional[_Efc]: def _efc_equality_tendon(m: Model, d: Data) -> Optional[_Efc]: """Calculates constraint rows for tendon equality constraints.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError( + '_efc_equality_tendon requires JAX backend implementation.' + ) eq_id = np.nonzero(m.eq_type == EqType.TENDON)[0] @@ -314,9 +320,9 @@ def _efc_equality_tendon(m: Model, d: Data) -> Optional[_Efc]: return jax.tree_util.tree_map(lambda x: x * active, efc) inv1, inv2 = m.tendon_invweight0[obj1id], m.tendon_invweight0[obj2id] - jac1, jac2 = d.ten_J[obj1id], d.ten_J[obj2id] - pos1 = d.ten_length[obj1id] - m.tendon_length0[obj1id] - pos2 = d.ten_length[obj2id] - m.tendon_length0[obj2id] + jac1, jac2 = d._impl.ten_J[obj1id], d._impl.ten_J[obj2id] + pos1 = d._impl.ten_length[obj1id] - m.tendon_length0[obj1id] + pos2 = d._impl.ten_length[obj2id] - m.tendon_length0[obj2id] invweight = inv1 + inv2 * (obj2id > -1) return rows( @@ -326,8 +332,11 @@ def _efc_equality_tendon(m: Model, d: Data) -> Optional[_Efc]: def _efc_friction(m: Model, d: Data) -> Optional[_Efc]: """Calculates constraint rows for dof frictionloss.""" - dof_id = np.nonzero(m.dof_hasfrictionloss)[0] - tendon_id = np.nonzero(m.tendon_hasfrictionloss)[0] + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('_efc_friction requires JAX backend implementation.') + + dof_id = np.nonzero(m._impl.dof_hasfrictionloss)[0] + tendon_id = np.nonzero(m._impl.tendon_hasfrictionloss)[0] size = dof_id.size + tendon_id.size if (m.opt.disableflags & DisableBit.FRICTIONLOSS) or (size == 0): @@ -337,7 +346,7 @@ def _efc_friction(m: Model, d: Data) -> Optional[_Efc]: args_dof += (m.dof_solimp,) args_dof = jax.tree_util.tree_map(lambda x: x[dof_id], args_dof) - args_ten = (d.ten_J, m.tendon_frictionloss, m.tendon_invweight0) + args_ten = (d._impl.ten_J, m.tendon_frictionloss, m.tendon_invweight0) args_ten += (m.tendon_solref_fri, m.tendon_solimp_fri) args_ten = jax.tree_util.tree_map(lambda x: x[tendon_id], args_ten) @@ -415,6 +424,9 @@ def _efc_limit_slide_hinge(m: Model, d: Data) -> Optional[_Efc]: def _efc_limit_tendon(m: Model, d: Data) -> Optional[_Efc]: """Calculates constraint rows for tendon limits.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('_efc_limit_tendon requires JAX backend implementation.') + tendon_id = np.nonzero(m.tendon_limited)[0] if (m.opt.disableflags & DisableBit.LIMIT) or tendon_id.size == 0: @@ -423,8 +435,8 @@ def _efc_limit_tendon(m: Model, d: Data) -> Optional[_Efc]: length, j, range_, margin, invweight, solref, solimp = jax.tree_util.tree_map( lambda x: x[tendon_id], ( - d.ten_length, - d.ten_J, + d._impl.ten_length, + d._impl.ten_J, m.tendon_range, m.tendon_margin, m.tendon_invweight0, @@ -446,8 +458,12 @@ def _efc_limit_tendon(m: Model, d: Data) -> Optional[_Efc]: def _efc_contact_frictionless(m: Model, d: Data) -> Optional[_Efc]: """Calculates constraint rows for frictionless contacts.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError( + '_efc_contact_frictionless requires JAX backend implementation.' + ) - con_id = np.nonzero(d.contact.dim == 1)[0] + con_id = np.nonzero(d._impl.contact.dim == 1)[0] if con_id.size == 0: return None @@ -473,15 +489,19 @@ def _efc_contact_frictionless(m: Model, d: Data) -> Optional[_Efc]: jp.zeros_like(pos), ) - contact = jax.tree_util.tree_map(lambda x: x[con_id], d.contact) + contact = jax.tree_util.tree_map(lambda x: x[con_id], d._impl.contact) return rows(contact) def _efc_contact_pyramidal(m: Model, d: Data, condim: int) -> Optional[_Efc]: """Calculates constraint rows for frictional pyramidal contacts.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError( + '_efc_contact_pyramidal requires JAX backend implementation.' + ) - con_id = np.nonzero(d.contact.dim == condim)[0] + con_id = np.nonzero(d._impl.contact.dim == condim)[0] if con_id.size == 0: return None @@ -518,15 +538,19 @@ def _efc_contact_pyramidal(m: Model, d: Data, condim: int) -> Optional[_Efc]: jp.zeros_like(pos), ) - contact = jax.tree_util.tree_map(lambda x: x[con_id], d.contact) + contact = jax.tree_util.tree_map(lambda x: x[con_id], d._impl.contact) # concatenate to drop row grouping return jax.tree_util.tree_map(jp.concatenate, rows(contact)) def _efc_contact_elliptic(m: Model, d: Data, condim: int) -> Optional[_Efc]: """Calculates constraint rows for frictional elliptic contacts.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError( + '_efc_contact_elliptic requires JAX backend implementation.' + ) - con_id = np.nonzero(d.contact.dim == condim)[0] + con_id = np.nonzero(d._impl.contact.dim == condim)[0] if con_id.size == 0: return None @@ -563,7 +587,7 @@ def _efc_contact_elliptic(m: Model, d: Data, condim: int) -> Optional[_Efc]: jp.zeros_like(pos), ) - contact = jax.tree_util.tree_map(lambda x: x[con_id], d.contact) + contact = jax.tree_util.tree_map(lambda x: x[con_id], d._impl.contact) # concatenate to drop row grouping return jax.tree_util.tree_map(jp.concatenate, rows(contact)) @@ -602,14 +626,14 @@ def make_efc_type( if not m.opt.disableflags & DisableBit.FRICTIONLOSS: nf_dof = ( - m.dof_hasfrictionloss.sum() - if isinstance(m, Model) + m._impl.dof_hasfrictionloss.sum() + if isinstance(m, Model) and isinstance(m._impl, ModelJAX) else (m.dof_frictionloss > 0).sum() ) efc_types += [ConstraintType.FRICTION_DOF] * nf_dof nf_tendon = ( - m.tendon_hasfrictionloss.sum() - if isinstance(m, Model) + m._impl.tendon_hasfrictionloss.sum() + if isinstance(m, Model) and isinstance(m._impl, ModelJAX) else (m.tendon_frictionloss > 0).sum() ) efc_types += [ConstraintType.FRICTION_TENDON] * nf_tendon @@ -683,10 +707,14 @@ def make_constraint(m: Model, d: Data) -> Data: if not efcs: z = jp.empty(0) - d = d.replace(efc_J=jp.empty((0, m.nv))) - d = d.replace( - efc_D=z, efc_aref=z, efc_frictionloss=z, efc_pos=z, efc_margin=z - ) + d = d.tree_replace({'_impl.efc_J': jp.empty((0, m.nv))}) + d = d.tree_replace({ + '_impl.efc_D': z, + '_impl.efc_aref': z, + '_impl.efc_frictionloss': z, + '_impl.efc_pos': z, + '_impl.efc_margin': z, + }) return d efc = jax.tree_util.tree_map(lambda *x: jp.concatenate(x), *efcs) @@ -699,9 +727,13 @@ def make_constraint(m: Model, d: Data) -> Data: return aref, r, efc.pos_aref + efc.margin, efc.margin, efc.frictionloss aref, r, pos, margin, frictionloss = fn(efc) - d = d.replace( - efc_J=efc.J, efc_D=1 / r, efc_aref=aref, efc_pos=pos, efc_margin=margin - ) - d = d.replace(efc_frictionloss=frictionloss) + d = d.tree_replace({ + '_impl.efc_J': efc.J, + '_impl.efc_D': 1 / r, + '_impl.efc_aref': aref, + '_impl.efc_pos': pos, + '_impl.efc_margin': margin, + }) + d = d.tree_replace({'_impl.efc_frictionloss': frictionloss}) return d diff --git a/mjx/mujoco/mjx/_src/constraint_test.py b/mjx/mujoco/mjx/_src/constraint_test.py index 4e53fb13..f94873ff 100644 --- a/mjx/mujoco/mjx/_src/constraint_test.py +++ b/mjx/mujoco/mjx/_src/constraint_test.py @@ -69,16 +69,16 @@ class ConstraintTest(parameterized.TestCase): order = test_util.efc_order(m, d, dx) d_efc_j = d.efc_J.reshape((-1, m.nv)) - _assert_eq(d_efc_j, dx.efc_J[order][: d.nefc], 'efc_J') - _assert_eq(0, dx.efc_J[order][d.nefc :], 'efc_J') - _assert_eq(d.efc_aref, dx.efc_aref[order][: d.nefc], 'efc_aref') - _assert_eq(0, dx.efc_aref[order][d.nefc :], 'efc_aref') - _assert_eq(d.efc_D, dx.efc_D[order][: d.nefc], 'efc_D') - _assert_eq(d.efc_pos, dx.efc_pos[order][: d.nefc], 'efc_pos') - _assert_eq(dx.efc_pos[order][d.nefc :], 0, 'efc_pos') + _assert_eq(d_efc_j, dx._impl.efc_J[order][: d.nefc], 'efc_J') + _assert_eq(0, dx._impl.efc_J[order][d.nefc :], 'efc_J') + _assert_eq(d.efc_aref, dx._impl.efc_aref[order][: d.nefc], 'efc_aref') + _assert_eq(0, dx._impl.efc_aref[order][d.nefc :], 'efc_aref') + _assert_eq(d.efc_D, dx._impl.efc_D[order][: d.nefc], 'efc_D') + _assert_eq(d.efc_pos, dx._impl.efc_pos[order][: d.nefc], 'efc_pos') + _assert_eq(dx._impl.efc_pos[order][d.nefc :], 0, 'efc_pos') _assert_eq( d.efc_frictionloss, - dx.efc_frictionloss[order][: d.nefc], + dx._impl.efc_frictionloss[order][: d.nefc], 'efc_frictionloss', ) @@ -104,7 +104,7 @@ class ConstraintTest(parameterized.TestCase): self.assertEqual(nl, 0) self.assertEqual(nc, 0) dx = constraint.make_constraint(mjx.put_model(m), mjx.make_data(m)) - self.assertEqual(dx.efc_J.shape[0], 0) + self.assertEqual(dx._impl.efc_J.shape[0], 0) def test_disable_equality(self): m = test_util.load_test_file('constraints.xml') @@ -115,7 +115,9 @@ class ConstraintTest(parameterized.TestCase): self.assertEqual(nl, 5) self.assertEqual(nc, 180) dx = constraint.make_constraint(mjx.put_model(m), mjx.make_data(m)) - self.assertEqual(dx.efc_J.shape[0], 187) # only joint/tendon limit, contact + self.assertEqual( + dx._impl.efc_J.shape[0], 187 + ) # only joint/tendon limit, contact def test_disable_contact(self): m = test_util.load_test_file('constraints.xml') @@ -126,7 +128,9 @@ class ConstraintTest(parameterized.TestCase): self.assertEqual(nl, 5) self.assertEqual(nc, 0) dx = constraint.make_constraint(mjx.put_model(m), mjx.make_data(m)) - self.assertEqual(dx.efc_J.shape[0], 27) # only equality, joint/tendon limit + self.assertEqual( + dx._impl.efc_J.shape[0], 27 + ) # only equality, joint/tendon limit def test_disable_frictionloss(self): m = test_util.load_test_file('constraints.xml') @@ -137,7 +141,7 @@ class ConstraintTest(parameterized.TestCase): self.assertEqual(nl, 5) self.assertEqual(nc, 180) dx = constraint.make_constraint(mjx.put_model(m), mjx.make_data(m)) - self.assertEqual(dx.efc_J.shape[0], 205) + self.assertEqual(dx._impl.efc_J.shape[0], 205) def test_margin(self): """Test margin.""" @@ -163,8 +167,8 @@ class ConstraintTest(parameterized.TestCase): dx = mjx.make_constraint(mx, dx) order = test_util.efc_order(m, d, dx) - _assert_eq(d.efc_pos, dx.efc_pos[order][: d.nefc], 'efc_pos') - _assert_eq(d.efc_margin, dx.efc_margin[order][: d.nefc], 'efc_margin') + _assert_eq(d.efc_pos, dx._impl.efc_pos[order][: d.nefc], 'efc_pos') + _assert_eq(d.efc_margin, dx._impl.efc_margin[order][: d.nefc], 'efc_margin') if __name__ == '__main__': diff --git a/mjx/mujoco/mjx/_src/dataclasses.py b/mjx/mujoco/mjx/_src/dataclasses.py index 85a18e60..c2d8fb63 100644 --- a/mjx/mujoco/mjx/_src/dataclasses.py +++ b/mjx/mujoco/mjx/_src/dataclasses.py @@ -18,6 +18,7 @@ import copy import dataclasses import typing from typing import Dict, Optional, Sequence, Tuple, TypeVar, Union +import warnings import jax import numpy as np @@ -49,6 +50,36 @@ def dataclass(clz: _T, register_as_pytree: bool) -> _T: """ data_clz = dataclasses.dataclass(frozen=True)(clz) data_clz.replace = dataclasses.replace + # def replace(self, **updates): + # """Returns a new object replacing the specified fields with new values.""" + # if not hasattr(self, '_impl'): + # return dataclasses.replace(self, **updates) + + # # Private fields under `_impl` are allowed to be replaced directly as if + # # they were on the base class. This logic will be removed in a future + # # release. + # impl_updates = {} + # for k in tuple(updates.keys()): + # # Recall that getattr is overridden for '_impl' fields. + # hasattr_ = k in self.__annotations__ + # if not hasattr_ and hasattr(self._impl, k): # pylint: disable=protected-access + # impl_updates[k] = updates[k] + # del updates[k] + + # if impl_updates: + # updates['_impl'] = self._impl.replace(**impl_updates) + # warnings.warn( + # f'Accessing/replacing fields `{tuple(impl_updates.keys())}` directly' + # f' from `{self.__class__.__name__}` will be deprecated. Refrain from' + # ' using private fields that were moved to' + # f' `{self.__class__.__name__}`._impl.', + # DeprecationWarning, + # stacklevel=2, + # ) + + # return dataclasses.replace(self, **updates) + + # data_clz.replace = replace if register_as_pytree: meta_fields, data_fields = [], [] diff --git a/mjx/mujoco/mjx/_src/derivative.py b/mjx/mujoco/mjx/_src/derivative.py index b0fc65f7..5ca84a42 100644 --- a/mjx/mujoco/mjx/_src/derivative.py +++ b/mjx/mujoco/mjx/_src/derivative.py @@ -40,7 +40,9 @@ def deriv_smooth_vel(m: Model, d: Data) -> Optional[jax.Array]: gain_vel = m.actuator_gainprm[:, 2] * affine_gain ctrl = d.ctrl.at[m.actuator_dyntype != DynType.NONE].set(d.act) vel = bias_vel + gain_vel * ctrl - qderiv = d.actuator_moment.T @ jax.vmap(jp.multiply)(d.actuator_moment, vel) + qderiv = d._impl.actuator_moment.T @ jax.vmap(jp.multiply)( + d._impl.actuator_moment, vel + ) # qDeriv += d qfrc_passive / d qvel if not m.opt.disableflags & DisableBit.PASSIVE: @@ -49,9 +51,9 @@ def deriv_smooth_vel(m: Model, d: Data) -> Optional[jax.Array]: else: qderiv -= jp.diag(m.dof_damping) if m.ntendon: - qderiv -= d.ten_J.T @ jp.diag(m.tendon_damping) @ d.ten_J + qderiv -= d._impl.ten_J.T @ jp.diag(m.tendon_damping) @ d._impl.ten_J # TODO(robotics-simulation): fluid drag model - if m.opt.has_fluid_params: + if m.opt.has_fluid_params: # pytype: disable=attribute-error raise NotImplementedError('fluid drag not supported for implicitfast') # TODO(team): rne derivative diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py index 967aa639..62535122 100644 --- a/mjx/mujoco/mjx/_src/forward.py +++ b/mjx/mujoco/mjx/_src/forward.py @@ -33,12 +33,14 @@ from mujoco.mjx._src import support # pylint: disable=g-importing-member from mujoco.mjx._src.types import BiasType from mujoco.mjx._src.types import Data +from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import DynType from mujoco.mjx._src.types import GainType from mujoco.mjx._src.types import IntegratorType from mujoco.mjx._src.types import JointType from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import ModelJAX from mujoco.mjx._src.types import TrnType # pylint: enable=g-importing-member import numpy as np @@ -81,10 +83,13 @@ def fwd_position(m: Model, d: Data) -> Data: @named_scope def fwd_velocity(m: Model, d: Data) -> Data: """Velocity-dependent computations.""" - d = d.replace( - actuator_velocity=d.actuator_moment @ d.qvel, - ten_velocity=d.ten_J @ d.qvel, - ) + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('fwd_velocity requires JAX backend implementation.') + + d = d.tree_replace({ + '_impl.actuator_velocity': d._impl.actuator_moment @ d.qvel, + '_impl.ten_velocity': d._impl.ten_J @ d.qvel, + }) d = smooth.com_vel(m, d) d = passive.passive(m, d) d = smooth.rne(m, d) @@ -173,8 +178,8 @@ def fwd_actuation(m: Model, d: Data) -> Data: m.actuator_gainprm, m.actuator_biastype, m.actuator_biasprm, - d.actuator_length, - d.actuator_velocity, + d._impl.actuator_length, + d._impl.actuator_velocity, ctrl_act, jp.array(m.actuator_lengthrange), jp.array(m.actuator_acc0), @@ -215,7 +220,7 @@ def fwd_actuation(m: Model, d: Data) -> Data: ) force = jp.clip(force, forcerange[:, 0], forcerange[:, 1]) - qfrc_actuator = d.actuator_moment.T @ force + qfrc_actuator = d._impl.actuator_moment.T @ force if m.ngravcomp: # actuator-level gravity compensation, skip if added as passive force @@ -331,13 +336,17 @@ def _advance( @named_scope def euler(m: Model, d: Data) -> Data: """Euler integrator, semi-implicit in velocity.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('euler requires JAX backend implementation.') + # integrate damping implicitly qacc = d.qacc if not m.opt.disableflags & DisableBit.EULERDAMP: if support.is_sparse(m): - dh = d.replace(qM=d.qM.at[m.dof_Madr].add(m.opt.timestep * m.dof_damping)) + qM = d._impl.qM.at[m.dof_Madr].add(m.opt.timestep * m.dof_damping) else: - dh = d.replace(qM=d.qM + jp.diag(m.opt.timestep * m.dof_damping)) + qM = d._impl.qM + jp.diag(m.opt.timestep * m.dof_damping) + dh = d.tree_replace({'_impl.qM': qM}) dh = smooth.factor_m(m, dh) qfrc = d.qfrc_smooth + d.qfrc_constraint qacc = smooth.solve_m(m, dh, qfrc) @@ -392,13 +401,15 @@ def rungekutta4(m: Model, d: Data) -> Data: @named_scope def implicit(m: Model, d: Data) -> Data: """Integrates fully implicit in velocity.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('implicit requires JAX backend implementation.') qderiv = derivative.deriv_smooth_vel(m, d) qacc = d.qacc if qderiv is not None: # TODO(robotics-simulation): use smooth.factor_m / solve_m here: - qm = support.full_m(m, d) if support.is_sparse(m) else d.qM + qm = support.full_m(m, d) if support.is_sparse(m) else d._impl.qM qm -= m.opt.timestep * qderiv qh, _ = jax.scipy.linalg.cho_factor(qm) qfrc = d.qfrc_smooth + d.qfrc_constraint @@ -410,6 +421,9 @@ def implicit(m: Model, d: Data) -> Data: @named_scope def forward(m: Model, d: Data) -> Data: """Forward dynamics.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('forward requires JAX backend implementation.') + d = fwd_position(m, d) d = sensor.sensor_pos(m, d) d = fwd_velocity(m, d) @@ -418,7 +432,7 @@ def forward(m: Model, d: Data) -> Data: d = fwd_acceleration(m, d) d = sensor.sensor_acc(m, d) - if d.efc_J.size == 0: + if d._impl.efc_J.size == 0: d = d.replace(qacc=d.qacc_smooth) return d diff --git a/mjx/mujoco/mjx/_src/inverse.py b/mjx/mujoco/mjx/_src/inverse.py index 5ad3e1ad..5df72de9 100644 --- a/mjx/mujoco/mjx/_src/inverse.py +++ b/mjx/mujoco/mjx/_src/inverse.py @@ -69,16 +69,16 @@ def inv_constraint(m: Model, d: Data) -> Data: """Inverse constraint solver.""" # no constraints - if d.efc_J.size == 0: + if d._impl.efc_J.size == 0: # pytype: disable=attribute-error return d.replace(qfrc_constraint=jp.zeros(m.nv)) # update ctx = solver.Context.create(m, d, grad=False) - return d.replace( - qfrc_constraint=ctx.qfrc_constraint, - efc_force=ctx.efc_force, - ) + return d.tree_replace({ + 'qfrc_constraint': ctx.qfrc_constraint, + '_impl.efc_force': ctx.efc_force, + }) def inverse(m: Model, d: Data) -> Data: diff --git a/mjx/mujoco/mjx/_src/inverse_test.py b/mjx/mujoco/mjx/_src/inverse_test.py index 556f2a7e..227406e3 100644 --- a/mjx/mujoco/mjx/_src/inverse_test.py +++ b/mjx/mujoco/mjx/_src/inverse_test.py @@ -15,12 +15,10 @@ """Tests for inverse dynamics functions.""" from absl.testing import absltest from absl.testing import parameterized -import jax from jax import numpy as jp import mujoco from mujoco import mjx from mujoco.mjx._src import support -from mujoco.mjx._src import test_util import numpy as np # tolerance for difference between MuJoCo and MJX calculations - mostly diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 77b1655c..1d2ad234 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -15,7 +15,10 @@ """Functions to initialize, load, or save data.""" import copy -from typing import List, Tuple, Union +import logging +import os +from typing import Any, Dict, List, Optional, Tuple, Union +import warnings import jax from jax import numpy as jp @@ -29,6 +32,130 @@ import numpy as np import scipy +def _is_cuda_gpu_device(device: jax.Device) -> bool: + try: + cuda_devices = jax.devices('cuda') + except RuntimeError: + logging.info('No CUDA GPU devices found in jax.devices("cuda").') + return False + return device in cuda_devices + + +def _resolve_backend_impl( + device: jax.Device, +) -> types.BackendImpl: + """Pick a default backend impl based on the device specified.""" + if _is_cuda_gpu_device(device): + # TODO(btaba): Remove flag once Warp is ready to launch. + mjx_warp_enabled = os.environ.get('MJX_WARP_ENABLED', 'f').lower() == 'true' + if mjx_warp_enabled: + logging.debug('Picking default backend implementation: Warp.') + return types.BackendImpl.WARP + logging.info('MJX Warp is disabled via MJX_WARP_ENABLED=false.') + + if device.platform in ('gpu', 'tpu'): + logging.debug('Picking default backend implementation: JAX.') + return types.BackendImpl.JAX + + if device.platform == 'cpu': + mjx_c_default = ( + os.environ.get('MJX_C_DEFAULT_ENABLED', 'f').lower() == 'true' + ) + if mjx_c_default: + logging.debug('Picking default backend implementation: C.') + return types.BackendImpl.C + return types.BackendImpl.JAX + + raise ValueError(f'Unsupported device: {device}') + + +def _resolve_device( + backend_impl: types.BackendImpl, +) -> jax.Device: + """Resolves a device based on the backend implementation.""" + backend_impl = types.BackendImpl(backend_impl) + if backend_impl == types.BackendImpl.JAX: + device_0 = jax.devices()[0] + logging.debug('Picking default device: %s.', device_0) + return device_0 + + if backend_impl == types.BackendImpl.C: + cpu_0 = jax.devices('cpu')[0] + logging.debug('Picking default device: %s', cpu_0) + return cpu_0 + + if backend_impl == types.BackendImpl.WARP: + # WARP backend requires a CUDA GPU. + cuda_gpus = [d for d in jax.devices('cuda')] + if not cuda_gpus: + raise AssertionError( + 'No CUDA GPU devices found in' + f' jax.devices("cuda")={jax.devices("cuda")}.' + ) + + logging.debug('Picking default device: %s', cuda_gpus[0]) + return cuda_gpus[0] + + raise ValueError(f'Unsupported backend implementation: {backend_impl}') + + +def _check_backend_impl_device_compatibility( + backend_impl: Union[str, types.BackendImpl], + device: jax.Device, +) -> None: + """Checks that the backend implementation is compatible with the device.""" + if backend_impl is None: + raise ValueError('No backend implementation specified.') + + backend_impl = types.BackendImpl(backend_impl) + + if backend_impl == types.BackendImpl.WARP: + if not _is_cuda_gpu_device(device): + raise AssertionError( + 'Warp backend implementation requires a CUDA GPU device, got ' + f'{device}.' + ) + + mjx_warp_enabled = os.environ.get('MJX_WARP_ENABLED', 'f').lower() == 'true' + if not mjx_warp_enabled: + raise AssertionError( + 'Warp backend implementation is disabled via MJX_WARP_ENABLED=false.' + ) + + is_cpu_device = device.platform == 'cpu' + if backend_impl == types.BackendImpl.C: + if not is_cpu_device: + raise AssertionError( + f'C backend implementation requires a CPU device, got {device}.' + ) + + # NB: JAX backend works with any device. + + +def _resolve_backend_impl_and_device( + backend_impl: Optional[Union[str, types.BackendImpl]], + device: Optional[jax.Device] = None, +) -> Tuple[types.BackendImpl, jax.Device]: + """Resolves a backend implementation and device.""" + if backend_impl: + backend_impl = types.BackendImpl(backend_impl) + + has_backend_impl, has_device = backend_impl is not None, device is not None + if (has_backend_impl, has_device) == (True, True): + pass + elif (has_backend_impl, has_device) == (True, False): + device = _resolve_device(backend_impl) + elif (has_backend_impl, has_device) == (False, True): + backend_impl = _resolve_backend_impl(device) + else: + device = jax.devices(jax.default_backend())[0] + logging.info('Using JAX default device: %s.', device) + backend_impl = _resolve_backend_impl(device) + + _check_backend_impl_device_compatibility(backend_impl, device) + return backend_impl, device # pytype: disable=bad-return-type + + def _strip_weak_type(tree): def f(leaf): if isinstance(leaf, jax.Array): @@ -38,32 +165,27 @@ def _strip_weak_type(tree): return jax.tree_util.tree_map(f, tree) -def _make_option( - o: mujoco.MjOption, _full_compat: bool = False # pylint: disable=invalid-name +def _put_option( + o: mujoco.MjOption, + backend_impl: types.BackendImpl, + impl_fields: Optional[dict[str, Any]] = None, ) -> types.Option: """Returns mjx.Option given mujoco.MjOption.""" - if not _full_compat: - if o.integrator not in set(types.IntegratorType): - raise NotImplementedError(f'{mujoco.mjtIntegrator(o.integrator)}') + if o.integrator not in set(types.IntegratorType): + raise NotImplementedError(f'{mujoco.mjtIntegrator(o.integrator)}') - if o.cone not in set(types.ConeType): - raise NotImplementedError(f'{mujoco.mjtCone(o.cone)}') + if o.cone not in set(types.ConeType): + raise NotImplementedError(f'{mujoco.mjtCone(o.cone)}') - if o.jacobian not in set(types.JacobianType): - raise NotImplementedError(f'{mujoco.mjtJacobian(o.jacobian)}') + if o.jacobian not in set(types.JacobianType): + raise NotImplementedError(f'{mujoco.mjtJacobian(o.jacobian)}') - if o.solver not in set(types.SolverType): - raise NotImplementedError(f'{mujoco.mjtSolver(o.solver)}') + if o.solver not in set(types.SolverType): + raise NotImplementedError(f'{mujoco.mjtSolver(o.solver)}') - for i in range(mujoco.mjtEnableBit.mjNENABLE): - if o.enableflags & 2**i and 2**i not in set(types.EnableBit): - raise NotImplementedError(f'{mujoco.mjtEnableBit(2**i)}') - - has_fluid_params = o.density > 0 or o.viscosity > 0 or o.wind.any() - implicitfast = o.integrator == mujoco.mjtIntegrator.mjINT_IMPLICITFAST - if not _full_compat: - if implicitfast and has_fluid_params: - raise NotImplementedError('implicitfast not implemented for fluid drag.') + for i in range(mujoco.mjtEnableBit.mjNENABLE): + if o.enableflags & 2**i and 2 ** i not in set(types.EnableBit): + raise NotImplementedError(f'{mujoco.mjtEnableBit(2**i)}') fields = {f.name: getattr(o, f.name, None) for f in types.Option.fields()} fields['integrator'] = types.IntegratorType(o.integrator) @@ -71,13 +193,25 @@ def _make_option( fields['jacobian'] = types.JacobianType(o.jacobian) fields['solver'] = types.SolverType(o.solver) fields['disableflags'] = types.DisableBit(o.disableflags) - fields['has_fluid_params'] = has_fluid_params fields['enableflags'] = types.EnableBit(o.enableflags) - return types.Option(**fields) + if backend_impl == types.BackendImpl.JAX: + has_fluid_params = o.density > 0 or o.viscosity > 0 or o.wind.any() + implicitfast = o.integrator == mujoco.mjtIntegrator.mjINT_IMPLICITFAST + if implicitfast and has_fluid_params: + raise NotImplementedError('implicitfast not implemented for fluid drag.') + fields['has_fluid_params'] = has_fluid_params + return types.OptionJAX(**fields, **(impl_fields or {})) + + if backend_impl == types.BackendImpl.C: + c_field_keys = types.OptionC.__annotations__.keys() - fields.keys() + c_fields = {k: getattr(o, k, None) for k in c_field_keys} + return types.OptionC(**fields, **c_fields, **(impl_fields or {})) + + raise NotImplementedError(f'Unsupported backend: {backend_impl}') -def _make_statistic(s: mujoco.MjStatistic) -> types.Statistic: +def _put_statistic(s: mujoco.MjStatistic) -> types.Statistic: """Puts mujoco.MjStatistic onto a device, resulting in mjx.Statistic.""" return types.Statistic( meaninertia=s.meaninertia, @@ -88,26 +222,16 @@ def _make_statistic(s: mujoco.MjStatistic) -> types.Statistic: ) -def put_model( - m: mujoco.MjModel, device=None, _full_compat: bool = False # pylint: disable=invalid-name +def _put_model_jax( + m: mujoco.MjModel, + device: Optional[jax.Device] = None, ) -> types.Model: - """Puts mujoco.MjModel onto a device, resulting in mjx.Model. - - Args: - m: the model to put onto device - device: which device to use - if unspecified picks the default device - _full_compat: put all MjModel fields onto device irrespective of MJX support - This is an experimental feature. Avoid using it for now. - - Returns: - an mjx.Model placed on device - """ - + """Puts mujoco.MjModel onto a device, resulting in mjx.Model.""" mesh_geomid = set() for g1, g2, ip in collision_driver.geom_pairs(m): t1, t2 = m.geom_type[[g1, g2]] # check collision function exists for type pair - if not collision_driver.has_collision_fn(t1, t2) and not _full_compat: + if not collision_driver.has_collision_fn(t1, t2): t1, t2 = mujoco.mjtGeom(t1), mujoco.mjtGeom(t2) raise NotImplementedError(f'({t1}, {t2}) collisions not implemented.') # margin/gap not supported for meshes and height fields @@ -117,7 +241,7 @@ def put_model( margin = m.pair_margin[ip] else: margin = m.geom_margin[g1] + m.geom_margin[g2] - if margin.any() and not _full_compat: + if margin.any(): t1, t2 = mujoco.mjtGeom(t1), mujoco.mjtGeom(t2) raise NotImplementedError(f'({t1}, {t2}) margin/gap not implemented.') for t, g in [(t1, g1), (t2, g2)]: @@ -151,38 +275,27 @@ def put_model( (m.wrap_type, types.WrapType, mujoco.mjtWrap), ): missing = set(enum_field) - set(enum_type) - if missing and not _full_compat: + if missing: raise NotImplementedError( f'{[mj_type(m) for m in missing]} not supported' ) - mj_field_names = { - f.name - for f in types.Model.fields() - if f.metadata.get('restricted_to') != 'mjx' - } + 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} - - # zero out fields restricted to MuJoCo - if not _full_compat: - for f in types.Model.fields(): - if f.metadata.get('restricted_to') == 'mujoco' and isinstance( - fields[f.name], np.ndarray - ): - fields[f.name] = np.zeros((0,), dtype=fields[f.name].dtype) - - fields['dof_hasfrictionloss'] = fields['dof_frictionloss'] > 0 - fields['tendon_hasfrictionloss'] = fields['tendon_frictionloss'] > 0 - fields['geom_rbound_hfield'] = fields['geom_rbound'] fields['cam_mat0'] = fields['cam_mat0'].reshape((-1, 3, 3)) - fields['opt'] = _make_option(m.opt, _full_compat=_full_compat) - fields['stat'] = _make_statistic(m.stat) + fields['opt'] = _put_option(m.opt, types.BackendImpl.JAX) + fields['stat'] = _put_statistic(m.stat) + + fields_jax = {} + fields_jax['dof_hasfrictionloss'] = fields['dof_frictionloss'] > 0 + fields_jax['tendon_hasfrictionloss'] = fields['tendon_frictionloss'] > 0 + fields_jax['geom_rbound_hfield'] = fields['geom_rbound'] # spatial tendon wrap inside - fields['wrap_inside_maxiter'] = 5 - fields['wrap_inside_tolerance'] = 1.0e-4 - fields['wrap_inside_z_init'] = 1.0 - 1.0e-5 - fields['is_wrap_inside'] = np.zeros(0, dtype=bool) + fields_jax['wrap_inside_maxiter'] = 5 + fields_jax['wrap_inside_tolerance'] = 1.0e-4 + fields_jax['wrap_inside_z_init'] = 1.0 - 1.0e-5 + fields_jax['is_wrap_inside'] = np.zeros(0, dtype=bool) if m.nsite: # find sphere or cylinder geoms (if any exist) (wrap_id_geom,) = np.nonzero( @@ -198,75 +311,98 @@ def put_model( side = m.site_pos[side_id] # wrap inside flag - fields['is_wrap_inside'] = np.array( + fields_jax['is_wrap_inside'] = np.array( (np.linalg.norm(side - geom_pos, axis=1) < geom_size) & (side_id >= 0) ) # Pre-compile meshes for MJX collisions. - fields['mesh_convex'] = [None] * m.nmesh - if not _full_compat: - for i in mesh_geomid: - dataid = m.geom_dataid[i] - if fields['mesh_convex'][dataid] is None: - fields['mesh_convex'][dataid] = mesh.convex(m, dataid) # pytype: disable=unsupported-operands - fields['mesh_convex'] = tuple(fields['mesh_convex']) + fields_jax['mesh_convex'] = [None] * m.nmesh + for i in mesh_geomid: + dataid = m.geom_dataid[i] + if fields_jax['mesh_convex'][dataid] is None: + fields_jax['mesh_convex'][dataid] = mesh.convex(m, dataid) # pytype: disable=unsupported-operands + fields_jax['mesh_convex'] = tuple(fields_jax['mesh_convex']) - model = types.Model(**{k: copy.copy(v) for k, v in fields.items()}) + jax_impl = types.ModelJAX(**fields_jax) + model = types.Model( + **{k: copy.copy(v) for k, v in fields.items()}, _impl=jax_impl + ) model = jax.device_put(model, device=device) return _strip_weak_type(model) -def make_data( - m: Union[types.Model, mujoco.MjModel], - device=None, +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, backend_impl=types.BackendImpl.C) + fields['stat'] = _put_statistic(m.stat) + + 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( + m: mujoco.MjModel, + device: Optional[jax.Device] = None, + backend_impl: Optional[Union[str, types.BackendImpl]] = None, _full_compat: bool = False, # pylint: disable=invalid-name -) -> types.Data: - """Allocate and initialize Data. +) -> types.Model: + """Puts mujoco.MjModel onto a device, resulting in mjx.Model. Args: - m: the model to use + m: the model to put onto device device: which device to use - if unspecified picks the default device - _full_compat: create all MjData fields on device irrespective of MJX support - This is an experimental feature. Avoid using it for now. If using this - flag, also use _full_compat for put_model. + backend_impl: backend implementation to use + _full_compat: put all MjModel fields onto device irrespective of MJX support + This is an experimental feature. Avoid using it for now. Returns: - an initialized mjx.Data placed on device + an mjx.Model placed on device + + Raises: + ValueError: if backend_impl is not supported + DeprecationWarning: if _full_compat is True """ - dim = collision_driver.make_condim(m) - 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 - contact = types.Contact( - dist=np.zeros((ncon,), dtype=float_), - pos=np.zeros((ncon, 3), dtype=float_), - frame=np.zeros((ncon, 3, 3), dtype=float_), - includemargin=np.zeros((ncon,), dtype=float_), - friction=np.zeros((ncon, 5), dtype=float_), - solref=np.zeros((ncon, mujoco.mjNREF), dtype=float_), - solreffriction=np.zeros((ncon, mujoco.mjNREF), dtype=float_), - solimp=np.zeros((ncon, mujoco.mjNIMP), dtype=float_), - dim=dim, - # let jax pick contact.geom int precision, for interop with - # jax_enable_x64 - geom1=np.full((ncon,), -1, dtype=int_), - geom2=np.full((ncon,), -1, dtype=int_), - geom=np.full((ncon, 2), -1, dtype=int_), - 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.' + if _full_compat: + warnings.warn( + 'mjx.put_model(..., _full_compat=True) is deprecated. Use' + ' mjx.put_model(..., backend_impl=types.BackendImpl.C) instead.', + DeprecationWarning, + stacklevel=2, ) + backend_impl = types.BackendImpl.C + backend_impl, device = _resolve_backend_impl_and_device(backend_impl, device) + if backend_impl == types.BackendImpl.JAX: + return _put_model_jax(m, device) + elif backend_impl == types.BackendImpl.C: + return _put_model_c(m, device) + elif backend_impl == types.BackendImpl.WARP: + raise NotImplementedError('Warp backend not implemented yet.') + else: + raise ValueError(f'Unsupported backend implementation: {backend_impl}') + + +def _make_data_public_fields(m: types.Model) -> Dict[str, Any]: + """Create public fields for the Data object.""" + float_ = jp.zeros(1, float).dtype zero_fields = { - 'solver_niter': (int_,), 'time': (float_,), 'qvel': (m.nv, float_), 'act': (m.na, float_), @@ -293,25 +429,183 @@ def make_data( 'site_xmat': (m.nsite, 3, 3, float_), 'cam_xpos': (m.ncam, 3, float_), 'cam_xmat': (m.ncam, 3, 3, float_), - 'light_xpos': (m.nlight, 3, float_), - 'light_xdir': (m.nlight, 3, float_), 'subtree_com': (m.nbody, 3, float_), + 'actuator_force': (m.nu, float_), + 'qfrc_bias': (m.nv, float_), + 'qfrc_gravcomp': (m.nv, float_), + 'qfrc_fluid': (m.nv, float_), + 'qfrc_passive': (m.nv, float_), + 'qfrc_actuator': (m.nv, float_), + 'qfrc_smooth': (m.nv, float_), + 'qacc_smooth': (m.nv, float_), + 'qfrc_constraint': (m.nv, float_), + 'qfrc_inverse': (m.nv, float_), + 'cvel': (m.nbody, 6, float_), + } + zero_fields = { + k: np.zeros(v[:-1], dtype=v[-1]) for k, v in zero_fields.items() + } + return zero_fields + + +def _make_data_contact_jax( + condim: np.ndarray, efc_address: np.ndarray +) -> types.Contact: + """Create contact for the Data object.""" + ncon = condim.size + float_ = jp.zeros(1, float).dtype + int_ = jp.zeros(1, int).dtype + contact = types.Contact( + dist=np.zeros((ncon,), dtype=float_), + pos=np.zeros((ncon, 3), dtype=float_), + frame=np.zeros((ncon, 3, 3), dtype=float_), + includemargin=np.zeros((ncon,), dtype=float_), + friction=np.zeros((ncon, 5), dtype=float_), + solref=np.zeros((ncon, mujoco.mjNREF), dtype=float_), + solreffriction=np.zeros((ncon, mujoco.mjNREF), dtype=float_), + solimp=np.zeros((ncon, mujoco.mjNIMP), dtype=float_), + dim=condim, + # let jax pick contact.geom int precision, for interop with + # jax_enable_x64 + geom1=np.full((ncon,), -1, dtype=int_), + geom2=np.full((ncon,), -1, dtype=int_), + geom=np.full((ncon, 2), -1, dtype=int_), + efc_address=efc_address, + ) + return contact + + +def _make_data_jax( + m: types.Model, + device: Optional[jax.Device] = None, +) -> types.Data: + """Allocate and initialize Data for the JAX implementation.""" + dim = collision_driver.make_condim(m) + efc_type = constraint.make_efc_type(m, dim) + ne, nf, nl, nc = constraint.counts(efc_type) + ncon, nefc = dim.size, ne + nf + nl + nc + efc_address = constraint.make_efc_address(m, dim, efc_type) + + float_ = jp.zeros(1, float).dtype + int_ = jp.zeros(1, int).dtype + contact = _make_data_contact_jax(dim, 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_impl_fields = { + 'solver_niter': (int_,), 'cdof': (m.nv, 6, float_), 'cinert': (m.nbody, 10, float_), - 'flexvert_xpos': (m.nflexvert, 3, float_), - 'flexelem_aabb': (m.nflexelem, 6, float_), - 'flexedge_J_rownnz': (m.nflexedge, np.int32), - 'flexedge_J_rowadr': (m.nflexedge, np.int32), - 'flexedge_J_colind': (m.nflexedge, m.nv, np.int32), - 'flexedge_J': (m.nflexedge, m.nv, float_), - 'flexedge_length': (m.nflexedge, float_), 'ten_wrapadr': (m.ntendon, np.int32), 'ten_wrapnum': (m.ntendon, np.int32), + 'ten_J': (m.ntendon, m.nv, float_), + 'ten_length': (m.ntendon, float_), + 'wrap_obj': (m.nwrap, 2, np.int32), + 'wrap_xpos': (m.nwrap, 6, float_), + 'actuator_length': (m.nu, float_), + 'actuator_moment': (m.nu, m.nv, float_), + 'crb': (m.nbody, 10, float_), + 'qM': (m.nM, float_) if support.is_sparse(m) else (m.nv, m.nv, float_), + 'qLD': (m.nM, float_) if support.is_sparse(m) else (m.nv, m.nv, float_), + 'qLDiagInv': (m.nv, float_) if support.is_sparse(m) else (0, float_), + 'ten_velocity': (m.ntendon, float_), + 'actuator_velocity': (m.nu, float_), + 'cdof_dot': (m.nv, 6, 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.DataJAX( + 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 _make_data_c( + m: types.Model, + 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) + 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) + + nflexvert = m.nflexvert + nflexedge = m.nflexedge + nflexelem = m.nflexelem + nbvh = m.nbvh + nbvhdynamic = m.nbvhdynamic + zero_impl_fields = { + 'solver_niter': (int_,), + 'cdof': (m.nv, 6, float_), + '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_rownnz': (nflexedge, np.int32), + 'flexedge_J_rowadr': (nflexedge, np.int32), + 'flexedge_J_colind': (nflexedge, m.nv, np.int32), + 'flexedge_J': (nflexedge, m.nv, float_), + 'flexedge_length': (nflexedge, float_), 'ten_J_rownnz': (m.ntendon, np.int32), 'ten_J_rowadr': (m.ntendon, np.int32), 'ten_J_colind': (m.ntendon, m.nv, np.int32), 'ten_J': (m.ntendon, m.nv, float_), 'ten_length': (m.ntendon, 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_), 'actuator_length': (m.nu, float_), @@ -319,27 +613,17 @@ def make_data( 'moment_rowadr': (m.nu, np.int32), 'moment_colind': (m.nJmom, np.int32), 'actuator_moment': (m.nu, m.nv, float_), + 'bvh_aabb_dyn': (nbvhdynamic, 6, float_), + 'bvh_active': (nbvh, np.uint8), + 'flexedge_velocity': (nflexedge, float_), 'crb': (m.nbody, 10, float_), - 'qM': (m.nM, float_) if support.is_sparse(m) else (m.nv, m.nv, float_), - 'qLD': (m.nM, float_) if support.is_sparse(m) else (m.nv, m.nv, float_), - 'qLDiagInv': (m.nv, float_) if support.is_sparse(m) else (0, float_), - 'bvh_aabb_dyn': (m.nbvhdynamic, 6, float_), - 'bvh_active': (m.nbvh, np.uint8), - 'flexedge_velocity': (m.nflexedge, float_), + 'qM': (m.nM, float_), + 'qLD': (m.nM, float_), + 'qH': (m.nM, float_), + 'qHDiagInv': (m.nv, float_), + 'qLDiagInv': (m.nv, float_), 'ten_velocity': (m.ntendon, float_), 'actuator_velocity': (m.nu, float_), - 'cvel': (m.nbody, 6, float_), - 'cdof_dot': (m.nv, 6, float_), - 'qfrc_bias': (m.nv, float_), - 'qfrc_spring': (m.nv, float_), - 'qfrc_damper': (m.nv, float_), - 'qfrc_gravcomp': (m.nv, float_), - 'qfrc_fluid': (m.nv, float_), - 'qfrc_passive': (m.nv, float_), - 'subtree_linvel': (m.nbody, 3, float_), - 'subtree_angmom': (m.nbody, 3, float_), - 'qH': (m.nM, float_) if support.is_sparse(m) else (m.nv, m.nv, float_), - 'qHDiagInv': (m.nv, float_), 'B_rownnz': (m.nbody, np.int32), 'B_rowadr': (m.nbody, np.int32), 'B_colind': (m.nB, np.int32), @@ -359,15 +643,14 @@ def make_data( 'mapD2M': (m.nM, np.int32), 'qDeriv': (m.nD, float_), 'qLU': (m.nD, float_), - 'actuator_force': (m.nu, float_), - 'qfrc_actuator': (m.nv, float_), - 'qfrc_smooth': (m.nv, float_), - 'qacc_smooth': (m.nv, float_), - 'qfrc_constraint': (m.nv, float_), - 'qfrc_inverse': (m.nv, float_), + 'qfrc_spring': (m.nv, float_), + 'qfrc_damper': (m.nv, float_), + 'cdof_dot': (m.nv, 6, 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_), @@ -375,31 +658,26 @@ def make_data( 'efc_D': (nefc, float_), 'efc_aref': (nefc, float_), 'efc_force': (nefc, float_), - '_qM_sparse': (m.nM, float_), - '_qLD_sparse': (m.nM, float_), - '_qLDiagInv_sparse': (m.nv, float_), } - - if not _full_compat: - for f in types.Data.fields(): - if f.metadata.get('restricted_to') in ('mujoco', 'mjx'): - zero_fields[f.name] = (0, zero_fields[f.name][-1]) - - zero_fields = { - k: np.zeros(v[:-1], dtype=v[-1]) for k, v in zero_fields.items() + zero_impl_fields = { + k: np.zeros(v[:-1], dtype=v[-1]) for k, v in zero_impl_fields.items() } - - d = types.Data( + impl = types.DataC( ne=ne, nf=nf, nl=nl, nefc=nefc, ncon=ncon, - qpos=jp.array(m.qpos0, dtype=float_), contact=contact, efc_type=efc_type, + **zero_impl_fields, + ) + + d = types.Data( + qpos=jp.array(m.qpos0, dtype=float_), eq_active=m.eq_active0, - **zero_fields, + _impl=impl, + **_make_data_public_fields(m), ) if m.nmocap: @@ -413,140 +691,62 @@ def make_data( ) d = jax.device_put(d, device=device) - return d -def _get_contact(c: mujoco._structs._MjContactList, cx: types.Contact): - """Converts mjx.Contact to mujoco._structs._MjContactList.""" - con_id = np.nonzero(cx.dist <= 0)[0] - for field in types.Contact.fields(): - value = getattr(cx, field.name)[con_id] - if field.name == 'frame': - value = value.reshape((-1, 9)) - getattr(c, field.name)[:] = value +def make_data( + m: Union[types.Model, mujoco.MjModel], + device: Optional[jax.Device] = None, + backend_impl: Optional[Union[str, types.BackendImpl]] = None, + _full_compat: bool = False, # pylint: disable=invalid-name +) -> types.Data: + """Allocate and initialize Data. + + Args: + m: the model to use + device: which device to use - if unspecified picks the default device + backend_impl: backend implementation to use + _full_compat: put all fields onto device irrespective of MJX support This is + an experimental feature. Avoid using it for now. If using this flag, also + use _full_compat for put_model. + + Returns: + an initialized mjx.Data placed on device + + Raises: + ValueError: if the model's backend_impl does not match the make_data + backend_impl + NotImplementedError: if the backend_impl is not implemented yet + DeprecationWarning: if _full_compat is used + """ + if _full_compat: + warnings.warn( + 'mjx.make_data(..., _full_compat=True) is deprecated. Use' + ' mjx.make_data(..., backend_impl=types.BackendImpl.C) instead.', + DeprecationWarning, + stacklevel=2, + ) + backend_impl = types.BackendImpl.C + + backend_impl, device = _resolve_backend_impl_and_device(backend_impl, device) + + if isinstance(m, types.Model) and m.backend_impl != backend_impl: + raise ValueError( + f'Model backend_impl {m.backend_impl} does not match make_data ' + f'backend_impl {backend_impl}.' + ) + + if backend_impl == types.BackendImpl.JAX: + return _make_data_jax(m, device) + elif backend_impl == types.BackendImpl.C: + return _make_data_c(m, device) + + raise NotImplementedError( + f'make_data for backend_impl "{backend_impl}" not implemented yet.' + ) -def get_data( - m: mujoco.MjModel, d: types.Data -) -> Union[mujoco.MjData, List[mujoco.MjData]]: - """Gets mjx.Data from a device, resulting in mujoco.MjData or List[MjData].""" - batched = len(d.qpos.shape) > 1 - batch_size = d.qpos.shape[0] if batched else 1 - - if batched: - result = [mujoco.MjData(m) for _ in range(batch_size)] - else: - result = mujoco.MjData(m) - - get_data_into(result, m, d) - - return result - - -def get_data_into( - result: Union[mujoco.MjData, List[mujoco.MjData]], - m: mujoco.MjModel, - d: types.Data, -): - """Gets mjx.Data from a device into an existing mujoco.MjData or list.""" - batched = isinstance(result, list) - if batched and len(d.qpos.shape) < 2: - raise ValueError('dst is a list, but d is not batched.') - if not batched and len(d.qpos.shape) >= 2: - raise ValueError('dst is a an MjData, but d is batched.') - - d = jax.device_get(d) - - batch_size = d.qpos.shape[0] if batched else 1 - - dof_i, dof_j = [], [] - for i in range(m.nv): - j = i - while j > -1: - dof_i.append(i) - dof_j.append(j) - j = m.dof_parentid[j] - - for i in range(batch_size): - d_i = jax.tree_util.tree_map(lambda x, i=i: x[i], d) if batched else d - result_i = result[i] if batched else result - ncon = (d_i.contact.dist <= 0).sum() - efc_active = (d_i.efc_J != 0).any(axis=1) - nefc = int(efc_active.sum()) - result_i.nJ = nefc * m.nv - if ncon != result_i.ncon or nefc != result_i.nefc: - mujoco._functions._realloc_con_efc(result_i, ncon=ncon, nefc=nefc) # pylint: disable=protected-access - result_i.efc_J_rownnz[:] = np.repeat(m.nv, nefc) - result_i.efc_J_rowadr[:] = np.arange(0, nefc * m.nv, m.nv) - result_i.efc_J_colind[:] = np.tile(np.arange(m.nv), nefc) - - for field in types.Data.fields(): - restricted_to = field.metadata.get('restricted_to') - if restricted_to == 'mjx': - continue - - if field.name == 'contact': - _get_contact(result_i.contact, d_i.contact) - # efc_address must be updated because rows were deleted above: - efc_map = np.cumsum(efc_active) - 1 - result_i.contact.efc_address[:] = efc_map[result_i.contact.efc_address] - continue - - # MuJoCo actuator_moment is sparse, MJX uses a dense representation. - if field.name == 'actuator_moment': - moment_rownnz = np.zeros(m.nu, dtype=np.int32) - moment_rowadr = np.zeros(m.nu, dtype=np.int32) - moment_colind = np.zeros(m.nJmom, dtype=np.int32) - actuator_moment = np.zeros(m.nJmom) - if m.nu: - mujoco.mju_dense2sparse( - actuator_moment, - d_i.actuator_moment, - moment_rownnz, - moment_rowadr, - moment_colind, - ) - result_i.moment_rownnz[:] = moment_rownnz - result_i.moment_rowadr[:] = moment_rowadr - result_i.moment_colind[:] = moment_colind - result_i.actuator_moment[:] = actuator_moment - continue - - value = getattr(d_i, field.name) - - if field.name in ('nefc', 'ncon'): - value = {'nefc': nefc, 'ncon': ncon}[field.name] - elif field.name.endswith('xmat') or field.name == 'ximat': - value = value.reshape((-1, 9)) - elif field.name.startswith('efc_'): - value = value[efc_active] - if field.name == 'efc_J': - value = value.reshape(-1) - elif field.name == 'qM' and not support.is_sparse(m): - value = value[dof_i, dof_j] - elif field.name == 'qLD' and not support.is_sparse(m): - # TODO(erikfrey): provide correct qLDs - value = np.zeros(m.nM) - elif field.name == 'qLDiagInv' and not support.is_sparse(m): - value = np.ones(m.nv) - - if isinstance(value, np.ndarray) and value.shape: - if restricted_to in ('mujoco', 'mjx'): - continue # don't copy fields that are mujoco-only or MJX-only - else: - result_field = getattr(result_i, field.name) - if result_field.shape != value.shape: - raise ValueError( - f'Input field {field.name} has shape {value.shape}, but output' - f' has shape {result_field.shape}' - ) - result_field[:] = value - else: - setattr(result_i, field.name, value) - - -def _make_contact( +def _put_contact( c: mujoco._structs._MjContactList, dim: np.ndarray, efc_address: np.ndarray, @@ -581,22 +781,24 @@ def _make_contact( return types.Contact(**fields), contact_map -def put_data( - m: mujoco.MjModel, d: mujoco.MjData, device=None, _full_compat: bool = False # pylint: disable=invalid-name +def _put_data_public_fields(d: mujoco.MjData) -> Dict[str, Any]: + """Returns public fields from mujoco.MjData in a dictionary.""" + fields = { + f.name: getattr(d, f.name) + for f in types.Data.fields() + if f.name != '_impl' + } + # MJX uses square matrices for these fields: + for fname in ('xmat', 'ximat', 'geom_xmat', 'site_xmat', 'cam_xmat'): + fields[fname] = fields[fname].reshape((-1, 3, 3)) + + return fields + + +def _put_data_jax( + m: mujoco.MjModel, d: mujoco.MjData, device: Optional[jax.Device] = None ) -> types.Data: - """Puts mujoco.MjData onto a device, resulting in mjx.Data. - - Args: - m: the model to use - d: the data to put on device - device: which device to use - if unspecified picks the default device - _full_compat: put all MjModel fields onto device irrespective of MJX support - This is an experimental feature. Avoid using it for now. If using this - flag, also use _full_compat for put_model. - - Returns: - an mjx.Data placed on device - """ + """Puts mujoco.MjData onto a device, resulting in mjx.Data.""" dim = collision_driver.make_condim(m) efc_type = constraint.make_efc_type(m, dim) efc_address = constraint.make_efc_address(m, dim, efc_type) @@ -613,18 +815,16 @@ def put_data( if d_val > val: raise ValueError(f'd.{name} too high, d.{name} = {d_val}, model = {val}') - fields = { + fields = _put_data_public_fields(d) + + # Implementation specific fields. + impl_fields = { f.name: getattr(d, f.name) - for f in types.Data.fields() - if f.metadata.get('restricted_to') != 'mjx' + for f in types.DataJAX.fields() + if hasattr(d, f.name) } - - # MJX prefers square matrices for these fields: - for fname in ('xmat', 'ximat', 'geom_xmat', 'site_xmat', 'cam_xmat'): - fields[fname] = fields[fname].reshape((-1, 3, 3)) - # MJX does not support islanding, so only transfer the first solver_niter - fields['solver_niter'] = fields['solver_niter'][0] + impl_fields['solver_niter'] = impl_fields['solver_niter'][0] # convert sparse representation of actuator_moment to dense matrix moment = np.zeros((m.nu, m.nv)) @@ -635,9 +835,9 @@ def put_data( d.moment_rowadr, d.moment_colind, ) - fields['actuator_moment'] = moment + impl_fields['actuator_moment'] = moment - contact, contact_map = _make_contact(d.contact, dim, efc_address) + contact, contact_map = _put_contact(d.contact, dim, efc_address) # pad efc fields: MuJoCo efc arrays are sparse for inactive constraints. # efc_J is also optionally column-sparse (typically for large nv). MJX is @@ -647,14 +847,16 @@ def put_data( efc_j = np.zeros((d.efc_J_rownnz.shape[0], m.nv)) mujoco.mju_sparse2dense( efc_j, - fields['efc_J'], + impl_fields['efc_J'], d.efc_J_rownnz, d.efc_J_rowadr, d.efc_J_colind, ) - fields['efc_J'] = efc_j + impl_fields['efc_J'] = efc_j else: - fields['efc_J'] = fields['efc_J'].reshape((-1 if m.nv else 0, m.nv)) + 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 ( @@ -671,7 +873,9 @@ def put_data( 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] = fields[fname][d_beg : d_beg + size] + 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): @@ -683,43 +887,348 @@ def put_data( 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] = fields[fname][efc_i : efc_i + num_rows] + value[efc_o : efc_o + num_rows] = impl_fields[fname][ + efc_i : efc_i + num_rows + ] - fields[fname] = value - - if _full_compat: - # full compatibility mode, we store sparse qM regardless of jacobian setting - fields['_qM_sparse'] = fields['qM'] - fields['_qLD_sparse'] = fields['qLD'] - fields['_qLDiagInv_sparse'] = fields['qLDiagInv'] - else: - fields['_qM_sparse'] = jp.zeros(0, dtype=float) - fields['_qLD_sparse'] = jp.zeros(0, dtype=float) - fields['_qLDiagInv_sparse'] = jp.zeros(0, dtype=float) - # otherwise clear out unused arrays - for f in types.Data.fields(): - if f.metadata.get('restricted_to') == 'mujoco' and isinstance( - fields[f.name], np.ndarray - ): - fields[f.name] = np.zeros(0, dtype=fields[f.name].dtype) + impl_fields[fname] = value # convert qM and qLD if jacobian is dense if not support.is_sparse(m): - fields['qM'] = np.zeros((m.nv, m.nv)) - mujoco.mj_fullM(m, fields['qM'], d.qM) + impl_fields['qM'] = np.zeros((m.nv, m.nv)) + mujoco.mj_fullM(m, impl_fields['qM'], d.qM) # TODO(erikfrey): derive L*L' from L'*D*L instead of recomputing try: - fields['qLD'], _ = scipy.linalg.cho_factor(fields['qM']) + impl_fields['qLD'], _ = scipy.linalg.cho_factor(impl_fields['qM']) except scipy.linalg.LinAlgError: # this happens when qM is empty or unstable simulation - fields['qLD'] = np.zeros((m.nv, m.nv)) - fields['qLDiagInv'] = np.zeros(0) + impl_fields['qLD'] = np.zeros((m.nv, m.nv)) + impl_fields['qLDiagInv'] = np.zeros(0) - fields['contact'] = contact - fields.update(ne=ne, nf=nf, nl=nl, nefc=nefc, ncon=ncon, efc_type=efc_type) + 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 = types.Data(**{k: copy.copy(v) for k, v in fields.items()}) + data_jax = types.DataJAX(**{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) + + +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) + 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) + } + + # TODO(stunya): support islanding via C backend impl. + impl_fields['solver_niter'] = impl_fields['solver_niter'][0] + + # TODO(btaba): remove dense actuator moment. + # convert sparse representation of actuator_moment to dense matrix + moment = np.zeros((m.nu, m.nv)) + mujoco.mju_sparse2dense( + moment, + d.actuator_moment, + d.moment_rownnz, + d.moment_rowadr, + d.moment_colind, + ) + impl_fields['actuator_moment'] = moment + + # TODO(btaba): remove reliance on JAX _put_contact. + contact, contact_map = _put_contact(d.contact, dim, efc_address) + + # TODO(btaba): 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) + + +def put_data( + m: mujoco.MjModel, + d: mujoco.MjData, + device: Optional[jax.Device] = None, + backend_impl: Optional[Union[str, types.BackendImpl]] = None, + _full_compat: bool = False, # pylint: disable=invalid-name +) -> types.Data: + """Puts mujoco.MjData onto a device, resulting in mjx.Data. + + Args: + m: the model to use + d: the data to put on device + device: which device to use - if unspecified picks the default device + backend_impl: backend implementation to use + _full_compat: put all MjModel fields onto device irrespective of MJX support + This is an experimental feature. Avoid using it for now. If using this + flag, also use _full_compat for put_model. + + Returns: + an mjx.Data placed on device + """ + if _full_compat: + warnings.warn( + 'mjx.put_data(..., _full_compat=True) is deprecated. Use' + ' mjx.put_data(..., backend_impl=types.BackendImpl.C) instead.', + DeprecationWarning, + stacklevel=2, + ) + backend_impl = types.BackendImpl.C + + backend_impl, device = _resolve_backend_impl_and_device(backend_impl, device) + if backend_impl == types.BackendImpl.JAX: + return _put_data_jax(m, d, device) + elif backend_impl == types.BackendImpl.C: + return _put_data_c(m, d, device) + + raise NotImplementedError( + f'put_data for backend_impl "{backend_impl}" not implemented yet.' + ) + + +def _get_contact(c: mujoco._structs._MjContactList, cx: types.Contact): + """Converts mjx.Contact to mujoco._structs._MjContactList.""" + con_id = np.nonzero(cx.dist <= 0)[0] + for field in types.Contact.fields(): + value = getattr(cx, field.name)[con_id] + if field.name == 'frame': + value = value.reshape((-1, 9)) + getattr(c, field.name)[:] = value + + +def _get_data_into( + result: Union[mujoco.MjData, List[mujoco.MjData]], + m: mujoco.MjModel, + d: types.Data, +): + """Gets mjx.Data from a device into an existing mujoco.MjData or list.""" + batched = isinstance(result, list) + d = jax.device_get(d) + batch_size = d.qpos.shape[0] if batched else 1 + + dof_i, dof_j = [], [] + if d.backend_impl == types.BackendImpl.JAX: + for i in range(m.nv): + j = i + while j > -1: + dof_i.append(i) + dof_j.append(j) + j = m.dof_parentid[j] + + for i in range(batch_size): + d_i = jax.tree_util.tree_map(lambda x, i=i: x[i], d) if batched else d + result_i = result[i] if batched else result + ncon = (d_i._impl.contact.dist <= 0).sum() + efc_active = (d_i._impl.efc_J != 0).any(axis=1) + nefc = int(efc_active.sum()) + result_i.nJ = nefc * m.nv + if ncon != result_i.ncon or nefc != result_i.nefc: + mujoco._functions._realloc_con_efc(result_i, ncon=ncon, nefc=nefc) # pylint: disable=protected-access + result_i.efc_J_rownnz[:] = np.repeat(m.nv, nefc) + result_i.efc_J_rowadr[:] = np.arange(0, nefc * m.nv, m.nv) + result_i.efc_J_colind[:] = np.tile(np.arange(m.nv), nefc) + + if d.backend_impl == types.BackendImpl.JAX: + all_fields = types.Data.fields() + types.DataJAX.fields() + elif d.backend_impl == types.BackendImpl.C: + all_fields = types.Data.fields() + types.DataC.fields() + else: + raise NotImplementedError( + f'get_data_into for backend_impl "{d.backend_impl}" not implemented' + ' yet.' + ) + + for field in all_fields: + if field.name not in mujoco.MjData.__dict__.keys(): + continue + + if field.name == 'contact': + _get_contact(result_i.contact, d_i._impl.contact) + # efc_address must be updated because rows were deleted above: + efc_map = np.cumsum(efc_active) - 1 + result_i.contact.efc_address[:] = efc_map[result_i.contact.efc_address] + continue + + # MuJoCo actuator_moment is sparse, MJX uses a dense representation. + if field.name == 'actuator_moment': + moment_rownnz = np.zeros(m.nu, dtype=np.int32) + moment_rowadr = np.zeros(m.nu, dtype=np.int32) + moment_colind = np.zeros(m.nJmom, dtype=np.int32) + actuator_moment = np.zeros(m.nJmom) + if m.nu: + mujoco.mju_dense2sparse( + actuator_moment, + d_i._impl.actuator_moment, + moment_rownnz, + moment_rowadr, + moment_colind, + ) + result_i.moment_rownnz[:] = moment_rownnz + result_i.moment_rowadr[:] = moment_rowadr + result_i.moment_colind[:] = moment_colind + result_i.actuator_moment[:] = actuator_moment + continue + + if hasattr(d_i._impl, field.name): + value = getattr(d_i._impl, field.name) + else: + value = getattr(d_i, field.name) + + if field.name in ('nefc', 'ncon'): + value = {'nefc': nefc, 'ncon': ncon}[field.name] + elif field.name.endswith('xmat') or field.name == 'ximat': + value = value.reshape((-1, 9)) + elif field.name.startswith('efc_'): + value = value[efc_active] + if field.name == 'efc_J': + value = value.reshape(-1) + if d.backend_impl == types.BackendImpl.JAX: + if field.name == 'qM' and not support.is_sparse(m): + value = value[dof_i, dof_j] + elif field.name == 'qLD' and not support.is_sparse(m): + # TODO(erikfrey): provide correct qLDs + value = np.zeros(m.nM) + elif field.name == 'qLDiagInv' and not support.is_sparse(m): + value = np.ones(m.nv) + + if isinstance(value, np.ndarray) and value.shape: + result_field = getattr(result_i, field.name) + if result_field.shape != value.shape: + raise ValueError( + f'Input field {field.name} has shape {value.shape}, but output' + f' has shape {result_field.shape}' + ) + result_field[:] = value + else: + setattr(result_i, field.name, value) + + +def get_data_into( + result: Union[mujoco.MjData, List[mujoco.MjData]], + m: mujoco.MjModel, + d: types.Data, +): + """Gets mjx.Data from a device into an existing mujoco.MjData or list.""" + is_batched = isinstance(result, list) + if is_batched and len(d.qpos.shape) < 2: + raise ValueError('destination is a list, but d is not batched.') + if not is_batched and len(d.qpos.shape) >= 2: + raise ValueError('destination is a an MjData, but d is batched.') + + d = jax.device_get(d) + + if d.backend_impl in (types.BackendImpl.JAX, types.BackendImpl.C): + # TODO(stunya): Split out _get_data_into once codepaths diverge enough. + return _get_data_into(result, m, d) + + raise NotImplementedError( + f'get_data_into for backend_impl "{d.backend_impl}" not implemented yet.' + ) + + +def get_data( + m: mujoco.MjModel, d: types.Data +) -> Union[mujoco.MjData, List[mujoco.MjData]]: + """Gets mjx.Data from a device, resulting in mujoco.MjData or List[MjData].""" + batched = len(d.qpos.shape) > 1 + batch_size = d.qpos.shape[0] if batched else 1 + + if batched: + result = [mujoco.MjData(m) for _ in range(batch_size)] + else: + result = mujoco.MjData(m) + + get_data_into(result, m, d) + + return result diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 494b2753..41d812d3 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -14,14 +14,19 @@ # ============================================================================== """Tests for io functions.""" +import os +from unittest import mock from absl.testing import absltest from absl.testing import parameterized import jax from jax import numpy as jp import mujoco from mujoco import mjx +from mujoco.mjx._src import io as mjx_io from mujoco.mjx._src import test_util + # pylint: disable=g-importing-member +from mujoco.mjx._src.types import BackendImpl from mujoco.mjx._src.types import ConeType # pylint: enable=g-importing-member import numpy as np @@ -95,10 +100,13 @@ _MULTIPLE_CONSTRAINTS = """ class ModelIOTest(parameterized.TestCase): """IO tests for mjx.Model.""" - @parameterized.parameters(_MULTIPLE_CONVEX_OBJECTS, _MULTIPLE_CONSTRAINTS) - def test_put_model(self, xml): + @parameterized.product( + xml=(_MULTIPLE_CONVEX_OBJECTS, _MULTIPLE_CONSTRAINTS), + backend_impl=('jax', 'c'), + ) + def test_put_model(self, xml, backend_impl): m = mujoco.MjModel.from_xml_string(xml) - mx = mjx.put_model(m) + mx = mjx.put_model(m, backend_impl=backend_impl) def assert_not_weak_type(x): if isinstance(x, jax.Array): @@ -120,8 +128,14 @@ class ModelIOTest(parameterized.TestCase): self.assertEqual(mx.nM, m.nM) self.assertAlmostEqual(mx.opt.timestep, m.opt.timestep) - # fields restricted to MuJoCo should not be populated - self.assertEqual(mx.bvh_aabb.shape, (0,)) + if backend_impl == 'jax': + # fields restricted to MuJoCo should not be populated + self.assertFalse(hasattr(mx, 'bvh_aabb')) + elif backend_impl == 'c': + # Options specific to C are populated. + self.assertEqual(mx.opt.apirate, m.opt.apirate) + # Fields private to C backend impl are populated. + self.assertTrue(hasattr(mx._impl, 'bvh_aabb')) np.testing.assert_allclose(mx.body_parentid, m.body_parentid) np.testing.assert_allclose(mx.geom_type, m.geom_type) @@ -150,7 +164,8 @@ class ModelIOTest(parameterized.TestCase): m = mjx.put_model( mujoco.MjModel.from_xml_string( '' - ) + ), + backend_impl='jax', ) self.assertTrue(m.opt.has_fluid_params) @@ -174,7 +189,8 @@ class ModelIOTest(parameterized.TestCase): def test_margin_gap_mesh_not_implemented(self): with self.assertRaises(NotImplementedError): - mjx.put_model(mujoco.MjModel.from_xml_string(""" + mjx.put_model( + mujoco.MjModel.from_xml_string(""" @@ -189,27 +205,32 @@ class ModelIOTest(parameterized.TestCase): - """)) + """), + backend_impl='jax', + ) def test_implicitfast_fluid_not_implemented(self): with self.assertRaises(NotImplementedError): - mjx.put_model(mujoco.MjModel.from_xml_string(""" + mjx.put_model( + mujoco.MjModel.from_xml_string(""" """)) + """), + backend_impl='jax', + ) def test_wrap_inside(self): m = test_util.load_test_file('tendon/wrap_sidesite.xml') - mx0 = mjx.put_model(m) + mx0 = mjx.put_model(m, backend_impl='jax') np.testing.assert_equal( - mx0.is_wrap_inside, + mx0._impl.is_wrap_inside, np.array([1, 0, 1, 0, 1, 1, 0]), ) m.site_pos[2] = m.site_pos[1] - mx1 = mjx.put_model(m) + mx1 = mjx.put_model(m, backend_impl='jax') np.testing.assert_equal( - mx1.is_wrap_inside, + mx1._impl.is_wrap_inside, np.array([0, 0, 1, 0, 1, 0, 0]), ) @@ -217,11 +238,11 @@ class ModelIOTest(parameterized.TestCase): class DataIOTest(parameterized.TestCase): """IO tests for mjx.Data.""" - def test_make_data(self): + @parameterized.parameters('jax', 'c') + def test_make_data(self, backend_impl: str): """Test that make_data returns the correct shapes.""" - m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONVEX_OBJECTS) - d = mjx.make_data(m) + d = mjx.make_data(m, backend_impl=backend_impl) nq = 22 nbody = 5 @@ -230,7 +251,7 @@ class DataIOTest(parameterized.TestCase): nv = 19 nefc = 185 - self.assertEqual(d.nefc, nefc) + self.assertEqual(d._impl.nefc, nefc) self.assertEqual(d.qpos.shape, (nq,)) self.assertEqual(d.qvel.shape, (nv,)) self.assertEqual(d.act.shape, (0,)) @@ -251,57 +272,68 @@ class DataIOTest(parameterized.TestCase): self.assertEqual(d.geom_xpos.shape, (6, 3)) self.assertEqual(d.geom_xmat.shape, (6, 3, 3)) self.assertEqual(d.subtree_com.shape, (nbody, 3)) - self.assertEqual(d.cdof.shape, (nv, 6)) - self.assertEqual(d.cinert.shape, (nbody, 10)) - self.assertEqual(d.crb.shape, (nbody, 10)) - self.assertEqual(d.actuator_length.shape, (1,)) - self.assertEqual(d.actuator_moment.shape, (1, nv)) - self.assertEqual(d.qM.shape, (nv, nv)) - self.assertEqual(d.qLD.shape, (nv, nv)) - self.assertEqual(d.qLDiagInv.shape, (0,)) - self.assertEqual(d.contact.dist.shape, (ncon,)) - self.assertEqual(d.contact.pos.shape, (ncon, 3)) - self.assertEqual(d.contact.frame.shape, (ncon, 3, 3)) - self.assertEqual(d.contact.solref.shape, (ncon, 2)) - self.assertEqual(d.contact.solimp.shape, (ncon, 5)) - self.assertEqual(d.contact.geom1.shape, (ncon,)) - self.assertEqual(d.contact.geom2.shape, (ncon,)) - self.assertEqual(d.efc_J.shape, (nefc, nv)) - self.assertEqual(d.efc_frictionloss.shape, (nefc,)) - self.assertEqual(d.efc_D.shape, (nefc,)) - self.assertEqual(d.actuator_velocity.shape, (1,)) + self.assertEqual(d._impl.cdof.shape, (nv, 6)) + self.assertEqual(d._impl.cinert.shape, (nbody, 10)) + self.assertEqual(d._impl.crb.shape, (nbody, 10)) + self.assertEqual(d._impl.actuator_length.shape, (1,)) + self.assertEqual(d._impl.actuator_moment.shape, (1, nv)) + 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)) + self.assertEqual(d._impl.contact.solref.shape, (ncon, 2)) + self.assertEqual(d._impl.contact.solimp.shape, (ncon, 5)) + self.assertEqual(d._impl.contact.geom1.shape, (ncon,)) + self.assertEqual(d._impl.contact.geom2.shape, (ncon,)) + self.assertEqual(d._impl.efc_J.shape, (nefc, nv)) + self.assertEqual(d._impl.efc_frictionloss.shape, (nefc,)) + self.assertEqual(d._impl.efc_D.shape, (nefc,)) + self.assertEqual(d._impl.actuator_velocity.shape, (1,)) self.assertEqual(d.cvel.shape, (nbody, 6)) - self.assertEqual(d.cdof_dot.shape, (nv, 6)) + self.assertEqual(d._impl.cdof_dot.shape, (nv, 6)) self.assertEqual(d.qfrc_bias.shape, (nv,)) self.assertEqual(d.qfrc_passive.shape, (nv,)) - self.assertEqual(d.efc_aref.shape, (nefc,)) + self.assertEqual(d._impl.efc_aref.shape, (nefc,)) self.assertEqual(d.qfrc_actuator.shape, (nv,)) self.assertEqual(d.qfrc_smooth.shape, (nv,)) self.assertEqual(d.qacc_smooth.shape, (nv,)) self.assertEqual(d.qfrc_constraint.shape, (nv,)) self.assertEqual(d.qfrc_inverse.shape, (nv,)) - self.assertEqual(d.efc_force.shape, (nefc,)) + self.assertEqual(d._impl.efc_force.shape, (nefc,)) + + if backend_impl == 'jax': + self.assertEqual(d._impl.qM.shape, (nv, nv)) + self.assertEqual(d._impl.qLD.shape, (nv, nv)) + self.assertEqual(d._impl.qLDiagInv.shape, (0,)) + elif backend_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 - d = mjx.make_data(m) - self.assertEqual(d.qM.shape, (nm,)) - self.assertEqual(d.qLD.shape, (nm,)) - self.assertEqual(d.qLDiagInv.shape, (nv,)) + d = mjx.make_data(m, backend_impl=backend_impl) + self.assertEqual(d._impl.qM.shape, (nm,)) + self.assertEqual(d._impl.qLD.shape, (nm,)) + self.assertEqual(d._impl.qLDiagInv.shape, (nv,)) - def test_put_data(self): + if backend_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,)) + + @parameterized.parameters('jax', 'c') + def test_put_data(self, backend_impl: str): """Test that put_data puts the correct data for dense and sparse.""" - m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) d = mujoco.MjData(m) mujoco.mj_step(m, d, 2) - dx = mjx.put_data(m, d) + dx = mjx.put_data(m, d, backend_impl=backend_impl) # check a few fields np.testing.assert_allclose(dx.qpos, d.qpos) np.testing.assert_allclose(dx.xpos, d.xpos) np.testing.assert_allclose(dx.cvel, d.cvel) - np.testing.assert_allclose(dx.cdof_dot, d.cdof_dot) + np.testing.assert_allclose(dx._impl.cdof_dot, d.cdof_dot) # check that there are no weak types self.assertFalse( @@ -312,21 +344,26 @@ class DataIOTest(parameterized.TestCase): ) ) - # check that qM is transformed properly - 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)) + if backend_impl == 'jax': + # check that qM is transformed properly + 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 backend_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) # 4 contacts, 2 for each capsule against the plane - self.assertEqual(dx.contact.dist.shape, (4,)) + self.assertEqual(dx._impl.contact.dist.shape, (4,)) self.assertEqual(d.ncon, 1) # however only 1 contact in this step - np.testing.assert_allclose(dx.contact.dist[0], d.contact.dist[0]) - self.assertTrue((dx.contact.dist[1:] > 0).all()) - self.assertEqual(dx.contact.frame.shape, (4, 3, 3)) + np.testing.assert_allclose(dx._impl.contact.dist[0], d.contact.dist[0]) + self.assertTrue((dx._impl.contact.dist[1:] > 0).all()) + self.assertEqual(dx._impl.contact.frame.shape, (4, 3, 3)) np.testing.assert_allclose( - dx.contact.frame[0].reshape(9), d.contact.frame[0] + dx._impl.contact.frame[0].reshape(9), d.contact.frame[0] ) - np.testing.assert_allclose(dx.contact.frame[1:], 0) + np.testing.assert_allclose(dx._impl.contact.frame[1:], 0) # xmat, ximat, geom_xmat are all shape transformed self.assertEqual(dx.xmat.shape, (3, 3, 3)) @@ -339,57 +376,65 @@ class DataIOTest(parameterized.TestCase): np.testing.assert_allclose(dx.site_xmat.reshape((1, 9)), d.site_xmat) # tendon data is correct - np.testing.assert_allclose(dx.ten_length, d.ten_length) - np.testing.assert_equal(dx.ten_wrapadr, np.zeros((1,))) - np.testing.assert_equal(dx.ten_wrapnum, np.zeros((1,))) - np.testing.assert_equal(dx.wrap_obj, np.zeros((2, 2))) - np.testing.assert_equal(dx.wrap_xpos, np.zeros((2, 6))) + np.testing.assert_allclose(dx._impl.ten_length, d.ten_length) + np.testing.assert_equal(dx._impl.ten_wrapadr, np.zeros((1,))) + np.testing.assert_equal(dx._impl.ten_wrapnum, np.zeros((1,))) + np.testing.assert_equal(dx._impl.wrap_obj, np.zeros((2, 2))) + np.testing.assert_equal(dx._impl.wrap_xpos, np.zeros((2, 6))) # efc_ are also shape transformed and padded - self.assertEqual(dx.efc_J.shape, (45, 8)) # nefc, nv + self.assertEqual(dx._impl.efc_J.shape, (45, 8)) # nefc, nv d_efc_j = d.efc_J.reshape((-1, 8)) - np.testing.assert_allclose(dx.efc_J[:3], d_efc_j[:3]) # connect eq - np.testing.assert_allclose(dx.efc_J[3], d_efc_j[3]) # one active limit - np.testing.assert_allclose(dx.efc_J[4], 0) # one inactive limit - np.testing.assert_allclose(dx.efc_J[5:15], d_efc_j[4:14]) # contact - np.testing.assert_allclose(dx.efc_J[15:], 0) # no contact + np.testing.assert_allclose(dx._impl.efc_J[:3], d_efc_j[:3]) # connect eq + np.testing.assert_allclose( + dx._impl.efc_J[3], d_efc_j[3] + ) # one active limit + np.testing.assert_allclose(dx._impl.efc_J[4], 0) # one inactive limit + np.testing.assert_allclose(dx._impl.efc_J[5:15], d_efc_j[4:14]) # contact + np.testing.assert_allclose(dx._impl.efc_J[15:], 0) # no contact # check another efc_ too - self.assertEqual(dx.efc_aref.shape, (45,)) # nefc - np.testing.assert_allclose(dx.efc_aref[:3], d.efc_aref[:3]) - np.testing.assert_allclose(dx.efc_aref[3], d.efc_aref[3]) - np.testing.assert_allclose(dx.efc_aref[4], 0) - np.testing.assert_allclose(dx.efc_aref[5:15], d.efc_aref[4:14]) - np.testing.assert_allclose(dx.efc_aref[15:], 0) + self.assertEqual(dx._impl.efc_aref.shape, (45,)) # nefc + np.testing.assert_allclose(dx._impl.efc_aref[:3], d.efc_aref[:3]) + np.testing.assert_allclose(dx._impl.efc_aref[3], d.efc_aref[3]) + np.testing.assert_allclose(dx._impl.efc_aref[4], 0) + np.testing.assert_allclose(dx._impl.efc_aref[5:15], d.efc_aref[4:14]) + np.testing.assert_allclose(dx._impl.efc_aref[15:], 0) # check sparse transform is correct m.opt.jacobian = mujoco.mjtJacobian.mjJAC_SPARSE d = mujoco.MjData(m) mujoco.mj_step(m, d, 2) - dx_sparse = mjx.put_data(m, d) - np.testing.assert_allclose(dx_sparse.efc_J, dx.efc_J, atol=1e-8) + dx_sparse = mjx.put_data(m, d, backend_impl=backend_impl) + np.testing.assert_allclose(dx_sparse._impl.efc_J, dx._impl.efc_J, atol=1e-8) # check sparse mass matrices are correct - np.testing.assert_allclose(dx_sparse.qM, d.qM, atol=1e-8) - np.testing.assert_allclose(dx_sparse.qLD, d.qLD, atol=1e-8) - np.testing.assert_allclose(dx_sparse.qLDiagInv, d.qLDiagInv, atol=1e-8) + np.testing.assert_allclose(dx_sparse._impl.qM, d.qM, atol=1e-8) + np.testing.assert_allclose(dx_sparse._impl.qLD, d.qLD, atol=1e-8) + np.testing.assert_allclose( + dx_sparse._impl.qLDiagInv, d.qLDiagInv, atol=1e-8 + ) # check dense mass matrices are correct m.opt.jacobian = mujoco.mjtJacobian.mjJAC_DENSE d = mujoco.MjData(m) mujoco.mj_step(m, d, 2) - dx_from_dense = mjx.put_data(m, d) - qm = np.zeros((m.nv, m.nv)) - mujoco.mj_fullM(m, qm, d.qM) - np.testing.assert_allclose(dx_from_dense.qM, qm, atol=1e-8) + dx_from_dense = mjx.put_data(m, d, backend_impl=backend_impl) + if backend_impl == 'jax': + 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 backend_impl == 'c': + np.testing.assert_allclose(dx_from_dense._impl.qM, d.qM, atol=1e-8) - def test_get_data(self): + @parameterized.parameters('jax', 'c') + def test_get_data(self, backend_impl: str): """Test that get_data makes correct MjData.""" m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) d = mujoco.MjData(m) mujoco.mj_step(m, d, 2) - dx = mjx.put_data(m, d) + dx = mjx.put_data(m, d, backend_impl=backend_impl) d_2: mujoco.MjData = mjx.get_data(m, dx) # check a few fields @@ -424,6 +469,10 @@ 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 backend_impl == 'c': + # check fields specific to the C implementation + np.testing.assert_allclose(d_2.M_rownnz, d.M_rownnz) + def test_get_data_runs(self): xml = """ @@ -444,13 +493,14 @@ class DataIOTest(parameterized.TestCase): dx = mjx.put_data(m, d) mjx.get_data(m, dx) - def test_get_data_batched(self): + @parameterized.parameters('jax', 'c') + def test_get_data_batched(self, backend_impl): """Test that get_data makes correct List[MjData] for batched Data.""" m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) d = mujoco.MjData(m) mujoco.mj_step(m, d, 2) - dx = mjx.put_data(m, d) + dx = mjx.put_data(m, d, backend_impl=backend_impl) # second data in batch has contact dist > 0, disables contact dx_b = jax.tree_util.tree_map(lambda x: jp.stack((x, x + 0.05)), dx) ds = mjx.get_data(m, dx_b) @@ -460,13 +510,14 @@ class DataIOTest(parameterized.TestCase): self.assertEqual(ds[0].ncon, 1) self.assertEqual(ds[1].ncon, 0) - def test_get_data_into(self): + @parameterized.parameters('jax', 'c') + def test_get_data_into(self, backend_impl): """Test that get_data_into correctly populates an MjData.""" m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) d = mujoco.MjData(m) mujoco.mj_step(m, d, 2) - dx = mjx.put_data(m, d) + dx = mjx.put_data(m, d, backend_impl=backend_impl) d_2 = mujoco.MjData(m) mjx.get_data_into(d_2, m, dx) @@ -482,32 +533,33 @@ class DataIOTest(parameterized.TestCase): self.assertEqual(d_2.contact.frame.shape, (1, 9)) np.testing.assert_allclose(d_2.contact.frame, d.contact.frame) - def test_get_data_into_wrong_shape(self): + @parameterized.parameters('jax', 'c') + def test_get_data_into_wrong_shape(self, backend_impl): """Tests that get_data_into throwsif input and output shapes don't match.""" m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) d = mujoco.MjData(m) mujoco.mj_step(m, d, 2) - dx = mjx.put_data(m, d) + dx = mjx.put_data(m, d, backend_impl=backend_impl) m_2 = mujoco.MjModel.from_xml_string(_MULTIPLE_CONVEX_OBJECTS) d_2 = mujoco.MjData(m_2) with self.assertRaisesRegex(ValueError, r'Input field.*has shape.*'): mjx.get_data_into(d_2, m, dx) - def test_make_matches_put(self): + @parameterized.parameters('jax', 'c') + def test_make_matches_put(self, backend_impl): """Test that make_data produces a pytree that matches put_data.""" - m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) d = mujoco.MjData(m) mujoco.mj_step(m, d, 2) - dx = mjx.put_data(m, d) + dx = mjx.put_data(m, d, backend_impl=backend_impl) step_fn = lambda d: d.replace(time=d.time + 1) step_fn_jit = jax.jit(step_fn).lower(dx).compile() # placing an MjData onto device should yield the same treedef mjx.Data as # calling make_data. they should be interchangeable for jax functions: - step_fn_jit(mjx.make_data(m)) + step_fn_jit(mjx.make_data(m, backend_impl=backend_impl)) def test_contact_elliptic_condim1(self): """Test that condim=1 with ConeType.ELLIPTIC is not implemented.""" @@ -556,7 +608,327 @@ class DataIOTest(parameterized.TestCase): """) with self.assertRaises(NotImplementedError): - mjx.put_model(m) + mjx.put_model(m, backend_impl='jax') + + +class FullCompatTest(parameterized.TestCase): + """Tests for the _full_compat flag.""" + + def test_full_compat_deprecated(self): + """Tests that _full_compat is deprecated.""" + xml = """ + + + + + + + + + + + + """ + m = mujoco.MjModel.from_xml_string(xml) + with self.assertWarns(DeprecationWarning): + out = mjx_io.put_model(m, _full_compat=True) + self.assertEqual(out.backend_impl, BackendImpl.C) + with self.assertWarns(DeprecationWarning): + out = mjx_io.make_data(m, _full_compat=True) + self.assertEqual(out.backend_impl, BackendImpl.C) + + +# Test cases for `_resolve_backend_impl_and_device` where the device is +# specified by the user and the device is available. +_DEVICE_TEST_CASES = [ + # Arguments use the following format: + # (device_type_str, backend_impl_str, + # (expected_device, expected_backend_impl))) + # No backend specified. + ('cpu', None, ('cpu', BackendImpl.C)), + ('gpu-notnvidia', None, ('gpu', BackendImpl.JAX)), + ('gpu-nvidia', None, ('gpu', BackendImpl.WARP)), + ('tpu', None, ('tpu', BackendImpl.JAX)), + # JAX backend specified. + ('cpu', 'jax', ('cpu', BackendImpl.JAX)), + ('gpu-notnvidia', 'jax', ('gpu', BackendImpl.JAX)), + ('gpu-nvidia', 'jax', ('gpu', BackendImpl.JAX)), + ('tpu', 'jax', ('tpu', BackendImpl.JAX)), + # WARP backend specified. + ('cpu', 'warp', ('cpu', 'error')), + ('gpu-notnvidia', 'warp', ('cpu', 'error')), + ('gpu-nvidia', 'warp', ('gpu', BackendImpl.WARP)), + ('tpu', 'warp', ('tpu', 'error')), + # C backend specified. + ('cpu', 'c', ('cpu', BackendImpl.C)), + ('gpu-notnvidia', 'c', ('cpu', 'error')), + ('gpu-nvidia', 'c', ('cpu', 'error')), + ('tpu', 'c', ('tpu', 'error')), +] + +# Test cases for `_resolve_backend_impl_and_device` where the user does NOT +# specify a device. We mock the JAX default device. +_DEFAULT_DEVICE_TEST_CASES = [ + # Arguments use the following format: + # (jax.default_device, backend_impl_str, + # (expected_device, expected_backend_impl)) + # No backend impl specified. + ('cpu', None, ('cpu', BackendImpl.C)), + ('gpu-notnvidia', None, ('gpu', BackendImpl.JAX)), + ('gpu-nvidia', None, ('gpu', BackendImpl.WARP)), + ('tpu', None, ('tpu', BackendImpl.JAX)), + # JAX backend impl specified. + ('cpu', 'jax', ('cpu', BackendImpl.JAX)), + ('gpu-notnvidia', 'jax', ('gpu', BackendImpl.JAX)), + ('gpu-nvidia', 'jax', ('gpu', BackendImpl.JAX)), + ('tpu', 'jax', ('tpu', BackendImpl.JAX)), + # WARP backend impl specified. + ('cpu', 'warp', ('cpu', 'error')), + ('gpu-notnvidia', 'warp', ('cpu', 'error')), + ('gpu-nvidia', 'warp', ('gpu', BackendImpl.WARP)), + ('tpu', 'warp', ('tpu', 'error')), + # C backend impl specified, CPU should always be available. + ('cpu', 'c', ('cpu', BackendImpl.C)), + ('gpu-notnvidia', 'c', ('cpu', BackendImpl.C)), + ('gpu-nvidia', 'c', ('cpu', BackendImpl.C)), + ('tpu', 'c', ('cpu', BackendImpl.C)), +] + + +class ResolveBackendImplAndDeviceTest(parameterized.TestCase): + """Tests for the _resolve_backend_impl_and_device function.""" + + def setUp(self): + super().setUp() + + # Create mock devices + self.mock_cpu = mock.Mock(spec=jax.Device) + self.mock_cpu.platform = 'cpu' + self.mock_cpu.device_kind = 'Mock CPU' + self.mock_cpu.id = 0 + + self.mock_nvidia_gpu = mock.Mock(spec=jax.Device) + self.mock_nvidia_gpu.platform = 'gpu' + self.mock_nvidia_gpu.device_kind = 'NVIDIA Mocked GPU' + self.mock_nvidia_gpu.id = 0 + + self.mock_other_gpu = mock.Mock(spec=jax.Device) + self.mock_other_gpu.platform = 'gpu' + self.mock_other_gpu.device_kind = 'Other Mocked GPU' + self.mock_other_gpu.id = 1 + + self.mock_tpu = mock.Mock(spec=jax.Device) + self.mock_tpu.platform = 'tpu' + self.mock_tpu.device_kind = 'Mock TPU' + self.mock_tpu.id = 0 + + # Patch jax.devices for the entire test class using enter_context + self.mock_jax_devices = self.enter_context(mock.patch('jax.devices')) + self.mock_default_backend = self.enter_context( + mock.patch('jax.default_backend') + ) + + @parameterized.named_parameters( + (f'{str(args[0])}_{str(args[1])}', *args) for args in _DEVICE_TEST_CASES + ) + @mock.patch.dict( + os.environ, {'MJX_WARP_ENABLED': 'true', 'MJX_C_DEFAULT_ENABLED': 'true'} + ) + def test_resolve_with_device( + self, + device_type_str, + backend_impl_str, + expected, + ): + """Tests various combinations of device and backend impls.""" + input_device = { + 'cpu': self.mock_cpu, + 'gpu-nvidia': self.mock_nvidia_gpu, + 'gpu-notnvidia': self.mock_other_gpu, + 'tpu': self.mock_tpu, + }[device_type_str] + + def devices_side_effect(backend=None): + # assume the user-specified device is always available + if backend == 'cpu': + return [self.mock_cpu] + elif backend == 'gpu': + if 'nvidia' in device_type_str: + return [self.mock_nvidia_gpu] + return [self.mock_other_gpu] + elif backend == 'tpu': + return [self.mock_tpu] + elif backend == 'cuda': + return [self.mock_nvidia_gpu] + + raise AssertionError('Should not be called.') + + self.mock_jax_devices.side_effect = devices_side_effect + + expected_device, expected_backend_impl = expected + if expected_backend_impl == 'error': + with self.assertRaises(AssertionError): + mjx_io._resolve_backend_impl_and_device( + backend_impl=backend_impl_str, device=input_device + ) + return + + actual_backend_impl, actual_device = ( + mjx_io._resolve_backend_impl_and_device( + backend_impl=backend_impl_str, device=input_device + ) + ) + + self.assertEqual(actual_backend_impl, expected_backend_impl) + self.assertIsNotNone(actual_device) + self.assertEqual(actual_device.platform, expected_device) + + @parameterized.named_parameters( + (f'{str(args[0])}_{str(args[1])}', *args) + for args in _DEFAULT_DEVICE_TEST_CASES + ) + @mock.patch.dict( + os.environ, {'MJX_WARP_ENABLED': 'true', 'MJX_C_DEFAULT_ENABLED': 'true'} + ) + def test_resolve_without_device( + self, + default_device_str, + backend_impl_str, + expected, + ): + """Tests various combinations of jax.default_device and backend impls.""" + default_devices = { + 'cpu': [self.mock_cpu], + 'gpu-nvidia': [self.mock_nvidia_gpu, self.mock_cpu], + 'gpu-notnvidia': [self.mock_other_gpu, self.mock_cpu], + 'tpu': [self.mock_tpu, self.mock_cpu], + }[default_device_str] + + def devices_side_effect(backend=None): + if backend == 'cpu': + return [self.mock_cpu] # CPU is always available + if backend == 'gpu' and default_device_str == 'gpu-notnvidia': + return [self.mock_other_gpu] + if backend == 'gpu' and default_device_str == 'gpu-nvidia': + return [self.mock_nvidia_gpu] + if backend == 'cuda' and default_device_str == 'gpu-nvidia': + return [self.mock_nvidia_gpu] + if backend == 'tpu' and default_device_str == 'tpu': + return [self.mock_tpu] + if backend is None: + return default_devices + if backend == 'cuda': + raise RuntimeError('cuda backend not supported') + raise AssertionError('jax.devices error') + + self.mock_jax_devices.side_effect = devices_side_effect + default_device_side_effect_str = { + 'cpu': 'cpu', + 'gpu-nvidia': 'gpu', + 'gpu-notnvidia': 'gpu', + 'tpu': 'tpu', + }[default_device_str] + self.mock_default_backend.side_effect = ( + lambda: default_device_side_effect_str + ) + + expected_device, expected_backend_impl = expected + if ( + expected_backend_impl == 'error' + and default_device_str != 'gpu-nvidia' + and backend_impl_str == 'warp' + ): + with self.assertRaisesRegex(RuntimeError, 'cuda backend not supported'): + mjx_io._resolve_backend_impl_and_device( + backend_impl=backend_impl_str, device=None + ) + return + + if expected_backend_impl == 'error': + with self.assertRaises(AssertionError): + mjx_io._resolve_backend_impl_and_device( + backend_impl=backend_impl_str, device=None + ) + return + + actual_backend_impl, actual_device = ( + mjx_io._resolve_backend_impl_and_device( + backend_impl=backend_impl_str, device=None + ) + ) + + self.assertEqual(actual_backend_impl, expected_backend_impl) + self.assertIsNotNone(actual_device) + self.assertEqual(actual_device.platform, expected_device) + + @mock.patch.dict(os.environ, {'MJX_WARP_ENABLED': 'false'}) + def test_resolve_warp_disabled(self): + """Tests behavior when MJX_WARP_ENABLED is false.""" + self.mock_jax_devices.side_effect = lambda backend=None: ( + [self.mock_nvidia_gpu, self.mock_cpu] + if backend is None + else ([self.mock_nvidia_gpu] if backend == 'gpu' else [self.mock_cpu]) + ) + self.mock_default_backend.side_effect = lambda: 'gpu' + + # Default to JAX instead of WARP on NVIDIA GPU. + backend_impl, device = mjx_io._resolve_backend_impl_and_device( + backend_impl=None, device=None + ) + self.assertEqual(backend_impl, BackendImpl.JAX) + self.assertEqual(device.platform, 'gpu') + + # Specifying an NVIDIA GPU should still choose JAX. + backend_impl, device = mjx_io._resolve_backend_impl_and_device( + backend_impl=None, device=self.mock_nvidia_gpu + ) + self.assertEqual(backend_impl, BackendImpl.JAX) + self.assertEqual(device.platform, 'gpu') + + # Requesting warp explicitly should fail since it is disabled. + with self.assertRaises(AssertionError): + mjx_io._resolve_backend_impl_and_device( + backend_impl='warp', device=self.mock_nvidia_gpu + ) + with self.assertRaises(AssertionError): + mjx_io._resolve_backend_impl_and_device(backend_impl='warp', device=None) + + @mock.patch.dict(os.environ, {'MJX_C_DEFAULT_ENABLED': 'false'}) + def test_resolve_c_disabled(self): + """Tests behavior when MJX_C_DEFAULT_ENABLED is false.""" + # Users expect that CPU defaults to the JAX impl. But in the future, it will + # default to the C backend implementation. This test checks that + # MJX_C_DEFAULT_ENABLED=false defaults to the old behavior, until the + # migration to MJEP-15 is complete. + self.mock_jax_devices.side_effect = lambda backend=None: ([self.mock_cpu]) + self.mock_default_backend.side_effect = lambda: 'cpu' + + # Default to JAX instead of C on CPU. + backend_impl, device = mjx_io._resolve_backend_impl_and_device( + backend_impl=None, device=None + ) + self.assertEqual(backend_impl, BackendImpl.JAX) + self.assertEqual(device.platform, 'cpu') + + # Specifing CPU should still choose JAX. + backend_impl, device = mjx_io._resolve_backend_impl_and_device( + backend_impl=None, device=self.mock_cpu + ) + self.assertEqual(backend_impl, BackendImpl.JAX) + self.assertEqual(device.platform, 'cpu') + + # Specifying C should choose C! + backend_impl, device = mjx_io._resolve_backend_impl_and_device( + backend_impl='c', device=None + ) + self.assertEqual(backend_impl, BackendImpl.C) + self.assertEqual(device.platform, 'cpu') + + backend_impl, device = mjx_io._resolve_backend_impl_and_device( + backend_impl='c', device=self.mock_cpu + ) + self.assertEqual(backend_impl, BackendImpl.C) + self.assertEqual(device.platform, 'cpu') if __name__ == '__main__': diff --git a/mjx/mujoco/mjx/_src/passive.py b/mjx/mujoco/mjx/_src/passive.py index c289690f..b4a8b674 100644 --- a/mjx/mujoco/mjx/_src/passive.py +++ b/mjx/mujoco/mjx/_src/passive.py @@ -23,9 +23,12 @@ from mujoco.mjx._src import scan from mujoco.mjx._src import support # pylint: disable=g-importing-member from mujoco.mjx._src.types import Data +from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import JointType from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import ModelJAX +from mujoco.mjx._src.types import OptionJAX # pylint: enable=g-importing-member @@ -72,11 +75,11 @@ def _spring_damper(m: Model, d: Data) -> jax.Array: qfrc -= m.dof_damping * d.qvel # tendon-level spring-dampers - below, above = m.tendon_lengthspring.T - d.ten_length + below, above = m.tendon_lengthspring.T - d._impl.ten_length frc_spring = jp.where(below > 0, m.tendon_stiffness * below, 0) frc_spring = jp.where(above < 0, m.tendon_stiffness * above, frc_spring) - frc_damper = -m.tendon_damping * d.ten_velocity - qfrc += d.ten_J.T @ (frc_spring + frc_damper) + frc_damper = -m.tendon_damping * d._impl.ten_velocity + qfrc += d._impl.ten_J.T @ (frc_spring + frc_damper) return qfrc @@ -113,6 +116,9 @@ def _fluid(m: Model, d: Data) -> jax.Array: def passive(m: Model, d: Data) -> Data: """Adds all passive forces.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('passive requires JAX backend implementation.') + if m.opt.disableflags & DisableBit.PASSIVE: return d.replace(qfrc_passive=jp.zeros(m.nv), qfrc_gravcomp=jp.zeros(m.nv)) @@ -124,7 +130,7 @@ def passive(m: Model, d: Data) -> Data: # add gravcomp unless added via actuators qfrc_passive += qfrc_gravcomp * (1 - m.jnt_actgravcomp[m.dof_jntid]) - if m.opt.has_fluid_params: + if m.opt.has_fluid_params: # pytype: disable=attribute-error qfrc_passive += _fluid(m, d) d = d.replace(qfrc_passive=qfrc_passive, qfrc_gravcomp=qfrc_gravcomp) diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py index 24471ccc..41751527 100644 --- a/mjx/mujoco/mjx/_src/sensor.py +++ b/mjx/mujoco/mjx/_src/sensor.py @@ -22,9 +22,12 @@ from mujoco.mjx._src import math from mujoco.mjx._src import ray from mujoco.mjx._src import smooth from mujoco.mjx._src import support +from mujoco.mjx._src.types import BackendImpl from mujoco.mjx._src.types import Data +from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import ModelJAX from mujoco.mjx._src.types import ObjType from mujoco.mjx._src.types import SensorType from mujoco.mjx._src.types import TrnType @@ -32,7 +35,7 @@ from mujoco.mjx._src.types import TrnType import numpy as np -def apply_cutoff( +def _apply_cutoff( sensor: jax.Array, cutoff: jax.Array, data_type: int ) -> jax.Array: """Clip sensor to cutoff value.""" @@ -51,6 +54,8 @@ def apply_cutoff( def sensor_pos(m: Model, d: Data) -> Data: """Compute position-dependent sensors values.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('sensor_pos requires JAX backend implementation.') if m.opt.disableflags & DisableBit.SENSOR: return d @@ -163,15 +168,15 @@ def sensor_pos(m: Model, d: Data) -> Data: sensor, _ = jax.vmap( ray.ray, in_axes=(None, None, 0, 0, None, None, None) )(m, d, site_xpos, site_mat, (), True, sid) - sensors.append(apply_cutoff(sensor, cutoffs, data_type[0])) + sensors.append(_apply_cutoff(sensor, cutoffs, data_type[0])) adrs.append(adr[idxs]) continue # avoid adding to sensors/adrs list a second time elif sensor_type == SensorType.JOINTPOS: sensor = d.qpos[m.jnt_qposadr[objid]] elif sensor_type == SensorType.TENDONPOS: - sensor = d.ten_length[objid] + sensor = d._impl.ten_length[objid] elif sensor_type == SensorType.ACTUATORPOS: - sensor = d.actuator_length[objid] + sensor = d._impl.actuator_length[objid] elif sensor_type == SensorType.BALLQUAT: jnt_qposadr = m.jnt_qposadr[objid, None] + np.arange(4)[None] quat = d.qpos[jnt_qposadr] @@ -194,7 +199,7 @@ def sensor_pos(m: Model, d: Data) -> Data: cutofft = cutoff[idxt] sensor = jax.vmap(_framepos)(xpos, xpos_ref, xmat_ref, refidt) adrt = adr[idxt, None] + np.arange(3)[None] - sensors.append(apply_cutoff(sensor, cutofft, data_type[0]).reshape(-1)) + sensors.append(_apply_cutoff(sensor, cutofft, data_type[0]).reshape(-1)) adrs.append(adrt.reshape(-1)) continue # avoid adding to sensors/adrs list a second time elif sensor_type in frame_axis: @@ -214,7 +219,7 @@ def sensor_pos(m: Model, d: Data) -> Data: cutofft = cutoff[idxt] sensor = jax.vmap(_frameaxis)(xmat, xmat_ref, refidt) adrt = adr[idxt, None] + np.arange(3)[None] - sensors.append(apply_cutoff(sensor, cutofft, data_type[0]).reshape(-1)) + sensors.append(_apply_cutoff(sensor, cutofft, data_type[0]).reshape(-1)) adrs.append(adrt.reshape(-1)) continue # avoid adding to sensors/adrs list a second time elif sensor_type == SensorType.FRAMEQUAT: @@ -255,7 +260,7 @@ def sensor_pos(m: Model, d: Data) -> Data: ) )(quat, refquat, refidt) adrt = adr[idxt, None] + np.arange(4)[None] - sensors.append(apply_cutoff(sensor, cutofft, data_type[0]).reshape(-1)) + sensors.append(_apply_cutoff(sensor, cutofft, data_type[0]).reshape(-1)) adrs.append(adrt.reshape(-1)) continue # avoid adding to sensors/adrs list a second time elif sensor_type == SensorType.SUBTREECOM: @@ -267,7 +272,7 @@ def sensor_pos(m: Model, d: Data) -> Data: # TODO(taylorhowell): raise error after adding sensor check to io.py continue # unsupported sensor type - sensors.append(apply_cutoff(sensor, cutoff, data_type[0]).reshape(-1)) + sensors.append(_apply_cutoff(sensor, cutoff, data_type[0]).reshape(-1)) adrs.append(adr) if not adrs: @@ -282,6 +287,8 @@ def sensor_pos(m: Model, d: Data) -> Data: def sensor_vel(m: Model, d: Data) -> Data: """Compute velocity-dependent sensors values.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('sensor_vel requires JAX backend implementation.') if m.opt.disableflags & DisableBit.SENSOR: return d @@ -333,9 +340,9 @@ def sensor_vel(m: Model, d: Data) -> Data: elif sensor_type == SensorType.JOINTVEL: sensor = d.qvel[m.jnt_dofadr[objid]] elif sensor_type == SensorType.TENDONVEL: - sensor = d.ten_velocity[objid] + sensor = d._impl.ten_velocity[objid] elif sensor_type == SensorType.ACTUATORVEL: - sensor = d.actuator_velocity[objid] + sensor = d._impl.actuator_velocity[objid] elif sensor_type == SensorType.BALLANGVEL: jnt_dotadr = m.jnt_dofadr[objid, None] + np.arange(3)[None] sensor = d.qvel[jnt_dotadr] @@ -393,20 +400,20 @@ def sensor_vel(m: Model, d: Data) -> Data: adrt = adr[idxt, None] + np.arange(3)[None] - sensors.append(apply_cutoff(sensor, cutofft, data_type[0]).reshape(-1)) + sensors.append(_apply_cutoff(sensor, cutofft, data_type[0]).reshape(-1)) adrs.append(adrt.reshape(-1)) continue # avoid adding to sensors/adrs list a second time elif sensor_type == SensorType.SUBTREELINVEL: - sensor = d.subtree_linvel[objid] + sensor = d._impl.subtree_linvel[objid] adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) elif sensor_type == SensorType.SUBTREEANGMOM: - sensor = d.subtree_angmom[objid] + sensor = d._impl.subtree_angmom[objid] adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) else: # TODO(taylorhowell): raise error after adding sensor check to io.py continue # unsupported sensor type - sensors.append(apply_cutoff(sensor, cutoff, data_type[0]).reshape(-1)) + sensors.append(_apply_cutoff(sensor, cutoff, data_type[0]).reshape(-1)) adrs.append(adr) if not adrs: @@ -421,6 +428,8 @@ def sensor_vel(m: Model, d: Data) -> Data: def sensor_acc(m: Model, d: Data) -> Data: """Compute acceleration/force-dependent sensors values.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('sensor_acc requires JAX backend implementation.') if m.opt.disableflags & DisableBit.SENSOR: return d @@ -460,14 +469,14 @@ def sensor_acc(m: Model, d: Data) -> Data: # compute contact forces forces = [] condim_ids = [] - for dim in set(d.contact.dim): + for dim in set(d._impl.contact.dim): force, condim_id = support.contact_force_dim(m, d, dim) forces.append(force) condim_ids.append(condim_id) forces = jp.concatenate(forces)[np.argsort(np.concatenate(condim_ids))] # get bodies of contact geoms - conbody = jp.array(m.geom_bodyid)[d.contact.geom] + conbody = jp.array(m.geom_bodyid)[d._impl.contact.geom] # get site information site_bodyid = m.site_bodyid[objid] @@ -477,12 +486,14 @@ def sensor_acc(m: Model, d: Data) -> Data: site_type = m.site_type[objid] conbody0 = site_bodyid[:, None] == conbody[:, 0] conbody1 = site_bodyid[:, None] == conbody[:, 1] - contacts = (d.contact.efc_address >= 0)[None] & (conbody0 | conbody1) + contacts = (d._impl.contact.efc_address >= 0)[None] & ( + conbody0 | conbody1 + ) # compute conray, flip if second body conray = jax.vmap( lambda frame, force: math.normalize(frame[0] * force[0]) - )(d.contact.frame, forces) + )(d._impl.contact.frame, forces) conray = jp.where(conbody1[..., None], -conray, conray) # compute distance, mapping over sites and contacts @@ -504,7 +515,7 @@ def sensor_acc(m: Model, d: Data) -> Data: site_xpos[dist_id_site], site_xmat[dist_id_site], st, - d.contact.pos, + d._impl.contact.pos, conray[dist_id_site], ) dist.append(jp.where(jp.isinf(dist_site), 0, dist_site)) @@ -526,14 +537,14 @@ def sensor_acc(m: Model, d: Data) -> Data: bodyid = m.site_bodyid[objid] rot = d.site_xmat[objid] cvel = d.cvel[bodyid] - cacc = d.cacc[bodyid] + cacc = d._impl.cacc[bodyid] dif = d.site_xpos[objid] - d.subtree_com[m.body_rootid[bodyid]] sensor = _accelerometer(cvel, cacc, dif, rot) adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) elif sensor_type == SensorType.FORCE: bodyid = m.site_bodyid[objid] - cfrc_int = d.cfrc_int[bodyid] + cfrc_int = d._impl.cfrc_int[bodyid] site_xmat = d.site_xmat[objid] sensor = jax.vmap(lambda mat, vec: mat.T @ vec)( site_xmat, cfrc_int[:, 3:] @@ -542,7 +553,7 @@ def sensor_acc(m: Model, d: Data) -> Data: elif sensor_type == SensorType.TORQUE: bodyid = m.site_bodyid[objid] rootid = m.body_rootid[bodyid] - cfrc_int = d.cfrc_int[bodyid] + cfrc_int = d._impl.cfrc_int[bodyid] site_xmat = d.site_xmat[objid] dif = d.site_xpos[objid] - d.subtree_com[rootid] sensor = jax.vmap( @@ -571,7 +582,7 @@ def sensor_acc(m: Model, d: Data) -> Data: pos, bodyid = objtype_data[ot] pos = pos[objidt] bodyid = bodyid[objidt] - cacc = d.cacc[bodyid] + cacc = d._impl.cacc[bodyid] if sensor_type == SensorType.FRAMELINACC: @@ -601,7 +612,7 @@ def sensor_acc(m: Model, d: Data) -> Data: # TODO(taylorhowell): raise error after adding sensor check to io.py continue # unsupported sensor type - sensors.append(apply_cutoff(sensor, cutoff, data_type[0]).reshape(-1)) + sensors.append(_apply_cutoff(sensor, cutoff, data_type[0]).reshape(-1)) adrs.append(adr) if not adrs: diff --git a/mjx/mujoco/mjx/_src/sensor_test.py b/mjx/mujoco/mjx/_src/sensor_test.py index fce6225c..cab50779 100644 --- a/mjx/mujoco/mjx/_src/sensor_test.py +++ b/mjx/mujoco/mjx/_src/sensor_test.py @@ -64,14 +64,14 @@ class SensorTest(parameterized.TestCase): mujoco.mj_forward(m, d) mx = mjx.put_model(m) - dx = mjx.put_data(m, d).replace( - sensordata=jp.zeros_like(d.sensordata), - subtree_linvel=jp.zeros_like(d.subtree_linvel), - subtree_angmom=jp.zeros_like(d.subtree_angmom), - cacc=jp.zeros_like(d.cacc), - cfrc_int=jp.zeros_like(d.cfrc_int), - cfrc_ext=jp.zeros_like(d.cfrc_ext), - ) + dx = mjx.put_data(m, d).tree_replace({ + 'sensordata': jp.zeros_like(d.sensordata), + '_impl.subtree_linvel': jp.zeros_like(d.subtree_linvel), + '_impl.subtree_angmom': jp.zeros_like(d.subtree_angmom), + '_impl.cacc': jp.zeros_like(d.cacc), + '_impl.cfrc_int': jp.zeros_like(d.cfrc_int), + '_impl.cfrc_ext': jp.zeros_like(d.cfrc_ext), + }) dx = jax.jit(mjx.sensor_pos)(mx, dx) dx = jax.jit(mjx.sensor_vel)(mx, dx) dx = jax.jit(mjx.sensor_acc)(mx, dx) diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py index 92e46702..f87c2ed3 100644 --- a/mjx/mujoco/mjx/_src/smooth.py +++ b/mjx/mujoco/mjx/_src/smooth.py @@ -23,10 +23,12 @@ from mujoco.mjx._src import support # pylint: disable=g-importing-member from mujoco.mjx._src.types import CamLightType from mujoco.mjx._src.types import Data +from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import EqType from mujoco.mjx._src.types import JointType from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import ModelJAX from mujoco.mjx._src.types import TrnType from mujoco.mjx._src.types import WrapType # pylint: enable=g-importing-member @@ -35,7 +37,6 @@ import numpy as np def kinematics(m: Model, d: Data) -> Data: """Converts position/velocity from generalized coordinates to maximal.""" - def fn(carry, jnt_typs, jnt_pos, jnt_axis, qpos, qpos0, pos, quat): # calculate joint anchors, axes, body pos and quat in global frame # also normalize qpos while we're at it @@ -131,6 +132,8 @@ def kinematics(m: Model, d: Data) -> Data: def com_pos(m: Model, d: Data) -> Data: """Maps inertias and motion dofs to global frame centered at subtree-CoM.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('com_pos requires JAX backend implementation.') # calculate center of mass of each subtree def subtree_sum(carry, xipos, body_mass): @@ -162,7 +165,7 @@ def com_pos(m: Model, d: Data) -> Data: root_com = subtree_com[m.body_rootid] offset = d.xipos - root_com cinert = inert_com(m.body_inertia, d.ximat, offset, m.body_mass) - d = d.replace(cinert=cinert) + d = d.tree_replace({'_impl.cinert': cinert}) # map motion dofs to global frame centered at subtree_com def cdof_fn(jnt_typs, root_com, xmat, xanchor, xaxis): @@ -201,13 +204,16 @@ def com_pos(m: Model, d: Data) -> Data: d.xanchor, d.xaxis, ) - d = d.replace(cdof=cdof) + d = d.tree_replace({'_impl.cdof': cdof}) return d def camlight(m: Model, d: Data) -> Data: """Computes camera and light positions and orientations.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('camlight requires JAX backend implementation.') + if m.ncam == 0: return d.replace(cam_xpos=jp.zeros((0, 3)), cam_xmat=jp.zeros((0, 3, 3))) @@ -278,32 +284,33 @@ def camlight(m: Model, d: Data) -> Data: def crb(m: Model, d: Data) -> Data: """Runs composite rigid body inertia algorithm.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('crb requires JAX backend implementation.') def crb_fn(crb_child, crb_body): if crb_child is not None: crb_body += crb_child return crb_body - crb_body = scan.body_tree(m, crb_fn, 'b', 'b', d.cinert, reverse=True) + crb_body = scan.body_tree(m, crb_fn, 'b', 'b', d._impl.cinert, reverse=True) crb_body = crb_body.at[0].set(0.0) - d = d.replace(crb=crb_body) + d = d.tree_replace({'_impl.crb': crb_body}) crb_dof = jp.take(crb_body, jp.array(m.dof_bodyid), axis=0) - crb_cdof = jax.vmap(math.inert_mul)(crb_dof, d.cdof) - qm = support.make_m(m, crb_cdof, d.cdof, m.dof_armature) - d = d.replace(qM=qm) - if support.is_sparse(m) and d._qM_sparse.size > 0: # pylint: disable=protected-access - d = d.replace(_qM_sparse=qm) - + crb_cdof = jax.vmap(math.inert_mul)(crb_dof, d._impl.cdof) + qm = support.make_m(m, crb_cdof, d._impl.cdof, m.dof_armature) + d = d.tree_replace({'_impl.qM': qm}) return d def factor_m(m: Model, d: Data) -> Data: """Gets factorizaton of inertia-like matrix M, assumed spd.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('factor_m requires JAX backend implementation.') if not support.is_sparse(m): - qh, _ = jax.scipy.linalg.cho_factor(d.qM) - d = d.replace(qLD=qh) + qh, _ = jax.scipy.linalg.cho_factor(d._impl.qM) + d = d.tree_replace({'_impl.qLD': qh}) return d # build up indices for where we will do backwards updates over qLD @@ -325,7 +332,7 @@ def factor_m(m: Model, d: Data) -> Data: (out_beg, out_end, madr_d, madr_ij) ) - qld = d.qM + qld = d._impl.qM for _, updates in sorted(updates.items(), reverse=True): # combine the updates into one update batch (per depth level) @@ -353,20 +360,17 @@ def factor_m(m: Model, d: Data) -> Data: qld_diag = qld[m.dof_Madr] qld = (qld / qld[jp.array(madr_ds)]).at[m.dof_Madr].set(qld_diag) - d = d.replace(qLD=qld, qLDiagInv=1 / qld_diag) - if d._qLD_sparse.size > 0: # pylint: disable=protected-access - d = d.replace(_qLD_sparse=d.qLD) - if d._qLDiagInv_sparse.size > 0: # pylint: disable=protected-access - d = d.replace(_qLDiagInv_sparse=d.qLDiagInv) - + d = d.tree_replace({'_impl.qLD': qld, '_impl.qLDiagInv': 1 / qld_diag}) return d def solve_m(m: Model, d: Data, x: jax.Array) -> jax.Array: """Computes sparse backsubstitution: x = inv(L'*D*L)*y .""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('solve_m requires JAX backend implementation.') if not support.is_sparse(m): - return jax.scipy.linalg.cho_solve((d.qLD, False), x) + return jax.scipy.linalg.cho_solve((d._impl.qLD, False), x) depth = [] for i in range(m.nv): @@ -385,21 +389,23 @@ def solve_m(m: Model, d: Data, x: jax.Array) -> jax.Array: # x <- inv(L') * x for _, vals in sorted(updates_j.items(), reverse=True): j, madr_ij, i = np.array(vals).T - x = x.at[j].add(-d.qLD[madr_ij] * x[i]) + x = x.at[j].add(-d._impl.qLD[madr_ij] * x[i]) # x <- inv(D) * x - x = x * d.qLDiagInv + x = x * d._impl.qLDiagInv # x <- inv(L) * x for _, vals in sorted(updates_i.items()): i, madr_ij, j = np.array(vals).T - x = x.at[i].add(-d.qLD[madr_ij] * x[j]) + x = x.at[i].add(-d._impl.qLD[madr_ij] * x[j]) return x def com_vel(m: Model, d: Data) -> Data: """Computes cvel, cdof_dot.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('com_vel requires JAX backend implementation.') # forward scan down tree: accumulate link center of mass velocity def fn(parent, jnt_typs, cdof, qvel): @@ -431,17 +437,19 @@ def com_vel(m: Model, d: Data) -> Data: 'jvv', 'bv', m.jnt_type, - d.cdof, + d._impl.cdof, d.qvel, ) - d = d.replace(cvel=cvel, cdof_dot=cdof_dot) + d = d.tree_replace({'cvel': cvel, '_impl.cdof_dot': cdof_dot}) return d def subtree_vel(m: Model, d: Data) -> Data: """Subtree linear velocity and angular momentum.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('subtree_vel requires JAX backend implementation.') # bodywise quantities def _forward(cvel, xipos, ximat, subtree_com_root, mass, inertia): @@ -529,7 +537,10 @@ def subtree_vel(m: Model, d: Data) -> Data: reverse=True, ) - return d.replace(subtree_linvel=subtree_linvel, subtree_angmom=subtree_angmom) + return d.tree_replace({ + '_impl.subtree_linvel': subtree_linvel, + '_impl.subtree_angmom': subtree_angmom, + }) def rne(m: Model, d: Data, flg_acc: bool = False) -> Data: @@ -537,6 +548,8 @@ def rne(m: Model, d: Data, flg_acc: bool = False) -> Data: flg_acc=False removes inertial term. """ + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('rne requires JAX backend implementation.') # forward scan over tree: accumulate link center of mass acceleration def cacc_fn(cacc, cdof_dot, qvel, cdof, qacc): @@ -555,7 +568,7 @@ def rne(m: Model, d: Data, flg_acc: bool = False) -> Data: return cacc cacc = scan.body_tree( - m, cacc_fn, 'vvvv', 'b', d.cdof_dot, d.qvel, d.cdof, d.qacc + m, cacc_fn, 'vvvv', 'b', d._impl.cdof_dot, d.qvel, d._impl.cdof, d.qacc ) def frc(cinert, cacc, cvel): @@ -564,7 +577,7 @@ def rne(m: Model, d: Data, flg_acc: bool = False) -> Data: return frc - loc_cfrc = jax.vmap(frc)(d.cinert, cacc, d.cvel) + loc_cfrc = jax.vmap(frc)(d._impl.cinert, cacc, d.cvel) # backward scan up tree: accumulate body forces def cfrc_fn(cfrc_child, cfrc): @@ -573,7 +586,7 @@ def rne(m: Model, d: Data, flg_acc: bool = False) -> Data: return cfrc cfrc = scan.body_tree(m, cfrc_fn, 'b', 'b', loc_cfrc, reverse=True) - qfrc_bias = jax.vmap(jp.dot)(d.cdof, cfrc[jp.array(m.dof_bodyid)]) + qfrc_bias = jax.vmap(jp.dot)(d._impl.cdof, cfrc[jp.array(m.dof_bodyid)]) d = d.replace(qfrc_bias=qfrc_bias) @@ -582,6 +595,8 @@ def rne(m: Model, d: Data, flg_acc: bool = False) -> Data: def rne_postconstraint(m: Model, d: Data) -> Data: """RNE with complete data: compute cacc, cfrc_ext, cfrc_int.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('rne_postconstraint requires JAX backend implementation.') def _transform_force(frc, offset): force, torque = jp.split(frc, 2) @@ -602,7 +617,7 @@ def rne_postconstraint(m: Model, d: Data) -> Data: # compute contact forces for each condim forces = [] condim_idx = [] - for dim in set(d.contact.dim): + for dim in set(d._impl.contact.dim): force, idx = support.contact_force_dim(m, d, dim) forces.append(force) condim_idx.append(idx) @@ -629,10 +644,10 @@ def rne_postconstraint(m: Model, d: Data) -> Data: ) condim_idx = jp.concatenate(condim_idx) - frame = d.contact.frame[condim_idx] - pos = d.contact.pos[condim_idx] - id1 = jp.array(m.geom_bodyid)[d.contact.geom[condim_idx, 0]] - id2 = jp.array(m.geom_bodyid)[d.contact.geom[condim_idx, 1]] + frame = d._impl.contact.frame[condim_idx] + pos = d._impl.contact.pos[condim_idx] + id1 = jp.array(m.geom_bodyid)[d._impl.contact.geom[condim_idx, 0]] + id2 = jp.array(m.geom_bodyid)[d._impl.contact.geom[condim_idx, 1]] com1 = d.subtree_com[jp.array(m.body_rootid)][id1] com2 = d.subtree_com[jp.array(m.body_rootid)][id2] @@ -668,8 +683,8 @@ def rne_postconstraint(m: Model, d: Data) -> Data: ) # cacc = cacc_parent + cdofdot * qvel + cdof * qacc - cacc_vel = d.cdof_dot.T @ (mask * d.qvel) - cacc_acc = d.cdof.T @ (mask * d.qacc) + cacc_vel = d._impl.cdof_dot.T @ (mask * d.qvel) + cacc_acc = d._impl.cdof.T @ (mask * d.qacc) cacc = cacc_parent + cacc_vel + cacc_acc # cfrc_body = cinert * cacc + cvel x (cinert * cvel) @@ -687,7 +702,7 @@ def rne_postconstraint(m: Model, d: Data) -> Data: 'bbbbb', 'bb', cfrc_ext, - d.cinert, + d._impl.cinert, d.cvel, jp.array(m.body_dofadr), jp.array(m.body_dofnum), @@ -704,11 +719,18 @@ def rne_postconstraint(m: Model, d: Data) -> Data: ) # update data - return d.replace(cacc=cacc, cfrc_int=cfrc_int, cfrc_ext=cfrc_ext) + return d.tree_replace({ + '_impl.cacc': cacc, + '_impl.cfrc_int': cfrc_int, + '_impl.cfrc_ext': cfrc_ext, + }) def tendon(m: Model, d: Data) -> Data: """Computes tendon lengths and moments.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('tendon requires JAX backend implementation.') + if not m.ntendon: return d @@ -836,8 +858,8 @@ def tendon(m: Model, d: Data) -> Data: # wrap inside # TODO(taylorhowell): check that is_wrap_inside is consistent with # site and geom relative positions - (wrap_inside_id,) = np.nonzero(m.is_wrap_inside) - (wrap_outside_id,) = np.nonzero(~m.is_wrap_inside) + (wrap_inside_id,) = np.nonzero(m._impl.is_wrap_inside) + (wrap_outside_id,) = np.nonzero(~m._impl.is_wrap_inside) # compute geom wrap length and connect points (if wrap occurs) v_wrap = jax.vmap( @@ -853,9 +875,9 @@ def tendon(m: Model, d: Data) -> Data: has_sidesite[wrap_inside_id], is_sphere[wrap_inside_id], True, - m.wrap_inside_maxiter, - m.wrap_inside_tolerance, - m.wrap_inside_z_init, + m._impl.wrap_inside_maxiter, + m._impl.wrap_inside_tolerance, + m._impl.wrap_inside_z_init, ) lengths_outside, pnt0_outside, pnt1_outside = v_wrap( @@ -868,9 +890,9 @@ def tendon(m: Model, d: Data) -> Data: has_sidesite[wrap_outside_id], is_sphere[wrap_outside_id], False, - m.wrap_inside_maxiter, - m.wrap_inside_tolerance, - m.wrap_inside_z_init, + m._impl.wrap_inside_maxiter, + m._impl.wrap_inside_tolerance, + m._impl.wrap_inside_z_init, ) wrap_id = np.argsort(np.concatenate([wrap_inside_id, wrap_outside_id])) @@ -952,12 +974,14 @@ def tendon(m: Model, d: Data) -> Data: ) # assemble length and moment - ten_length = jp.zeros_like(d.ten_length).at[tendon_id_jnt].set(length_jnt) + ten_length = ( + jp.zeros_like(d._impl.ten_length).at[tendon_id_jnt].set(length_jnt) + ) ten_length = ten_length.at[tendon_id_site].add(length_site) ten_length = ten_length.at[tendon_id_geom].add(length_geom) ten_moment = ( - jp.zeros_like(d.ten_J) + jp.zeros_like(d._impl.ten_J) .at[adr_moment_jnt, dofadr_moment_jnt] .set(moment_jnt) ) @@ -1020,14 +1044,14 @@ def tendon(m: Model, d: Data) -> Data: [wrap_obj[sort], jp.zeros(2 * m.nwrap - count, dtype=int)] ).reshape((m.nwrap, 2)) - return d.replace( - ten_length=ten_length, - ten_J=ten_moment, - ten_wrapadr=jp.array(ten_wrapadr, dtype=int), - ten_wrapnum=jp.array(ten_wrapnum, dtype=int), - wrap_xpos=wrap_xpos, - wrap_obj=jp.array(wrap_obj, dtype=int), - ) + return d.tree_replace({ + '_impl.ten_length': ten_length, + '_impl.ten_J': ten_moment, + '_impl.ten_wrapadr': jp.array(ten_wrapadr, dtype=int), + '_impl.ten_wrapnum': jp.array(ten_wrapnum, dtype=int), + '_impl.wrap_xpos': wrap_xpos, + '_impl.wrap_obj': jp.array(wrap_obj, dtype=int), + }) def _site_dof_mask(m: Model) -> np.ndarray: @@ -1061,6 +1085,9 @@ def _site_dof_mask(m: Model) -> np.ndarray: def transmission(m: Model, d: Data) -> Data: """Computes actuator/transmission lengths and moments.""" + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('transmission requires JAX backend implementation.') + if not m.nu: return d @@ -1120,8 +1147,8 @@ def transmission(m: Model, d: Data) -> Data: wrench = jp.concatenate((frame_xmat @ gear[:3], frame_xmat @ gear[3:])) moment = jac @ wrench elif trntype == TrnType.TENDON: - length = d.ten_length[trnid[0]] * gear[:1] - moment = d.ten_J[trnid[0]] * gear[0] + length = d._impl.ten_length[trnid[0]] * gear[:1] + moment = d._impl.ten_J[trnid[0]] * gear[0] else: raise RuntimeError(f'unrecognized trntype: {TrnType(trntype)}') @@ -1153,5 +1180,7 @@ def transmission(m: Model, d: Data) -> Data: length = length.reshape((m.nu,)) moment = moment.reshape((m.nu, m.nv)) - d = d.replace(actuator_length=length, actuator_moment=moment) + d = d.tree_replace( + {'_impl.actuator_length': length, '_impl.actuator_moment': moment} + ) return d diff --git a/mjx/mujoco/mjx/_src/smooth_test.py b/mjx/mujoco/mjx/_src/smooth_test.py index eca05068..4f182a75 100644 --- a/mjx/mujoco/mjx/_src/smooth_test.py +++ b/mjx/mujoco/mjx/_src/smooth_test.py @@ -20,7 +20,7 @@ import jax import mujoco from mujoco import mjx from mujoco.mjx._src import test_util -from mujoco.mjx._src.types import ConeType +from mujoco.mjx._src.types import ConeType # pylint: disable=g-importing-member import numpy as np # tolerance for difference between MuJoCo and MJX smooth calculations - mostly @@ -78,30 +78,27 @@ class SmoothTest(absltest.TestCase): # com_pos dx = jax.jit(mjx.com_pos)(mx, mjx.put_data(m, d)) _assert_attr_eq(d, dx, 'subtree_com') - _assert_attr_eq(d, dx, 'cinert') - _assert_attr_eq(d, dx, 'cdof') + _assert_attr_eq(d, dx._impl, 'cinert') + _assert_attr_eq(d, dx._impl, 'cdof') # camlight dx = jax.jit(mjx.camlight)(mx, mjx.put_data(m, d)) _assert_attr_eq(d, dx, 'cam_xpos') _assert_eq(d.cam_xmat.reshape((-1, 3, 3)), dx.cam_xmat, 'cam_xmat') # crb dx = jax.jit(mjx.crb)(mx, mjx.put_data(m, d)) - _assert_attr_eq(d, dx, 'crb') - _assert_attr_eq(d, dx, 'qM') - _assert_eq(dx._qM_sparse, np.zeros(0), '_qM_sparse') + _assert_attr_eq(d, dx._impl, 'crb') + _assert_attr_eq(d, dx._impl, 'qM') # factor_m dx = jax.jit(mjx.factor_m)(mx, mjx.put_data(m, d)) qLDLegacy = np.zeros(mx.nM) # pylint:disable=invalid-name for i in range(m.nM): qLDLegacy[d.mapM2M[i]] = d.qLD[i] - _assert_eq(qLDLegacy, dx.qLD, 'qLD') - _assert_attr_eq(d, dx, 'qLDiagInv') - _assert_eq(dx._qLD_sparse, np.zeros(0), '_qLD_sparse') - _assert_eq(dx._qLDiagInv_sparse, np.zeros(0), '_qLDiagInv_sparse') + _assert_eq(qLDLegacy, dx._impl.qLD, 'qLD') + _assert_attr_eq(d, dx._impl, 'qLDiagInv') # com_vel dx = jax.jit(mjx.com_vel)(mx, mjx.put_data(m, d)) _assert_attr_eq(d, dx, 'cvel') - _assert_attr_eq(d, dx, 'cdof_dot') + _assert_attr_eq(d, dx._impl, 'cdof_dot') # rne dx = jax.jit(mjx.rne)(mx, mjx.put_data(m, d)) _assert_attr_eq(d, dx, 'qfrc_bias') @@ -122,11 +119,11 @@ class SmoothTest(absltest.TestCase): mujoco.mj_forward(m, d) # tendon dx = jax.jit(mjx.tendon)(mx, mjx.put_data(m, d)) - _assert_attr_eq(d, dx, 'ten_J') - _assert_attr_eq(d, dx, 'ten_length') + _assert_attr_eq(d, dx._impl, 'ten_J') + _assert_attr_eq(d, dx._impl, 'ten_length') # transmission dx = jax.jit(mjx.transmission)(mx, dx) - _assert_attr_eq(d, dx, 'actuator_length') + _assert_attr_eq(d, dx._impl, 'actuator_length') # convert sparse actuator_moment to dense representation moment = np.zeros((m.nu, m.nv)) @@ -137,7 +134,7 @@ class SmoothTest(absltest.TestCase): d.moment_rowadr, d.moment_colind, ) - _assert_eq(moment, dx.actuator_moment, 'actuator_moment') + _assert_eq(moment, dx._impl.actuator_moment, 'actuator_moment') def test_disable_gravity(self): m = mujoco.MjModel.from_xml_string(""" @@ -197,7 +194,7 @@ class SmoothTest(absltest.TestCase): mujoco.mj_transmission(m, d) dx = jax.jit(mjx.transmission)(mx, dx) - _assert_attr_eq(d, dx, 'actuator_length') + _assert_attr_eq(d, dx._impl, 'actuator_length') # convert sparse actuator_moment to dense representation moment = np.zeros((m.nu, m.nv)) @@ -208,7 +205,7 @@ class SmoothTest(absltest.TestCase): d.moment_rowadr, d.moment_colind, ) - _assert_eq(moment, dx.actuator_moment, 'actuator_moment') + _assert_eq(moment, dx._impl.actuator_moment, 'actuator_moment') def test_subtree_vel(self): """Tests MJX subtree_vel function matches MuJoCo mj_subtreeVel.""" @@ -226,8 +223,8 @@ class SmoothTest(absltest.TestCase): mujoco.mj_subtreeVel(m, d) dx = jax.jit(mjx.subtree_vel)(mx, dx) - _assert_attr_eq(d, dx, 'subtree_linvel') - _assert_attr_eq(d, dx, 'subtree_angmom') + _assert_attr_eq(d, dx._impl, 'subtree_linvel') + _assert_attr_eq(d, dx._impl, 'subtree_angmom') class RnePostConstraintTest(parameterized.TestCase): @@ -278,9 +275,9 @@ class RnePostConstraintTest(parameterized.TestCase): mujoco.mj_rnePostConstraint(m, d) dx = jax.jit(mjx.rne_postconstraint)(mx, dx) - _assert_eq(d.cacc, dx.cacc, 'cacc') - _assert_eq(d.cfrc_ext, dx.cfrc_ext, 'cfrc_ext') - _assert_eq(d.cfrc_int, dx.cfrc_int, 'cfrc_int') + _assert_eq(d.cacc, dx._impl.cacc, 'cacc') + _assert_eq(d.cfrc_ext, dx._impl.cfrc_ext, 'cfrc_ext') + _assert_eq(d.cfrc_int, dx._impl.cfrc_int, 'cfrc_int') class TendonTest(parameterized.TestCase): @@ -312,12 +309,12 @@ class TendonTest(parameterized.TestCase): mujoco.mj_forward(m, d) dx = jax.jit(mjx.forward)(mx, dx) - _assert_eq(d.ten_length, dx.ten_length, 'ten_length') - _assert_eq(d.ten_J, dx.ten_J, 'ten_J') - _assert_eq(d.ten_wrapnum, dx.ten_wrapnum, 'ten_wrapnum') - _assert_eq(d.ten_wrapadr, dx.ten_wrapadr, 'ten_wrapadr') - _assert_eq(d.wrap_obj, dx.wrap_obj, 'wrap_obj') - _assert_eq(d.wrap_xpos, dx.wrap_xpos, 'wrap_xpos') + _assert_eq(d.ten_length, dx._impl.ten_length, 'ten_length') + _assert_eq(d.ten_J, dx._impl.ten_J, 'ten_J') + _assert_eq(d.ten_wrapnum, dx._impl.ten_wrapnum, 'ten_wrapnum') + _assert_eq(d.ten_wrapadr, dx._impl.ten_wrapadr, 'ten_wrapadr') + _assert_eq(d.wrap_obj, dx._impl.wrap_obj, 'wrap_obj') + _assert_eq(d.wrap_xpos, dx._impl.wrap_xpos, 'wrap_xpos') if __name__ == '__main__': diff --git a/mjx/mujoco/mjx/_src/solver.py b/mjx/mujoco/mjx/_src/solver.py index efb58c41..98ec3d27 100644 --- a/mjx/mujoco/mjx/_src/solver.py +++ b/mjx/mujoco/mjx/_src/solver.py @@ -24,8 +24,10 @@ from mujoco.mjx._src import support from mujoco.mjx._src.dataclasses import PyTreeNode from mujoco.mjx._src.types import ConeType from mujoco.mjx._src.types import Data +from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import ModelJAX from mujoco.mjx._src.types import SolverType # pylint: enable=g-importing-member @@ -73,14 +75,19 @@ class Context(PyTreeNode): @classmethod def create(cls, m: Model, d: Data, grad: bool = True) -> 'Context': - jaref = d.efc_J @ d.qacc - d.efc_aref + if not isinstance(d._impl, DataJAX): + raise ValueError( + 'Constraint context requires JAX backend implementation.' + ) + + jaref = d._impl.efc_J @ d.qacc - d._impl.efc_aref # TODO(robotics-team): determine nv at which sparse mul is faster ma = support.mul_m(m, d, d.qacc) nv_0 = jp.zeros(m.nv) fri = 0.0 if m.opt.cone == ConeType.ELLIPTIC: - friction = d.contact.friction[d.contact.dim > 1] - dim = d.contact.dim[d.contact.dim > 1] + friction = d._impl.contact.friction[d._impl.contact.dim > 1] + dim = d._impl.contact.dim[d._impl.contact.dim > 1] mu = friction[:, 0] / jp.sqrt(m.opt.impratio) fri = jp.concatenate((mu[:, None], friction), axis=1) for condim in (3, 4, 6): @@ -90,7 +97,7 @@ class Context(PyTreeNode): qacc=d.qacc, qfrc_constraint=d.qfrc_constraint, Jaref=jaref, - efc_force=d.efc_force, + efc_force=d._impl.efc_force, Ma=ma, grad=nv_0, Mgrad=nv_0, @@ -145,18 +152,20 @@ class _LSPoint(PyTreeNode): ) -> '_LSPoint': """Creates a linesearch point with first and second derivatives.""" # roughly corresponds to CGEval in mujoco/src/engine/engine_solver.c + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('LSPoint requires JAX backend implementation.') cost, deriv_0, deriv_1 = 0.0, 0.0, 0.0 quad_total = quad_gauss x = ctx.Jaref + alpha * jv - active = (x < 0).at[: d.ne + d.nf].set(True) + active = (x < 0).at[: d._impl.ne + d._impl.nf].set(True) - dof_fl, ten_fl = m.dof_hasfrictionloss, m.tendon_hasfrictionloss + dof_fl, ten_fl = m._impl.dof_hasfrictionloss, m._impl.tendon_hasfrictionloss if (dof_fl.any() or ten_fl.any()) and not ( m.opt.disableflags & DisableBit.FRICTIONLOSS ): - f = d.efc_frictionloss - r = 1.0 / (d.efc_D + (d.efc_D == 0.0) * mujoco.mjMINVAL) + f = d._impl.efc_frictionloss + r = 1.0 / (d._impl.efc_D + (d._impl.efc_D == 0.0) * mujoco.mjMINVAL) rf, z = r * f, jp.zeros_like(f) linear_neg = (x <= -rf)[:, None] linear_pos = (x >= rf)[:, None] @@ -174,13 +183,13 @@ class _LSPoint(PyTreeNode): middle_zone = (tsqr > 0) & (n < (mu * t)) & ((mu * n + t) > 0) # quadratic cost for equality, friction, limits, frictionless contacts - dim1 = d.contact.efc_address[d.contact.dim == 1] - nefl = d.ne + d.nf + d.nl + dim1 = d._impl.contact.efc_address[d._impl.contact.dim == 1] + nefl = d._impl.ne + d._impl.nf + d._impl.nl active = active.at[nefl:].set(False).at[dim1].set(active[dim1]) quad_efld = jax.vmap(jp.multiply)(quad, active) quad_total += jp.sum(quad_efld, axis=0) # elliptic bottom zone: quadratic cost - efc_elliptic = d.contact.efc_address[d.contact.dim > 1] + efc_elliptic = d._impl.contact.efc_address[d._impl.contact.dim > 1] quad_c = jax.vmap(jp.multiply)(quad[efc_elliptic], bottom_zone) quad_total += jp.sum(quad_c, axis=0) # elliptic middle zone @@ -254,17 +263,20 @@ def _update_constraint(m: Model, d: Data, ctx: Context) -> Context: Returns: context with new constraint force and costs """ + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('_update_constraint requires JAX backend implementation.') + # ne constraints are always active, nf are conditionally active, others are # non-negative constraints. - active = (ctx.Jaref < 0).at[: d.ne + d.nf].set(True) + active = (ctx.Jaref < 0).at[: d._impl.ne + d._impl.nf].set(True) - floss_force, floss_cost = jp.zeros(d.nefc), 0.0 - dof_fl, ten_fl = m.dof_hasfrictionloss, m.tendon_hasfrictionloss + floss_force, floss_cost = jp.zeros(d._impl.nefc), 0.0 + dof_fl, ten_fl = m._impl.dof_hasfrictionloss, m._impl.tendon_hasfrictionloss if (dof_fl.any() or ten_fl.any()) and not ( m.opt.disableflags & DisableBit.FRICTIONLOSS ): - f = d.efc_frictionloss - r = 1.0 / (d.efc_D + (d.efc_D == 0.0) * mujoco.mjMINVAL) + f = d._impl.efc_frictionloss + r = 1.0 / (d._impl.efc_D + (d._impl.efc_D == 0.0) * mujoco.mjMINVAL) linear_neg = (ctx.Jaref <= -r * f) * (f > 0) linear_pos = (ctx.Jaref >= r * f) * (f > 0) active = active & ~linear_neg & ~linear_pos @@ -274,13 +286,13 @@ def _update_constraint(m: Model, d: Data, ctx: Context) -> Context: floss_cost = floss_cost.sum() if m.opt.cone == ConeType.PYRAMIDAL: - efc_force = d.efc_D * -ctx.Jaref * active + floss_force - cost = 0.5 * jp.sum(d.efc_D * ctx.Jaref * ctx.Jaref * active) + efc_force = d._impl.efc_D * -ctx.Jaref * active + floss_force + cost = 0.5 * jp.sum(d._impl.efc_D * ctx.Jaref * ctx.Jaref * active) dm, u, h = 0.0, 0.0, 0.0 elif m.opt.cone == ConeType.ELLIPTIC: - 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] + friction = d._impl.contact.friction[d._impl.contact.dim > 1] + efc_address = d._impl.contact.efc_address[d._impl.contact.dim > 1] + dim = d._impl.contact.dim[d._impl.contact.dim > 1] # to prevent out of range append zeros to ctx.Jaref slice_fn = jax.vmap( lambda x: jax.lax.dynamic_slice( @@ -297,12 +309,12 @@ def _update_constraint(m: Model, d: Data, ctx: Context) -> Context: adr_i.extend(range(addr, addr + condim)) adr_j.extend([i] * condim) active = active.at[jp.array(adr_i)].set(bottom_zone[jp.array(adr_j)]) - efc_force = d.efc_D * -ctx.Jaref * active + floss_force - cost = 0.5 * jp.sum(d.efc_D * ctx.Jaref * ctx.Jaref * active) + efc_force = d._impl.efc_D * -ctx.Jaref * active + floss_force + cost = 0.5 * jp.sum(d._impl.efc_D * ctx.Jaref * ctx.Jaref * active) # middle zone: cone middle_zone = (t > 0) & (n < (mu * t)) & ((mu * n + t) > 0) - dm = d.efc_D[efc_address] / jp.maximum( + dm = d._impl.efc_D[efc_address] / jp.maximum( mu * mu * (1 + mu * mu), mujoco.mjMINVAL ) nmt = n - mu * t @@ -339,7 +351,7 @@ def _update_constraint(m: Model, d: Data, ctx: Context) -> Context: else: raise NotImplementedError(f'unsupported cone type: {m.opt.cone}') - qfrc_constraint = d.efc_J.T @ efc_force + qfrc_constraint = d._impl.efc_J.T @ efc_force gauss = 0.5 * jp.dot(ctx.Ma - d.qfrc_smooth, ctx.qacc - d.qacc_smooth) ctx = ctx.replace( qfrc_constraint=qfrc_constraint, @@ -371,6 +383,8 @@ def _update_gradient(m: Model, d: Data, ctx: Context) -> Context: Raises: NotImplementedError: for unsupported solver type """ + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('_update_gradient requires JAX backend implementation.') grad = ctx.Ma - d.qfrc_smooth - ctx.qfrc_constraint @@ -378,16 +392,16 @@ def _update_gradient(m: Model, d: Data, ctx: Context) -> Context: mgrad = smooth.solve_m(m, d, grad) elif m.opt.solver == SolverType.NEWTON: if m.opt.cone == ConeType.ELLIPTIC: - cm = jp.diag(d.efc_D * ctx.active) - efc_address = d.contact.efc_address[d.contact.dim > 1] - dim = d.contact.dim[d.contact.dim > 1] + cm = jp.diag(d._impl.efc_D * ctx.active) + efc_address = d._impl.contact.efc_address[d._impl.contact.dim > 1] + dim = d._impl.contact.dim[d._impl.contact.dim > 1] # set efc of cone H along diagonal for i, (condim, addr) in enumerate(zip(dim, efc_address)): h_cone = ctx.h[i, :condim, :condim] cm = cm.at[addr : addr + condim, addr : addr + condim].add(h_cone) - h = d.efc_J.T @ cm @ d.efc_J + h = d._impl.efc_J.T @ cm @ d._impl.efc_J else: - h = (d.efc_J.T * d.efc_D * ctx.active) @ d.efc_J + h = (d._impl.efc_J.T * d._impl.efc_D * ctx.active) @ d._impl.efc_J h = support.full_m(m, d) + h h_ = jax.scipy.linalg.cho_factor(h) mgrad = jax.scipy.linalg.cho_solve(h_, grad) @@ -414,12 +428,15 @@ def _linesearch(m: Model, d: Data, ctx: Context) -> Context: Returns: updated context with new qacc, Ma, Jaref """ + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + raise ValueError('_lineasearch requires JAX backend implementation.') + smag = math.norm(ctx.search) * m.stat.meaninertia * max(1, m.nv) gtol = m.opt.tolerance * m.opt.ls_tolerance * smag # compute Mv, Jv mv = support.mul_m(m, d, ctx.search) - jv = d.efc_J @ ctx.search + jv = d._impl.efc_J @ ctx.search # prepare quadratics quad_gauss = jp.stack(( @@ -428,13 +445,15 @@ def _linesearch(m: Model, d: Data, ctx: Context) -> Context: 0.5 * jp.dot(ctx.search, mv), )) quad = jp.stack((0.5 * ctx.Jaref * ctx.Jaref, jv * ctx.Jaref, 0.5 * jv * jv)) - quad = (quad * d.efc_D).T + quad = (quad * d._impl.efc_D).T uu, v0, uv, vv = 0.0, 0.0, 0.0, 0.0 if m.opt.cone == ConeType.ELLIPTIC: - mask = d.contact.dim > 1 + mask = d._impl.contact.dim > 1 # complete vector quadratic (for bottom zone) efc_con, efc_fri = [], [] - for condim, addr in zip(d.contact.dim[mask], d.contact.efc_address[mask]): + for condim, addr in zip( + d._impl.contact.dim[mask], d._impl.contact.efc_address[mask] + ): efc_con.extend([addr] * (condim - 1)) efc_fri.extend(range(addr + 1, addr + condim)) quad = quad.at[jp.array(efc_con)].add(quad[jp.array(efc_fri)]) @@ -446,7 +465,7 @@ def _linesearch(m: Model, d: Data, ctx: Context) -> Context: jp.concatenate((jv, jp.zeros(3))), (x,), (6,) ) ) - efc_elliptic = d.contact.efc_address[mask] + efc_elliptic = d._impl.contact.efc_address[mask] v = jv_fn(efc_elliptic) * ctx.fri uu = jp.sum(ctx.u[:, 1:] * ctx.u[:, 1:], axis=1) v0 = v[:, 0] @@ -571,11 +590,11 @@ def solve(m: Model, d: Data) -> Data: else: ctx = jax.lax.while_loop(cond, body, ctx) - d = d.replace( - qacc_warmstart=ctx.qacc, - qacc=ctx.qacc, - qfrc_constraint=ctx.qfrc_constraint, - efc_force=ctx.efc_force, - ) + d = d.tree_replace({ + 'qacc_warmstart': ctx.qacc, + 'qfrc_constraint': ctx.qfrc_constraint, + 'qacc': ctx.qacc, + '_impl.efc_force': ctx.efc_force, + }) return d diff --git a/mjx/mujoco/mjx/_src/solver_test.py b/mjx/mujoco/mjx/_src/solver_test.py index 6da06526..0b1b3b5a 100644 --- a/mjx/mujoco/mjx/_src/solver_test.py +++ b/mjx/mujoco/mjx/_src/solver_test.py @@ -87,8 +87,8 @@ class SolverTest(parameterized.TestCase): # MJX finds very similar solutions with the newton solver if solver_ == mujoco.mjtSolver.mjSOL_NEWTON: - nnz = dx.efc_J.any(axis=1) - _assert_eq(d.efc_force, dx.efc_force[nnz], 'efc_force') + nnz = dx._impl.efc_J.any(axis=1) + _assert_eq(d.efc_force, dx._impl.efc_force[nnz], 'efc_force') _assert_attr_eq(d, dx, 'qfrc_constraint') _assert_attr_eq(d, dx, 'qacc') @@ -108,9 +108,9 @@ class SolverTest(parameterized.TestCase): mujoco.mj_forward(m, d) mx = mjx.put_model(m) dx = jax.jit(mjx.solve)(mx, mjx.put_data(m, d)) - nnz = dx.efc_J.any(axis=1) + nnz = dx._impl.efc_J.any(axis=1) # even without warmstart, newton converges quickly - _assert_eq(d.efc_force, dx.efc_force[nnz], 'efc_force', tol=2e-4) + _assert_eq(d.efc_force, dx._impl.efc_force[nnz], 'efc_force', tol=2e-4) def test_sparse(self): """Test solver works with sparse mass matrices.""" @@ -130,8 +130,8 @@ class SolverTest(parameterized.TestCase): _assert_attr_eq(d, dx, 'qacc') _assert_attr_eq(d, dx, 'qfrc_constraint') - nnz = dx.efc_J.any(axis=1) - _assert_eq(d.efc_force, dx.efc_force[nnz], 'efc_force') + nnz = dx._impl.efc_J.any(axis=1) + _assert_eq(d.efc_force, dx._impl.efc_force[nnz], 'efc_force') def test_quad_frictionloss(self): """Test a case with quadratic frictionloss constraints.""" @@ -144,8 +144,8 @@ class SolverTest(parameterized.TestCase): _assert_attr_eq(d, dx, 'qacc') _assert_attr_eq(d, dx, 'qfrc_constraint') - nnz = dx.efc_J.any(axis=1) - _assert_eq(d.efc_force, dx.efc_force[nnz], 'efc_force') + nnz = dx._impl.efc_J.any(axis=1) + _assert_eq(d.efc_force, dx._impl.efc_force[nnz], 'efc_force') # TODO(taylorhowell): condim=1 with ConeType.ELLIPTIC @parameterized.product(condim=(3, 4, 6), cone=tuple(ConeType)) diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py index 90c517ef..d68c0407 100644 --- a/mjx/mujoco/mjx/_src/support.py +++ b/mjx/mujoco/mjx/_src/support.py @@ -13,6 +13,7 @@ # limitations under the License. # ============================================================================== """Engine support functions.""" + from collections.abc import Iterable, Sequence from typing import Optional, Tuple, Union @@ -92,7 +93,7 @@ def full_m(m: Model, d: Data) -> jax.Array: """Reconstitute dense mass matrix from qM.""" if not is_sparse(m): - return d.qM + return d._impl.qM # pytype: disable=attribute-error ij = [] for i in range(m.nv): @@ -103,7 +104,7 @@ def full_m(m: Model, d: Data) -> jax.Array: i, j = (jp.array(x) for x in zip(*ij)) - mat = jp.zeros((m.nv, m.nv)).at[(i, j)].set(d.qM) + mat = jp.zeros((m.nv, m.nv)).at[(i, j)].set(d._impl.qM) # pytype: disable=attribute-error # also set upper triangular mat = mat + jp.tril(mat, -1).T @@ -115,9 +116,9 @@ def mul_m(m: Model, d: Data, vec: jax.Array) -> jax.Array: """Multiply vector by inertia matrix.""" if not is_sparse(m): - return d.qM @ vec + return d._impl.qM @ vec # pytype: disable=attribute-error - diag_mul = d.qM[jp.array(m.dof_Madr)] * vec + diag_mul = d._impl.qM[jp.array(m.dof_Madr)] * vec # pytype: disable=attribute-error is_, js, madr_ijs = [], [], [] for i in range(m.nv): @@ -131,8 +132,8 @@ def mul_m(m: Model, d: Data, vec: jax.Array) -> jax.Array: i, j, madr_ij = (jp.array(x, dtype=jp.int32) for x in (is_, js, madr_ijs)) - out = diag_mul.at[i].add(d.qM[madr_ij] * vec[j]) - out = out.at[j].add(d.qM[madr_ij] * vec[i]) + out = diag_mul.at[i].add(d._impl.qM[madr_ij] * vec[j]) # pytype: disable=attribute-error + out = out.at[j].add(d._impl.qM[madr_ij] * vec[i]) # pytype: disable=attribute-error return out @@ -147,9 +148,9 @@ def jac( mask = mask[jp.array(m.dof_bodyid)] > 0 offset = point - d.subtree_com[jp.array(m.body_rootid)[body_id]] - jacp = jax.vmap(lambda a, b=offset: a[3:] + jp.cross(a[:3], b))(d.cdof) + jacp = jax.vmap(lambda a, b=offset: a[3:] + jp.cross(a[:3], b))(d._impl.cdof) # pytype: disable=attribute-error jacp = jax.vmap(jp.multiply)(jacp, mask) - jacr = jax.vmap(jp.multiply)(d.cdof[:, :3], mask) + jacr = jax.vmap(jp.multiply)(d._impl.cdof[:, :3], mask) # pytype: disable=attribute-error return jacp, jacr @@ -548,20 +549,20 @@ def contact_force( m: Model, d: Data, contact_id: int, to_world_frame: bool = False ) -> jax.Array: """Extract 6D force:torque for one contact, in contact frame by default.""" - efc_address = d.contact.efc_address[contact_id] - condim = d.contact.dim[contact_id] + efc_address = d._impl.contact.efc_address[contact_id] # pytype: disable=attribute-error + condim = d._impl.contact.dim[contact_id] # pytype: disable=attribute-error if m.opt.cone == ConeType.PYRAMIDAL: force = _decode_pyramid( - d.efc_force[efc_address:], d.contact.friction[contact_id], condim + d._impl.efc_force[efc_address:], d._impl.contact.friction[contact_id], condim # pytype: disable=attribute-error ) elif m.opt.cone == ConeType.ELLIPTIC: - force = d.efc_force[efc_address : efc_address + condim] + force = d._impl.efc_force[efc_address : efc_address + condim] # pytype: disable=attribute-error force = jp.concatenate([force, jp.zeros((6 - condim))]) else: raise ValueError(f'Unknown cone type: {m.opt.cone}') if to_world_frame: - force = force.reshape((-1, 3)) @ d.contact.frame[contact_id] + force = force.reshape((-1, 3)) @ d._impl.contact.frame[contact_id] # pytype: disable=attribute-error force = force.reshape(-1) return force * (efc_address >= 0) @@ -572,21 +573,21 @@ def contact_force_dim( ) -> Tuple[jax.Array, np.ndarray]: """Extract 6D force:torque for contacts with dimension dim.""" # valid contact and condim indices - idx_dim = (d.contact.efc_address >= 0) & (d.contact.dim == dim) + idx_dim = (d._impl.contact.efc_address >= 0) & (d._impl.contact.dim == dim) # pytype: disable=attribute-error # contact force from efc if m.opt.cone == ConeType.PYRAMIDAL: efc_address = ( - d.contact.efc_address[idx_dim, None] + d._impl.contact.efc_address[idx_dim, None] # pytype: disable=attribute-error + np.arange(np.where(dim == 1, 1, 2 * (dim - 1)))[None] ) - efc_force = d.efc_force[efc_address] + efc_force = d._impl.efc_force[efc_address] # pytype: disable=attribute-error force = jax.vmap(_decode_pyramid, in_axes=(0, 0, None))( - efc_force, d.contact.friction[idx_dim], dim + efc_force, d._impl.contact.friction[idx_dim], dim # pytype: disable=attribute-error ) elif m.opt.cone == ConeType.ELLIPTIC: - efc_address = d.contact.efc_address[idx_dim, None] + np.arange(dim)[None] - force = d.efc_force[efc_address] + efc_address = d._impl.contact.efc_address[idx_dim, None] + np.arange(dim)[None] # pytype: disable=attribute-error + force = d._impl.efc_force[efc_address] # pytype: disable=attribute-error force = jp.hstack([force, jp.zeros((force.shape[0], 6 - dim))]) else: raise ValueError(f'Unknown cone type: {m.opt.cone}.') diff --git a/mjx/mujoco/mjx/_src/support_test.py b/mjx/mujoco/mjx/_src/support_test.py index 1a0c904c..5f060bc2 100644 --- a/mjx/mujoco/mjx/_src/support_test.py +++ b/mjx/mujoco/mjx/_src/support_test.py @@ -385,7 +385,7 @@ class SupportTest(parameterized.TestCase): # map MJX contacts to MJ ones def _find(g): - val = (g == dx.contact.geom).sum(axis=1) + val = (g == dx._impl.contact.geom).sum(axis=1) return np.where(val == 2)[0][0] contact_id_map = {i: _find(d.contact.geom[i]) for i in range(d.ncon)} @@ -399,7 +399,7 @@ class SupportTest(parameterized.TestCase): np.testing.assert_allclose(result, force, rtol=1e-5, atol=2) # check for zeros after first condim elements - condim = dx.contact.dim[j] + condim = dx._impl.contact.dim[j] if condim < 6: np.testing.assert_allclose(force[condim:], 0, rtol=1e-5, atol=1e-5) @@ -412,8 +412,8 @@ class SupportTest(parameterized.TestCase): ), )(mx, dx, j, True) # back to contact frame - force = force.at[:3].set(dx.contact.frame[j] @ force[:3]) - force = force.at[3:].set(dx.contact.frame[j] @ force[3:]) + force = force.at[:3].set(dx._impl.contact.frame[j] @ force[:3]) + force = force.at[3:].set(dx._impl.contact.frame[j] @ force[3:]) np.testing.assert_allclose(result, force, rtol=1e-5, atol=2) def test_wrap_inside(self): diff --git a/mjx/mujoco/mjx/_src/test_util.py b/mjx/mujoco/mjx/_src/test_util.py index f91ee9af..142a2eaa 100644 --- a/mjx/mujoco/mjx/_src/test_util.py +++ b/mjx/mujoco/mjx/_src/test_util.py @@ -106,20 +106,20 @@ def benchmark( def efc_order(m: mujoco.MjModel, d: mujoco.MjData, dx: Data) -> np.ndarray: - """Returns a sort order such that dx.efc_*[order][:d.nefc] == d.efc_*.""" + """Returns a sort order such that dx.efc_*[order][:d._impl.nefc] == d.efc_*.""" # pytype: disable=attribute-error # reorder efc rows to skip inactive constraints and match contact order - efl = dx.ne + dx.nf + dx.nl + efl = dx._impl.ne + dx._impl.nf + dx._impl.nl # pytype: disable=attribute-error order = np.arange(efl) - order[(dx.efc_J[:efl] == 0).all(axis=1)] = 2**16 # move empty rows to end - for i in range(dx.ncon): - num_rows = dx.contact.dim[i] - if dx.contact.dim[i] > 1 and m.opt.cone == mujoco.mjtCone.mjCONE_PYRAMIDAL: - num_rows = (dx.contact.dim[i] - 1) * 2 - if dx.contact.dist[i] > 0: # move empty contacts to end + order[(dx._impl.efc_J[:efl] == 0).all(axis=1)] = 2**16 # move empty rows to end # pytype: disable=attribute-error + for i in range(dx._impl.ncon): # pytype: disable=attribute-error + num_rows = dx._impl.contact.dim[i] # pytype: disable=attribute-error + if dx._impl.contact.dim[i] > 1 and m.opt.cone == mujoco.mjtCone.mjCONE_PYRAMIDAL: # pytype: disable=attribute-error + num_rows = (dx._impl.contact.dim[i] - 1) * 2 # pytype: disable=attribute-error + if dx._impl.contact.dist[i] > 0: # move empty contacts to end # pytype: disable=attribute-error order = np.append(order, np.repeat(2**16, num_rows)) continue - contact_match = (d.contact.geom == dx.contact.geom[i]).all(axis=-1) - contact_match &= (d.contact.pos == dx.contact.pos[i]).all(axis=-1) + contact_match = (d.contact.geom == dx._impl.contact.geom[i]).all(axis=-1) # pytype: disable=attribute-error + contact_match &= (d.contact.pos == dx._impl.contact.pos[i]).all(axis=-1) # pytype: disable=attribute-error assert contact_match.any(), f'contact {i} not found' contact_id = np.nonzero(contact_match)[0][0] order = np.append(order, np.repeat(efl + contact_id, num_rows)) diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 657cd524..c70d71a2 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -14,20 +14,33 @@ # ============================================================================== """Base types used in MJX.""" -import dataclasses import enum -from typing import Tuple +from typing import Tuple, Union +import warnings + import jax import mujoco from mujoco.mjx._src.dataclasses import PyTreeNode # pylint: disable=g-importing-member import numpy as np -def _restricted_to(platform: str): - """Specifies whether a field exists in only MuJoCo or MJX.""" - if platform not in ('mujoco', 'mjx'): - raise ValueError(f'unknown platform: {platform}') - return dataclasses.field(metadata={'restricted_to': platform}) +class BackendImpl(enum.Enum): + """Backend implementation to use.""" + + C = 'c' + JAX = 'jax' + WARP = 'warp' + + @classmethod + def _missing_(cls, value): + # This method is called only when lookup by value fails + # (e.g., BackendImpl('JAX') fails initially because 'JAX' != 'jax') + if not isinstance(value, str): + return None + for member in cls: + if member.value == value.lower(): + return member + return None class DisableBit(enum.IntFlag): @@ -430,77 +443,6 @@ class ObjType(PyTreeNode): CAMERA = mujoco.mjtObj.mjOBJ_CAMERA -class Option(PyTreeNode): - """Physics options. - - Attributes: - timestep: timestep - apirate: update rate for remote API (Hz) (not used) - impratio: ratio of friction-to-normal contact impedance - tolerance: main solver tolerance - ls_tolerance: CG/Newton linesearch tolerance - noslip_tolerance: noslip solver tolerance (not used) - ccd_tolerance: CCD solver tolerance (not used) - gravity: gravitational acceleration (3,) - wind: wind (for lift, drag and viscosity) - magnetic: global magnetic flux (not used) - density: density of medium - viscosity: viscosity of medium - o_margin: contact solver override: margin (not used) - o_solref: contact solver override: solref (not used) - o_solimp: contact solver override: solimp (not used) - o_friction[5]: contact solver override: friction (not used) - has_fluid_params: automatically set by mjx if wind/density/viscosity are - nonzero. Not used by mj - integrator: integration mode - cone: type of friction cone - jacobian: matrix layout for mass matrices (dense or sparse) - (note that this is different from MuJoCo, where jacobian - specifies whether efc_J and its accompanying matrices - are dense or sparse. - solver: solver algorithm - iterations: number of main solver iterations - ls_iterations: maximum number of CG/Newton linesearch iterations - noslip_iterations: maximum number of noslip solver iterations (not used) - ccd_iterations: maximum number of CCD solver iterations (not used) - disableflags: bit flags for disabling standard features - enableflags: bit flags for enabling optional features (not used) - disableactuator: bit flags for disabling actuators by group id (not used) - sdf_initpoints: number of starting points for gradient descent (not used) - sdf_iterations: max number of iterations for gradient descent (not used) - """ # fmt: skip - timestep: jax.Array - apirate: jax.Array = _restricted_to('mujoco') - impratio: jax.Array - tolerance: jax.Array - ls_tolerance: jax.Array - noslip_tolerance: jax.Array = _restricted_to('mujoco') - ccd_tolerance: jax.Array = _restricted_to('mujoco') - gravity: jax.Array - wind: jax.Array - magnetic: jax.Array - density: jax.Array - viscosity: jax.Array - o_margin: jax.Array - o_solref: jax.Array - o_solimp: jax.Array - o_friction: jax.Array - has_fluid_params: bool = _restricted_to('mjx') - integrator: IntegratorType - cone: ConeType - jacobian: JacobianType - solver: SolverType - iterations: int - ls_iterations: int - noslip_iterations: int = _restricted_to('mujoco') - ccd_iterations: int = _restricted_to('mujoco') - disableflags: DisableBit - enableflags: int - disableactuator: int - sdf_initpoints: int = _restricted_to('mujoco') - sdf_iterations: int = _restricted_to('mujoco') - - class Statistic(PyTreeNode): """Model statistics (in qpos0). @@ -519,371 +461,142 @@ class Statistic(PyTreeNode): center: jax.Array -class Model(PyTreeNode): - """Static model of the scene that remains unchanged with each physics step. +class Option(PyTreeNode): + """Physics options.""" # fmt: skip + timestep: jax.Array + impratio: jax.Array + tolerance: jax.Array + ls_tolerance: jax.Array + gravity: jax.Array + wind: jax.Array + magnetic: jax.Array + density: jax.Array + viscosity: jax.Array + o_margin: jax.Array + o_solref: jax.Array + o_solimp: jax.Array + o_friction: jax.Array + integrator: IntegratorType + cone: ConeType + jacobian: JacobianType + solver: SolverType + iterations: int + ls_iterations: int + disableflags: DisableBit + enableflags: int + disableactuator: int - Attributes: - nq: number of generalized coordinates = dim(qpos) - nv: number of degrees of freedom = dim(qvel) - nu: number of actuators/controls = dim(ctrl) - na: number of activation states = dim(act) - nbody: number of bodies - nbvh: number of total bounding volumes in all bodies - nbvhstatic: number of static bounding volumes (aabb stored in mjModel) - nbvhdynamic: number of dynamic bounding volumes (aabb stored in mjData) - njnt: number of joints - ngeom: number of geoms - nsite: number of sites - ncam: number of cameras - nlight: number of lights - nflex: number of flexes - nflexvert: number of vertices in all flexes - nflexedge: number of edges in all flexes - nflexelem: number of elements in all flexes - nflexelemdata: number of element vertex ids in all flexes - nflexshelldata: number of shell fragment vertex ids in all flexes - nflexevpair: number of element-vertex pairs in all flexes - nflextexcoord: number of vertices with texture coordinates - nmesh: number of meshes - nmeshvert: number of vertices in all meshes - nmeshnormal: number of normals in all meshes - nmeshtexcoord: number of texcoords in all meshes - nmeshface: number of triangular faces in all meshes - nmeshgraph: number of ints in mesh auxiliary data - nhfield: number of heightfields - nhfielddata: number of data points in all heightfields - ntex: number of textures - ntexdata: number of bytes in texture rgb data - nmat: number of materials - npair: number of predefined geom pairs - nexclude: number of excluded geom pairs - neq: number of equality constraints - ntendon: number of tendons - nwrap: number of wrap objects in all tendon paths - nsensor: number of sensors - nnumeric: number of numeric custom fields - ntuple: number of tuple custom fields - nkey: number of keyframes - nmocap: number of mocap bodies - nM: number of non-zeros in sparse inertia matrix - nD: number of non-zeros in sparse dof-dof matrix - nB: number of non-zeros in sparse body-dof matrix - nC: number of non-zeros in sparse reduced dof-dof matrix - nD: number of non-zeros in sparse dof-dof matrix - nJmom: number of non-zeros in sparse actuator_moment matrix - ntree: number of kinematic trees under world body - ngravcomp: number of bodies with nonzero gravcomp - nuserdata: size of userdata array - nsensordata: number of mjtNums in sensor data vector - narena: number of bytes in the mjData arena (inclusive of stack) - opt: physics options - stat: model statistics - qpos0: qpos values at default pose (nq,) - qpos_spring: reference pose for springs (nq,) - body_parentid: id of body's parent (nbody,) - body_rootid: id of root above body (nbody,) - body_weldid: id of body that this body is welded to (nbody,) - body_jntnum: number of joints for this body (nbody,) - body_jntadr: start addr of joints; -1: no joints (nbody,) - body_dofnum: number of motion degrees of freedom (nbody,) - body_dofadr: start addr of dofs; -1: no dofs (nbody,) - body_treeid: id of body's kinematic tree; -1: static (nbody,) - body_geomnum: number of geoms (nbody,) - body_geomadr: start addr of geoms; -1: no geoms (nbody,) - body_simple: 1: diag M; 2: diag M, sliders only (nbody,) - body_pos: position offset rel. to parent body (nbody, 3) - body_quat: orientation offset rel. to parent body (nbody, 4) - body_ipos: local position of center of mass (nbody, 3) - body_iquat: local orientation of inertia ellipsoid (nbody, 4) - body_mass: mass (nbody,) - body_subtreemass: mass of subtree starting at this body (nbody,) - body_inertia: diagonal inertia in ipos/iquat frame (nbody, 3) - body_gravcomp: antigravity force, units of body weight (nbody,) - body_margin: MAX over all geom margins (nbody,) - body_contype: OR over all geom contypes (nbody,) - body_conaffinity: OR over all geom conaffinities (nbody,) - body_bvhadr: address of bvh root (nbody,) - body_bvhnum: number of bounding volumes (nbody,) - bvh_child: left and right children in tree (nbvh, 2) - bvh_nodeid: geom or elem id of node; -1: non-leaf (nbvh,) - bvh_aabb: local bounding box (center, size) (nbvhstatic, 6) - body_invweight0: mean inv inert in qpos0 (trn, rot) (nbody, 2) - jnt_type: type of joint (mjtJoint) (njnt,) - jnt_qposadr: start addr in 'qpos' for joint's data (njnt,) - jnt_dofadr: start addr in 'qvel' for joint's data (njnt,) - jnt_bodyid: id of joint's body (njnt,) - jnt_group: group for visibility (njnt,) - jnt_limited: does joint have limits (njnt,) - jnt_actfrclimited: does joint have actuator force limits (njnt,) - jnt_actgravcomp: is gravcomp force applied via actuators (njnt,) - jnt_solref: constraint solver reference: limit (njnt, mjNREF) - jnt_solimp: constraint solver impedance: limit (njnt, mjNIMP) - jnt_pos: local anchor position (njnt, 3) - jnt_axis: local joint axis (njnt, 3) - jnt_stiffness: stiffness coefficient (njnt,) - jnt_range: joint limits (njnt, 2) - jnt_actfrcrange: range of total actuator force (njnt, 2) - jnt_margin: min distance for limit detection (njnt,) - dof_bodyid: id of dof's body (nv,) - dof_jntid: id of dof's joint (nv,) - dof_parentid: id of dof's parent; -1: none (nv,) - dof_treeid: id of dof's kinematic tree (nv,) - dof_Madr: dof address in M-diagonal (nv,) - dof_simplenum: number of consecutive simple dofs (nv,) - dof_solref: constraint solver reference:frictionloss (nv, mjNREF) - dof_solimp: constraint solver impedance:frictionloss (nv, mjNIMP) - dof_frictionloss: dof friction loss (nv,) - dof_hasfrictionloss: dof has >0 frictionloss (MJX) (nv,) - dof_armature: dof armature inertia/mass (nv,) - dof_damping: damping coefficient (nv,) - dof_invweight0: diag. inverse inertia in qpos0 (nv,) - dof_M0: diag. inertia in qpos0 (nv,) - geom_type: geometric type (mjtGeom) (ngeom,) - geom_contype: geom contact type (ngeom,) - geom_conaffinity: geom contact affinity (ngeom,) - geom_condim: contact dimensionality (1, 3, 4, 6) (ngeom,) - geom_bodyid: id of geom's body (ngeom,) - geom_dataid: id of geom's mesh/hfield; -1: none (ngeom,) - geom_group: group for visibility (ngeom,) - geom_matid: material id for rendering (ngeom,) - geom_priority: geom contact priority (ngeom,) - geom_solmix: mixing coef for solref/imp in geom pair (ngeom,) - geom_solref: constraint solver reference: contact (ngeom, mjNREF) - geom_solimp: constraint solver impedance: contact (ngeom, mjNIMP) - geom_size: geom-specific size parameters (ngeom, 3) - geom_aabb: bounding box, (center, size) (ngeom, 6) - geom_rbound: radius of bounding sphere (ngeom,) - geom_rbound_hfield: static rbound for hfield grid bounds (ngeom,) - geom_pos: local position offset rel. to body (ngeom, 3) - geom_quat: local orientation offset rel. to body (ngeom, 4) - geom_friction: friction for (slide, spin, roll) (ngeom, 3) - geom_margin: include in solver if dist0 frictionloss (MJX) (ntendon,) - wrap_type: wrap object type (mjtWrap) (nwrap,) - wrap_objid: object id: geom, site, joint (nwrap,) - wrap_prm: divisor, joint coef, or site id (nwrap,) - wrap_inside_maxiter: maximum iterations for wrap_inside - wrap_inside_tolerance: tolerance for wrap_inside - wrap_inside_z_init: initialization for wrap_inside - is_wrap_inside: spatial tendon sidesite inside geom (nwrapinside,) - actuator_trntype: transmission type (mjtTrn) (nu,) - actuator_dyntype: dynamics type (mjtDyn) (nu,) - actuator_gaintype: gain type (mjtGain) (nu,) - actuator_biastype: bias type (mjtBias) (nu,) - actuator_trnid: transmission id: joint, tendon, site (nu, 2) - actuator_actadr: first activation address; -1: stateless (nu,) - actuator_actnum: number of activation variables (nu,) - actuator_group: group for visibility (nu,) - actuator_ctrllimited: is control limited (nu,) - actuator_forcelimited: is force limited (nu,) - actuator_actlimited: is activation limited (nu,) - actuator_dynprm: dynamics parameters (nu, mjNDYN) - actuator_gainprm: gain parameters (nu, mjNGAIN) - actuator_biasprm: bias parameters (nu, mjNBIAS) - actuator_actearly: step activation before force (nu,) - actuator_ctrlrange: range of controls (nu, 2) - actuator_forcerange: range of forces (nu, 2) - actuator_actrange: range of activations (nu, 2) - actuator_gear: scale length and transmitted force (nu, 6) - actuator_cranklength: crank length for slider-crank (nu,) - actuator_acc0: acceleration from unit force in qpos0 (nu,) - actuator_lengthrange: feasible actuator length range (nu, 2) - sensor_type: sensor type (mjtSensor) (nsensor,) - sensor_datatype: numeric data type (mjtDataType) (nsensor,) - sensor_needstage: required compute stage (mjtStage) (nsensor,) - sensor_objtype: type of sensorized object (mjtObj) (nsensor,) - sensor_objid: id of sensorized object (nsensor,) - sensor_reftype: type of reference frame (mjtObj) (nsensor,) - sensor_refid: id of reference frame; -1: global frame (nsensor,) - sensor_dim: number of scalar outputs (nsensor,) - sensor_adr: address in sensor array (nsensor,) - sensor_cutoff: cutoff for real and positive; 0: ignore (nsensor,) - numeric_adr: address of field in numeric_data (nnumeric,) - numeric_data: array of all numeric fields (nnumericdata,) - tuple_adr: address of text in text_data (ntuple,) - tuple_size: number of objects in tuple (ntuple,) - tuple_objtype: array of object types in all tuples (ntupledata,) - tuple_objid: array of object ids in all tuples (ntupledata,) - tuple_objprm: array of object params in all tuples (ntupledata,) - key_time: key time (nkey,) - key_qpos: key position (nkey, nq) - key_qvel: key velocity (nkey, nv) - key_act: key activation (nkey, na) - key_mpos: key mocap position (nkey, nmocap, 3) - key_mquat: key mocap quaternion (nkey, nmocap, 4) - key_ctrl: key control (nkey, nu) - name_bodyadr: body name pointers (nbody,) - name_jntadr: joint name pointers (njnt,) - name_geomadr: geom name pointers (ngeom,) - name_siteadr: site name pointers (nsite,) - name_camadr: camera name pointers (ncam,) - name_meshadr: mesh name pointers (nmesh,) - name_pairadr: geom pair name pointers (npair,) - name_eqadr: equality constraint name pointers (neq,) - name_tendonadr: tendon name pointers (ntendon,) - name_actuatoradr: actuator name pointers (nu,) - name_sensoradr: sensor name pointers (nsensor,) - name_numericadr: numeric name pointers (nnumeric,) - name_tupleadr: tuple name pointers (ntuple,) - name_keyadr: keyframe name pointers (nkey,) - names: names of all objects, 0-terminated (nnames,) - signature: compilation signature - """ + +class OptionC(Option): + """C-specific option.""" + + apirate: jax.Array + noslip_tolerance: jax.Array + ccd_tolerance: jax.Array + noslip_iterations: int + ccd_iterations: int + sdf_initpoints: int + sdf_iterations: int + + +class OptionJAX(Option): + """JAX-specific option.""" + + has_fluid_params: bool + + +class ModelC(PyTreeNode): + """CPU-specific model data.""" + + nbvh: jax.Array + nbvhstatic: jax.Array + nbvhdynamic: 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 + ntree: 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 + 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_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 + actuator_plugin: jax.Array + + +class ModelJAX(PyTreeNode): + """JAX-specific model data.""" + + dof_hasfrictionloss: np.ndarray + geom_rbound_hfield: np.ndarray + mesh_convex: Tuple[ConvexMesh, ...] + tendon_hasfrictionloss: np.ndarray + wrap_inside_maxiter: int + wrap_inside_tolerance: float + wrap_inside_z_init: float + is_wrap_inside: np.ndarray + + +class Model(PyTreeNode): + """Static model of the scene that remains unchanged with each physics step.""" nq: int nv: int nu: int na: int nbody: int - nbvh: int = _restricted_to('mujoco') - nbvhstatic: int = _restricted_to('mujoco') - nbvhdynamic: int = _restricted_to('mujoco') njnt: int ngeom: int nsite: int ncam: int nlight: int - nflex: int = _restricted_to('mujoco') - nflexvert: int = _restricted_to('mujoco') - nflexedge: int = _restricted_to('mujoco') - nflexelem: int = _restricted_to('mujoco') - nflexelemdata: int = _restricted_to('mujoco') - nflexshelldata: int = _restricted_to('mujoco') - nflexevpair: int = _restricted_to('mujoco') - nflextexcoord: int = _restricted_to('mujoco') nmesh: int nmeshvert: int nmeshnormal: int @@ -910,11 +623,9 @@ class Model(PyTreeNode): nC: int # pylint:disable=invalid-name nD: int # pylint:disable=invalid-name nJmom: int # pylint:disable=invalid-name - ntree: int = _restricted_to('mujoco') ngravcomp: int nuserdata: int nsensordata: int - narena: int = _restricted_to('mujoco') opt: Option stat: Statistic qpos0: jax.Array @@ -943,11 +654,6 @@ class Model(PyTreeNode): body_margin: np.ndarray body_contype: np.ndarray body_conaffinity: np.ndarray - body_bvhadr: np.ndarray = _restricted_to('mujoco') - body_bvhnum: np.ndarray = _restricted_to('mujoco') - bvh_child: np.ndarray = _restricted_to('mujoco') - bvh_nodeid: np.ndarray = _restricted_to('mujoco') - bvh_aabb: np.ndarray = _restricted_to('mujoco') body_invweight0: jax.Array jnt_type: np.ndarray jnt_qposadr: np.ndarray @@ -973,7 +679,6 @@ class Model(PyTreeNode): dof_solref: jax.Array dof_solimp: jax.Array dof_frictionloss: jax.Array - dof_hasfrictionloss: np.ndarray = _restricted_to('mjx') dof_armature: jax.Array dof_damping: jax.Array dof_invweight0: jax.Array @@ -994,7 +699,6 @@ class Model(PyTreeNode): geom_size: jax.Array geom_aabb: np.ndarray geom_rbound: jax.Array - geom_rbound_hfield: np.ndarray = _restricted_to('mjx') geom_pos: jax.Array geom_quat: jax.Array geom_friction: jax.Array @@ -1021,8 +725,6 @@ class Model(PyTreeNode): cam_sensorsize: np.ndarray cam_intrinsic: np.ndarray light_mode: np.ndarray - light_bodyid: np.ndarray = _restricted_to('mujoco') - light_targetbodyid: np.ndarray = _restricted_to('mujoco') light_directional: jax.Array light_castshadow: jax.Array light_pos: jax.Array @@ -1031,46 +733,6 @@ class Model(PyTreeNode): light_pos0: np.ndarray light_dir0: np.ndarray light_cutoff: jax.Array - flex_contype: np.ndarray = _restricted_to('mujoco') - flex_conaffinity: np.ndarray = _restricted_to('mujoco') - flex_condim: np.ndarray = _restricted_to('mujoco') - flex_priority: np.ndarray = _restricted_to('mujoco') - flex_solmix: np.ndarray = _restricted_to('mujoco') - flex_solref: np.ndarray = _restricted_to('mujoco') - flex_solimp: np.ndarray = _restricted_to('mujoco') - flex_friction: np.ndarray = _restricted_to('mujoco') - flex_margin: np.ndarray = _restricted_to('mujoco') - flex_gap: np.ndarray = _restricted_to('mujoco') - flex_internal: np.ndarray = _restricted_to('mujoco') - flex_selfcollide: np.ndarray = _restricted_to('mujoco') - flex_activelayers: np.ndarray = _restricted_to('mujoco') - flex_dim: np.ndarray = _restricted_to('mujoco') - flex_vertadr: np.ndarray = _restricted_to('mujoco') - flex_vertnum: np.ndarray = _restricted_to('mujoco') - flex_edgeadr: np.ndarray = _restricted_to('mujoco') - flex_edgenum: np.ndarray = _restricted_to('mujoco') - flex_elemadr: np.ndarray = _restricted_to('mujoco') - flex_elemnum: np.ndarray = _restricted_to('mujoco') - flex_elemdataadr: np.ndarray = _restricted_to('mujoco') - flex_evpairadr: np.ndarray = _restricted_to('mujoco') - flex_evpairnum: np.ndarray = _restricted_to('mujoco') - flex_vertbodyid: np.ndarray = _restricted_to('mujoco') - flex_edge: np.ndarray = _restricted_to('mujoco') - flex_elem: np.ndarray = _restricted_to('mujoco') - flex_elemlayer: np.ndarray = _restricted_to('mujoco') - flex_evpair: np.ndarray = _restricted_to('mujoco') - flex_vert: np.ndarray = _restricted_to('mujoco') - flexedge_length0: np.ndarray = _restricted_to('mujoco') - flexedge_invweight0: np.ndarray = _restricted_to('mujoco') - flex_radius: np.ndarray = _restricted_to('mujoco') - flex_edgestiffness: np.ndarray = _restricted_to('mujoco') - flex_edgedamping: np.ndarray = _restricted_to('mujoco') - flex_edgeequality: np.ndarray = _restricted_to('mujoco') - flex_rigid: np.ndarray = _restricted_to('mujoco') - flexedge_rigid: np.ndarray = _restricted_to('mujoco') - flex_centered: np.ndarray = _restricted_to('mujoco') - flex_bvhadr: np.ndarray = _restricted_to('mujoco') - flex_bvhnum: np.ndarray = _restricted_to('mujoco') mesh_vertadr: np.ndarray mesh_vertnum: np.ndarray mesh_faceadr: np.ndarray @@ -1082,7 +744,6 @@ class Model(PyTreeNode): mesh_graph: np.ndarray mesh_pos: np.ndarray mesh_quat: np.ndarray - mesh_convex: Tuple[ConvexMesh, ...] = _restricted_to('mjx') mesh_texcoordadr: np.ndarray mesh_texcoordnum: np.ndarray mesh_texcoord: np.ndarray @@ -1136,14 +797,9 @@ class Model(PyTreeNode): tendon_lengthspring: jax.Array tendon_length0: jax.Array tendon_invweight0: jax.Array - tendon_hasfrictionloss: np.ndarray = _restricted_to('mjx') wrap_type: np.ndarray wrap_objid: np.ndarray wrap_prm: np.ndarray - wrap_inside_maxiter: int = _restricted_to('mjx') - wrap_inside_tolerance: float = _restricted_to('mjx') - wrap_inside_z_init: float = _restricted_to('mjx') - is_wrap_inside: np.ndarray = _restricted_to('mjx') actuator_trntype: np.ndarray actuator_dyntype: np.ndarray actuator_gaintype: np.ndarray @@ -1166,7 +822,6 @@ class Model(PyTreeNode): actuator_cranklength: np.ndarray actuator_acc0: jax.Array actuator_lengthrange: np.ndarray - actuator_plugin: np.ndarray = _restricted_to('mujoco') sensor_type: np.ndarray sensor_datatype: np.ndarray sensor_needstage: np.ndarray @@ -1209,6 +864,36 @@ class Model(PyTreeNode): names: bytes signature: np.uint64 _sizes: jax.Array + _impl: Union[ModelC, ModelJAX] + + @property + def backend_impl(self) -> BackendImpl: + return { + ModelC: BackendImpl.C, + ModelJAX: BackendImpl.JAX, + }[type(self._impl)] + + def __getattr__(self, name: str): + if name == 'value': + # Special case for NNX, the value attribute may not exist on the parent + # PyTreeNode, before it exists on the child PyTreeNode. Thanks NNX. + return object.__getattribute__(self, 'value') + + try: + impl_instsance = object.__getattribute__(self, '_impl') + val = getattr(impl_instsance, name) + warnings.warn( + f'Accessing `{name}` directly from `Model` is deprecated. ' + f'Access it via `model._impl.{name}` instead.', + DeprecationWarning, + stacklevel=2, + ) + except AttributeError: + # raise the standard exception + raise AttributeError( # pylint: disable=raise-missing-from + f"'{type(self).__name__}' object has no attribute '{name}'" + ) + return val class Contact(PyTreeNode): @@ -1246,144 +931,141 @@ class Contact(PyTreeNode): efc_address: np.ndarray -class Data(PyTreeNode): - r"""\Dynamic state that updates each step. +class DataC(PyTreeNode): + """C-specific data.""" - Attributes: - ne: number of equality constraints - nf: number of friction constraints - nl: number of limit constraints - nefc: number of constraints - ncon: number of contacts - solver_niter: number of solver iterations - time: simulation time - qpos: position (nq,) - qvel: velocity (nv,) - act: actuator activation (na,) - qacc_warmstart: acceleration used for warmstart (nv,) - ctrl: control (nu,) - qfrc_applied: applied generalized force (nv,) - xfrc_applied: applied Cartesian force/torque (nbody, 6) - eq_active: enable/disable constraints (neq,) - mocap_pos: positions of mocap bodies (nmocap x 3) - mocap_quat: orientations of mocap bodies (nmocap x 4) - qacc: acceleration (nv,) - act_dot: time-derivative of actuator activation (na,) - userdata: user data, not touched by engine (nuserdata,) - sensordata: sensor data array (nsensordata,) - xpos: Cartesian position of body frame (nbody, 3) - xquat: Cartesian orientation of body frame (nbody, 4) - xmat: Cartesian orientation of body frame (nbody, 3, 3) - xipos: Cartesian position of body com (nbody, 3) - ximat: Cartesian orientation of body inertia (nbody, 3, 3) - xanchor: Cartesian position of joint anchor (njnt, 3) - xaxis: Cartesian joint axis (njnt, 3) - geom_xpos: Cartesian geom position (ngeom, 3) - geom_xmat: Cartesian geom orientation (ngeom, 3, 3) - site_xpos: Cartesian site position (nsite, 3) - site_xmat: Cartesian site orientation (nsite, 3, 3) - cam_xpos: Cartesian camera position (ncam, 3) - cam_xmat: Cartesian camera orientation (ncam, 3, 3) - light_xpos: Cartesian light position (nlight, 3) - light_xdir: Cartesian light direction (nlight, 3) - subtree_com: center of mass of each subtree (nbody, 3) - cdof: com-based motion axis of each dof (nv, 6) - cinert: com-based body inertia and mass (nbody, 10) - flexvert_xpos: Cartesian flex vertex positions (nflexvert, 3) - flexelem_aabb: flex element bounding boxes (center, size) (nflexelem, 6) - flexedge_J_rownnz: number of non-zeros in Jacobian row (nflexedge,) - flexedge_J_rowadr: row start address in colind array (nflexedge,) - flexedge_J_colind: column indices in sparse Jacobian (nflexedge, nv) - flexedge_J: flex edge Jacobian (nflexedge, nv) - flexedge_length: flex edge lengths (nflexedge,) - ten_wrapadr: start address of tendon's path (ntendon,) - ten_wrapnum: number of wrap points in path (ntendon,) - ten_J_rownnz: number of non-zeros in Jacobian row (ntendon,) - ten_J_rowadr: row start address in colind array (ntendon,) - ten_J_colind: column indices in sparse Jacobian (ntendon, nv) - ten_J: tendon Jacobian (ntendon, nv) - ten_length: tendon lengths (ntendon,) - wrap_obj: geom id; -1: site; -2: pulley (nwrap*2,) - wrap_xpos: Cartesian 3D points in all path (nwrap*2, 3) - actuator_length: actuator lengths (nu,) - moment_rownnz: number of non-zeros in actuator_moment row (nu,) - moment_rowadr: row start address in colind array (nu,) - moment_colind: column indices in sparse Jacobian (nJmom,) - actuator_moment: actuator moments (nJmom,) - crb: com-based composite inertia and mass (nbody, 10) - qM: total inertia if sparse: (nM,) - if dense: (nv, nv) - qLD: L'*D*L (or Cholesky) factorization of M. if sparse: (nM,) - if dense: (nv, nv) - qLDiagInv: 1/diag(D) if sparse: (nv,) - if dense: (0,) - bvh_aabb_dyn: global bounding box (center, size) (nbvhdynamic, 6) - bvh_active: volume has been added to collisions (nbvh,) - flexedge_velocity: flex edge velocities (nflexedge,) - ten_velocity: tendon velocities (ntendon,) - actuator_velocity: actuator velocities (nu,) - cvel: com-based velocity [3D rot; 3D tran] (nbody, 6) - cdof_dot: time-derivative of cdof (nv, 6) - qfrc_bias: C(qpos,qvel) (nv,) - qfrc_spring: passive spring force (nv,) - qfrc_damper: passive damper force (nv,) - qfrc_gravcomp: passive gravity compensation force (nv,) - qfrc_fluid: passive fluid force (nv,) - qfrc_passive: total passive force (nv,) - subtree_linvel: linear velocity of subtree com (nbody, 3) - subtree_angmom: angular momentum about subtree com (nbody, 3) - qH: L'*D*L factorization of modified M (nM,) - qHDiagInv: 1/diag(D) of modified M (nv,) - B_rownnz: body-dof: non-zeros in each row (nbody,) - B_rowadr: body-dof: address of each row in B_colind (nbody,) - B_colind: body-dof: column indices of non-zeros (nB,) - M_rownnz: inertia: non-zeros in each row (nv,) - M_rowadr: inertia: address of each row in M_colind (nv,) - M_colind: inertia: column indices of non-zeros (nM,) - mapM2M: index mapping from M (legacy) to M (CSR) (nM,) - C_rownnz: reduced dof-dof: non-zeros in each row (nv,) - C_rowadr: reduced dof-dof: address of each row in C_colind (nv,) - C_colind: reduced dof-dof: column indices of non-zeros (nC,) - mapM2C: index mapping from M to C (nC,) - D_rownnz: dof-dof: non-zeros in each row (nv,) - D_rowadr: dof-dof: address of each row in D_colind (nv,) - D_diag: dof-dof: index of diagonal element (nv,) - D_colind: dof-dof: column indices of non-zeros (nD,) - mapM2D: index mapping from M to D (nD,) - mapD2M: index mapping from D to M (nM,) - qDeriv: d (passive + actuator - bias) / d qvel (nD,) - qLU: sparse LU of (qM - dt*qDeriv) (nD,) - actuator_force: actuator force in actuation space (nu,) - qfrc_actuator: actuator force (nv,) - qfrc_smooth: net unconstrained force (nv,) - qacc_smooth: unconstrained acceleration (nv,) - qfrc_constraint: constraint force (nv,) - qfrc_inverse: net external force; should equal: (nv,) - qfrc_applied + J'*xfrc_applied + qfrc_actuator - cacc: com-based acceleration (nbody, 6) - cfrc_int: com-based interaction force with parent (nbody, 6) - cfrc_ext: com-based external force on body (nbody, 6) - contact: all detected contacts (ncon,) - efc_type: constraint type (nefc,) - efc_J: constraint Jacobian (nefc, nv) - efc_pos: constraint position (equality, contact) (nefc,) - efc_margin: inclusion margin (contact) (nefc,) - efc_frictionloss: frictionloss (friction) (nefc,) - efc_D: constraint mass (nefc,) - efc_aref: reference pseudo-acceleration (nefc,) - efc_force: constraint force in constraint space (nefc,) - _qM_sparse: qM in sparse representation (nM,) - _qLD_sparse: qLD in sparse representation (nM,) - _qLDiagInv_sparse: qLDiagInv in sparse representation (nv,) - """ # fmt: skip # constant sizes: + # TODO(stunya): make these sizes jax.Array? ne: int nf: int nl: int nefc: int ncon: int - # solver statistics: + # TODO(stunya): remove most of these fields solver_niter: jax.Array + cdof: jax.Array + cinert: jax.Array + light_xpos: jax.Array + light_xdir: jax.Array + flexvert_xpos: jax.Array + flexelem_aabb: jax.Array + flexedge_J_rownnz: jax.Array # pylint:disable=invalid-name + flexedge_J_rowadr: jax.Array # pylint:disable=invalid-name + flexedge_J_colind: jax.Array # pylint:disable=invalid-name + flexedge_J: jax.Array # pylint:disable=invalid-name + flexedge_length: 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 + ten_length: jax.Array + wrap_obj: jax.Array + wrap_xpos: jax.Array + actuator_length: 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 + qLD: jax.Array # pylint:disable=invalid-name + qLDiagInv: jax.Array # pylint:disable=invalid-name + bvh_aabb_dyn: jax.Array + bvh_active: jax.Array + # position, velocity dependent: + flexedge_velocity: jax.Array + ten_velocity: jax.Array + actuator_velocity: jax.Array + cdof_dot: jax.Array + qH: jax.Array # pylint:disable=invalid-name + qHDiagInv: jax.Array # pylint:disable=invalid-name + 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 + C_rownnz: jax.Array # pylint:disable=invalid-name + C_rowadr: jax.Array # pylint:disable=invalid-name + C_colind: jax.Array # pylint:disable=invalid-name + mapM2C: 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 + qDeriv: jax.Array # pylint:disable=invalid-name + qLU: jax.Array # pylint:disable=invalid-name + qfrc_spring: jax.Array + qfrc_damper: jax.Array + cacc: jax.Array + cfrc_int: jax.Array + cfrc_ext: jax.Array + subtree_linvel: jax.Array + subtree_angmom: 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 + efc_aref: jax.Array + efc_force: jax.Array + + +class DataJAX(PyTreeNode): + """JAX-specific data.""" + + ne: int + nf: int + nl: int + nefc: int + ncon: int + solver_niter: jax.Array + cdof: jax.Array + cinert: jax.Array + ten_wrapadr: jax.Array + ten_wrapnum: jax.Array + ten_J: jax.Array # pylint:disable=invalid-name + ten_length: jax.Array + wrap_obj: jax.Array + wrap_xpos: jax.Array + actuator_length: jax.Array + actuator_moment: jax.Array + crb: jax.Array + qM: jax.Array # pylint:disable=invalid-name + qLD: jax.Array # pylint:disable=invalid-name + qLDiagInv: jax.Array # pylint:disable=invalid-name + ten_velocity: jax.Array + actuator_velocity: jax.Array + cdof_dot: jax.Array + cacc: jax.Array + cfrc_int: jax.Array + cfrc_ext: jax.Array + subtree_linvel: jax.Array + subtree_angmom: jax.Array + # dynamically sized data which are made static due to JAX limitations + 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 + efc_aref: jax.Array + efc_force: jax.Array + + +class Data(PyTreeNode): + """Dynamic state that updates each step.""" + # global properties: time: jax.Array # state: @@ -1419,98 +1101,45 @@ class Data(PyTreeNode): site_xmat: jax.Array cam_xpos: jax.Array cam_xmat: jax.Array - light_xpos: jax.Array = _restricted_to('mujoco') - light_xdir: jax.Array = _restricted_to('mujoco') subtree_com: jax.Array - cdof: jax.Array - cinert: jax.Array - flexvert_xpos: jax.Array = _restricted_to('mujoco') - flexelem_aabb: jax.Array - flexedge_J_rownnz: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - flexedge_J_rowadr: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - flexedge_J_colind: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - flexedge_J: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - flexedge_length: jax.Array = _restricted_to('mujoco') - ten_wrapadr: jax.Array - ten_wrapnum: jax.Array - ten_J_rownnz: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - ten_J_rowadr: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - ten_J_colind: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - ten_J: jax.Array # pylint:disable=invalid-name - ten_length: jax.Array - wrap_obj: jax.Array - wrap_xpos: jax.Array - actuator_length: jax.Array - moment_rownnz: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - moment_rowadr: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - moment_colind: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - actuator_moment: jax.Array - crb: jax.Array - qM: jax.Array # pylint:disable=invalid-name - qLD: jax.Array # pylint:disable=invalid-name - qLDiagInv: jax.Array # pylint:disable=invalid-name - bvh_aabb_dyn: jax.Array = _restricted_to('mujoco') - bvh_active: jax.Array = _restricted_to('mujoco') - # position, velocity dependent: - flexedge_velocity: jax.Array = _restricted_to('mujoco') - ten_velocity: jax.Array - actuator_velocity: jax.Array cvel: jax.Array - cdof_dot: jax.Array qfrc_bias: jax.Array - qfrc_spring: jax.Array = _restricted_to('mujoco') - qfrc_damper: jax.Array = _restricted_to('mujoco') qfrc_gravcomp: jax.Array qfrc_fluid: jax.Array qfrc_passive: jax.Array - subtree_linvel: jax.Array - subtree_angmom: jax.Array - qH: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - qHDiagInv: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - B_rownnz: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - B_rowadr: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - B_colind: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - M_rownnz: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - M_rowadr: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - M_colind: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - mapM2M: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - C_rownnz: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - C_rowadr: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - C_colind: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - mapM2C: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - D_rownnz: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - D_rowadr: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - D_diag: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - D_colind: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - mapM2D: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - mapD2M: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - qDeriv: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - qLU: jax.Array = _restricted_to('mujoco') # pylint:disable=invalid-name - # position, velocity, control & acceleration dependent: qfrc_actuator: jax.Array actuator_force: jax.Array qfrc_smooth: jax.Array qacc_smooth: jax.Array qfrc_constraint: jax.Array qfrc_inverse: jax.Array - cacc: jax.Array - cfrc_int: jax.Array - cfrc_ext: jax.Array - # dynamically sized - contact: Contact - # dynamically sized - position dependent: - 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 - # dynamically sized - position & velocity dependent: - efc_aref: jax.Array - # dynamically sized - position, velocity, control & acceleration dependent: - efc_force: jax.Array - # sparse representation of qM, qLD, qLDiagInv, for compatibility with MuJoCo - # when in dense mode - _qM_sparse: jax.Array = _restricted_to('mjx') # pylint:disable=invalid-name - _qLD_sparse: jax.Array = _restricted_to('mjx') # pylint:disable=invalid-name - _qLDiagInv_sparse: jax.Array = _restricted_to('mjx') # pylint:disable=invalid-name + _impl: Union[DataC, DataJAX] + + @property + def backend_impl(self) -> BackendImpl: + return { + DataC: BackendImpl.C, + DataJAX: BackendImpl.JAX, + }[type(self._impl)] + + def __getattr__(self, name: str): + if name == 'value': + # Special case for NNX, the value attribute may not exist on the parent + # PyTreeNode, before it exists on the child PyTreeNode. Thanks NNX. + return object.__getattribute__(self, 'value') + + try: + impl_instsance = object.__getattribute__(self, '_impl') + val = getattr(impl_instsance, name) + warnings.warn( + f'Accessing `{name}` directly from `Data` is deprecated. ' + f'Access it via `data._impl.{name}` instead.', + DeprecationWarning, + stacklevel=2, + ) + except AttributeError: + # raise the standard exception + raise AttributeError( # pylint: disable=raise-missing-from + f"'{type(self).__name__}' object has no attribute '{name}'" + ) + return val