diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py
index 088728ff..02afca7a 100644
--- a/mjx/mujoco/mjx/_src/forward.py
+++ b/mjx/mujoco/mjx/_src/forward.py
@@ -103,6 +103,8 @@ def fwd_velocity(m: Model, d: Data) -> Data:
@named_scope
def fwd_actuation(m: Model, d: Data) -> Data:
"""Actuation-dependent computations."""
+ if not isinstance(d._impl, DataJAX):
+ raise ValueError('fwd_actuation requires JAX backend implementation.')
if not m.nu or m.opt.disableflags & DisableBit.ACTUATION:
return d.replace(
act_dot=jp.zeros((m.na,)),
diff --git a/mjx/mujoco/mjx/_src/inverse.py b/mjx/mujoco/mjx/_src/inverse.py
index 086dd101..2b045ea4 100644
--- a/mjx/mujoco/mjx/_src/inverse.py
+++ b/mjx/mujoco/mjx/_src/inverse.py
@@ -25,6 +25,7 @@ from mujoco.mjx._src import support
from mujoco.mjx._src.types import Data
from mujoco.mjx._src.types import DisableBit
from mujoco.mjx._src.types import EnableBit
+from mujoco.mjx._src.types import Impl
from mujoco.mjx._src.types import IntegratorType
from mujoco.mjx._src.types import Model
@@ -83,6 +84,7 @@ def inv_constraint(m: Model, d: Data) -> Data:
def inverse(m: Model, d: Data) -> Data:
"""Inverse dynamics."""
+
d = forward.fwd_position(m, d)
d = sensor.sensor_pos(m, d)
d = forward.fwd_velocity(m, d)
diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py
index c1012dc0..abf5f952 100644
--- a/mjx/mujoco/mjx/_src/io.py
+++ b/mjx/mujoco/mjx/_src/io.py
@@ -21,6 +21,7 @@ from typing import Any, Dict, List, Optional, Tuple, Union
import warnings
import jax
+import jax.experimental
from jax import numpy as jp
from jax.extend import backend
import mujoco
@@ -82,7 +83,7 @@ def _resolve_device(
logging.debug('Picking default device: %s.', device_0)
return device_0
- if impl == types.Impl.C:
+ if impl == types.Impl.C or impl == types.Impl.CPP:
cpu_0 = jax.devices('cpu')[0]
logging.debug('Picking default device: %s', cpu_0)
return cpu_0
@@ -124,7 +125,7 @@ def _check_impl_device_compatibility(
)
is_cpu_device = device.platform == 'cpu'
- if impl == types.Impl.C:
+ if impl == types.Impl.C or impl == types.Impl.CPP:
if not is_cpu_device:
raise AssertionError(
f'C implementation requires a CPU device, got {device}.'
@@ -489,6 +490,39 @@ def _put_model_warp(
return _strip_weak_type(model)
+def _put_model_cpp(
+ m: mujoco.MjModel,
+ device: Optional[jax.Device] = None,
+) -> types.Model:
+ """Puts mujoco.MjModel onto a device, resulting in mjx.Model."""
+
+ mj_field_names = {f.name for f in types.Model.fields() if f.name != '_impl'}
+ fields = {f: getattr(m, f) for f in mj_field_names}
+ fields['cam_mat0'] = fields['cam_mat0'].reshape((-1, 3, 3))
+ fields['opt'] = _put_option(m.opt, impl=types.Impl.C)
+ fields['stat'] = _put_statistic(m.stat, impl=types.Impl.C)
+
+ # get the pointer address
+ # we use a 0-d array
+ addr = m._address # pytype: disable=attribute-error
+ # To ensure that we retain the full pointer even if jax.config.enable_x64 is
+ # set to True, we store the pointer as two 32-bit values. In the FFI call,
+ # we combine the two values into a single pointer value.
+ pointer_lo = jp.array(addr & 0xFFFFFFFF, dtype=jp.uint32)
+ pointer_hi = jp.array(addr >> 32, dtype=jp.uint32)
+ c_pointers_impl = types.ModelCPP(
+ pointer_lo=pointer_lo,
+ pointer_hi=pointer_hi,
+ _model=m,
+ )
+
+ model = types.Model(
+ **{k: copy.copy(v) for k, v in fields.items()}, _impl=c_pointers_impl
+ )
+ model = jax.device_put(model, device=device)
+ return _strip_weak_type(model)
+
+
def put_model(
m: mujoco.MjModel,
device: Optional[jax.Device] = None,
@@ -515,6 +549,8 @@ def put_model(
return _put_model_c(m, device)
elif impl == types.Impl.WARP:
return _put_model_warp(m, device)
+ elif impl == types.Impl.CPP:
+ return _put_model_cpp(m, device)
else:
raise ValueError(f'Unsupported implementation: {impl}')
@@ -612,7 +648,7 @@ def _make_data_jax(
efc_address = constraint.make_efc_address(m, dim, efc_type)
float_ = jp.zeros(1, float).dtype
- int_ = jp.zeros(1, int).dtype
+ int_ = np.int32
contact = _make_data_contact_jax(dim, efc_address)
if m.opt.cone == types.ConeType.ELLIPTIC and np.any(contact.dim == 1):
@@ -1216,6 +1252,62 @@ def _put_data_c(
return _strip_weak_type(data)
+def _put_data_cpp(
+ m: mujoco.MjModel,
+ d: mujoco.MjData,
+ device: Optional[jax.Device] = None,
+ dummy_arg_for_batching: Optional[jax.Array] = None,
+) -> types.Data:
+ """Puts mujoco.MjData onto a device, resulting in mjx.Data."""
+
+ data_list = []
+
+ def _copy_and_get_addr(unused_jax_array):
+ # We use the input to the callback as a dummy dependency to ensure
+ # io_callback runs for each element in the batch.
+ try:
+ new_d = mujoco.MjData(m)
+ except mujoco.FatalError as e:
+ raise ValueError('Failed to create new MjData') from e
+ mujoco.mj_copyState(m, d, new_d, int(mujoco.mjtState.mjSTATE_FULLPHYSICS))
+ mujoco.mj_forward(m, new_d)
+ data_list.append(new_d)
+ addr = new_d._address
+ # To ensure that we retain the full pointer even if jax.config.enable_x64 is
+ # set to True, we store the pointer as two 32-bit values. In the FFI call,
+ # we combine the two values into a single pointer value.
+ return (
+ np.array(addr & 0xFFFFFFFF, dtype=np.uint32),
+ np.array(addr >> 32, dtype=np.uint32),
+ )
+
+ # Pass a dummy dependency to ensure io_callback runs across the batch.
+ pointer_lo, pointer_hi = jax.experimental.io_callback(
+ _copy_and_get_addr,
+ (
+ jax.ShapeDtypeStruct((), jp.uint32),
+ jax.ShapeDtypeStruct((), jp.uint32),
+ ),
+ dummy_arg_for_batching,
+ )
+
+ new_d = data_list[0]
+ fields = _put_data_public_fields(new_d)
+
+ c_pointers_impl = types.DataCPP(
+ pointer_lo=pointer_lo,
+ pointer_hi=pointer_hi,
+ _data=data_list,
+ )
+
+ data = types.Data(
+ _impl=c_pointers_impl,
+ **fields,
+ )
+ data = jax.device_put(data, device=device)
+ return _strip_weak_type(data)
+
+
def put_data(
m: mujoco.MjModel,
d: mujoco.MjData,
@@ -1224,6 +1316,7 @@ def put_data(
nconmax: Optional[int] = None,
naconmax: Optional[int] = None,
njmax: Optional[int] = None,
+ dummy_arg_for_batching: Optional[jax.Array] = None,
) -> types.Data:
"""Puts mujoco.MjData onto a device, resulting in mjx.Data.
@@ -1238,6 +1331,8 @@ def put_data(
`naconmax` argument to set the upper bound for the number of contacts
across all worlds, rather than the `nconmax` argument from MuJoCo Warp.
njmax: maximum number of constraints to allocate for warp
+ dummy_arg_for_batching: dummy argument to use for batching in cpp
+ implementation
Returns:
an mjx.Data placed on device
@@ -1256,6 +1351,10 @@ def put_data(
return _put_data_jax(m, d, device)
elif impl == types.Impl.C:
return _put_data_c(m, d, device)
+ elif impl == types.Impl.CPP:
+ return _put_data_cpp(
+ m, d, device, dummy_arg_for_batching=dummy_arg_for_batching
+ )
# TODO(robotics-team): implement put_data_warp
diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py
index 63d74bdf..cd76b40d 100644
--- a/mjx/mujoco/mjx/_src/io_test.py
+++ b/mjx/mujoco/mjx/_src/io_test.py
@@ -136,7 +136,7 @@ class ModelIOTest(parameterized.TestCase):
@parameterized.product(
xml=(_MULTIPLE_CONVEX_OBJECTS, _MULTIPLE_CONSTRAINTS),
- impl=('jax', 'c', 'warp'),
+ impl=('jax', 'c', 'warp', 'cpp'),
)
@mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'true'})
def test_put_model(self, xml, impl):
@@ -181,6 +181,9 @@ class ModelIOTest(parameterized.TestCase):
self.assertTrue(hasattr(mx.opt._impl, 'ls_parallel'))
# Fields private to Warp backend impl are populated.
self.assertTrue(hasattr(mx._impl, 'nxn_geom_pair'))
+ elif impl == 'cpp':
+ self.assertTrue(hasattr(mx._impl, 'pointer_lo'))
+ self.assertTrue(hasattr(mx._impl, 'pointer_hi'))
np.testing.assert_allclose(mx.body_parentid, m.body_parentid)
np.testing.assert_allclose(mx.geom_type, m.geom_type)
@@ -410,6 +413,10 @@ class DataIOTest(parameterized.TestCase):
self.assertEqual(d._impl.actuator_moment.shape, (1, nv))
elif impl == 'c':
self.assertEqual(d._impl.actuator_moment.shape, (m.nJmom,))
+ elif impl == 'cpp':
+ self.assertTrue(hasattr(d._impl, 'pointer_lo'))
+ self.assertTrue(hasattr(d._impl, 'pointer_hi'))
+ return # cpp does not populate other fields in _impl
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))
@@ -497,6 +504,10 @@ class DataIOTest(parameterized.TestCase):
np.testing.assert_allclose(dx._impl.qM, d.qM)
np.testing.assert_allclose(dx._impl.qLD, d.qLD)
np.testing.assert_allclose(dx._impl.qLDiagInv, d.qLDiagInv)
+ elif impl == 'cpp':
+ self.assertTrue(hasattr(dx._impl, 'pointer_lo'))
+ self.assertTrue(hasattr(dx._impl, 'pointer_hi'))
+ return # cpp does not populate other fields in _impl
# 4 contacts, 2 for each capsule against the plane
self.assertEqual(dx._impl.contact.dist.shape, (4,))
@@ -816,6 +827,39 @@ class DataIOTest(parameterized.TestCase):
self.assertEqual(dx._impl.contact__dist.shape, (dx._impl.naconmax,))
self.assertEqual(dx[0]._impl.contact__dist.shape, (dx._impl.naconmax,))
+ def test_put_data_cpp(self):
+ m = mujoco.MjModel.from_xml_string("""
+
+
+
+
+
+
+
+
+ """)
+ d = mujoco.MjData(m)
+ d.qpos[0] = 1.0
+ unused_mjx_data = mjx_io.put_data(m, d, impl='cpp')
+
+ def put_data(dummy_arg_for_batching):
+ dx = mjx_io.put_data(
+ m, d, impl='cpp', dummy_arg_for_batching=dummy_arg_for_batching
+ )
+ return dx
+
+ vmjx_data = jax.vmap(put_data, in_axes=0, out_axes=0)(
+ jp.zeros(2, dtype=jp.uint32)
+ )
+
+ self.assertEqual(vmjx_data.qpos.shape, (2, m.nq))
+ self.assertEqual(len(vmjx_data._impl._data), 2)
+ # check that the data pointers in fact point to different datas
+ self.assertNotEqual(
+ vmjx_data._impl._data[0]._address,
+ vmjx_data._impl._data[1]._address,
+ )
+
# Test cases for `_resolve_impl_and_device` where the device is
# specified by the user and the device is available.
@@ -1117,6 +1161,39 @@ class StateIOTest(parameterized.TestCase):
mx = mjx.put_model(m)
self.assertEqual(mjx.state_size(mx, spec), mujoco.mj_stateSize(m, spec))
+ def test_put_data_cpp(self):
+ m = mujoco.MjModel.from_xml_string("""
+
+
+
+
+
+
+
+
+ """)
+ d = mujoco.MjData(m)
+ d.qpos[0] = 1.0
+ unused_mjx_data = mjx_io.put_data(m, d, impl='cpp')
+
+ def put_data(dummy_arg_for_batching):
+ dx = mjx_io.put_data(
+ m, d, impl='cpp', dummy_arg_for_batching=dummy_arg_for_batching
+ )
+ return dx
+
+ vmjx_data = jax.vmap(put_data, in_axes=0, out_axes=0)(
+ jp.zeros(2, dtype=jp.uint32)
+ )
+
+ self.assertEqual(vmjx_data.qpos.shape, (2, m.nq))
+ self.assertEqual(len(vmjx_data._impl._data), 2)
+ # check that the data pointers in fact point to different datas
+ self.assertNotEqual(
+ vmjx_data._impl._data[0]._address,
+ vmjx_data._impl._data[1]._address,
+ )
+
def test_get_set_state(self):
m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS)
d = mujoco.MjData(m)
diff --git a/mjx/mujoco/mjx/_src/passive.py b/mjx/mujoco/mjx/_src/passive.py
index 61bcd9a4..dc574fbf 100644
--- a/mjx/mujoco/mjx/_src/passive.py
+++ b/mjx/mujoco/mjx/_src/passive.py
@@ -34,6 +34,9 @@ from mujoco.mjx._src.types import OptionJAX
def _spring_damper(m: Model, d: Data) -> jax.Array:
"""Applies joint level spring and damping forces."""
+ if not isinstance(d._impl, DataJAX) and not isinstance(m._impl, ModelJAX):
+ raise ValueError('_spring_damper requires JAX backend implementation.')
+ assert isinstance(d._impl, DataJAX) and isinstance(m._impl, ModelJAX)
def fn(jnt_typs, stiffness, qpos_spring, qpos):
qpos_i = 0
diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py
index ab8fcea6..55437fbf 100644
--- a/mjx/mujoco/mjx/_src/smooth.py
+++ b/mjx/mujoco/mjx/_src/smooth.py
@@ -44,7 +44,6 @@ def kinematics(m: Model, d: Data) -> Data:
from mujoco.mjx.warp import smooth as mjxw_smooth # pylint: disable=g-import-not-at-top # pytype: disable=import-error
return mjxw_smooth.kinematics(m, d)
-
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
@@ -140,6 +139,7 @@ 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.')
@@ -412,6 +412,7 @@ def solve_m(m: Model, d: Data, x: jax.Array) -> jax.Array:
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.')
diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py
index e13a0c5a..40f2d099 100644
--- a/mjx/mujoco/mjx/_src/types.py
+++ b/mjx/mujoco/mjx/_src/types.py
@@ -14,8 +14,9 @@
# ==============================================================================
"""Base types used in MJX."""
+import dataclasses
import enum
-from typing import Tuple, Union
+from typing import Any, Tuple, Union
import warnings
import jax
@@ -29,6 +30,7 @@ class Impl(enum.Enum):
"""Implementation to use."""
C = 'c'
+ CPP = 'cpp'
JAX = 'jax'
WARP = 'warp'
@@ -529,6 +531,26 @@ class Option(PyTreeNode):
_impl: Union[OptionJAX, OptionC, mjxw_types.OptionWarp]
+class ModelCPP(PyTreeNode):
+ """Minimal Model implementation holding only the pointer."""
+ # To ensure that we retain the full pointer even if jax.config.enable_x64 is
+ # set to True, we store the pointer as two 32-bit values. In the FFI call,
+ # we combine the two values into a single pointer value.
+ pointer_lo: jax.Array
+ pointer_hi: jax.Array
+ _model: mujoco.MjModel
+
+
+class DataCPP(PyTreeNode):
+ """Minimal Data implementation holding only the pointer."""
+ # To ensure that we retain the full pointer even if jax.config.enable_x64 is
+ # set to True, we store the pointer as two 32-bit values. In the FFI call,
+ # we combine the two values into a single pointer value.
+ pointer_lo: jax.Array
+ pointer_hi: jax.Array
+ _data: list[Any] = dataclasses.field(default_factory=list, repr=False)
+
+
class ModelC(PyTreeNode):
"""CPU-specific model data."""
@@ -943,6 +965,7 @@ class Model(PyTreeNode):
def impl(self) -> Impl:
return {
ModelC: Impl.C,
+ ModelCPP: Impl.CPP,
ModelJAX: Impl.JAX,
mjxw_types.ModelWarp: Impl.WARP,
}[type(self._impl)]
@@ -1180,12 +1203,13 @@ class Data(PyTreeNode):
qacc_smooth: jax.Array
qfrc_constraint: jax.Array
qfrc_inverse: jax.Array
- _impl: Union[DataC, DataJAX, mjxw_types.DataWarp]
+ _impl: Union[DataC, DataCPP, DataJAX, mjxw_types.DataWarp]
@property
def impl(self) -> Impl:
return {
DataC: Impl.C,
+ DataCPP: Impl.CPP,
DataJAX: Impl.JAX,
mjxw_types.DataWarp: Impl.WARP,
}[type(self._impl)]