diff --git a/doc/changelog.rst b/doc/changelog.rst index 69851a63..ac0cd959 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -14,20 +14,29 @@ General MJX ^^^ -3. Added cylinder plane collisions. -4. Added ``efc_type`` to ``mjx.Data`` and ``dim``, ``efc_address`` to ``mjx.Contact``. -5. Added ``geom`` to ``mjx.Contact`` and marked ``geom1``, ``geom2`` deprecated. -6. Added ``ne``, ``nf``, ``nl``, ``nefc``, and ``ncon`` to ``mjx.Data`` to match ``mujoco.MjData``. -7. Given the above added fields, removed ``mjx.get_params``, ``mjx.ncon``, and ``mjx.count_constraints``. -8. Changed the way meshes are organized on device to speed up collision detection when a mesh is replicated for many +.. admonition:: Breaking API changes + :class: attention + + 3. Removed deprecated ``mjx.device_get_into`` and ```mjx.device_put``` functions as they lack critical new + functionality. + + **Migration:** Use ``mjx.get_data_into`` instead of ``mjx.device_get_into``, and ``mjx.put_data`` instead of + ``mjx.device_put``. + +4. Added cylinder plane collisions. +5. Added ``efc_type`` to ``mjx.Data`` and ``dim``, ``efc_address`` to ``mjx.Contact``. +6. Added ``geom`` to ``mjx.Contact`` and marked ``geom1``, ``geom2`` deprecated. +7. Added ``ne``, ``nf``, ``nl``, ``nefc``, and ``ncon`` to ``mjx.Data`` to match ``mujoco.MjData``. +8. Given the above added fields, removed ``mjx.get_params``, ``mjx.ncon``, and ``mjx.count_constraints``. +9. Changed the way meshes are organized on device to speed up collision detection when a mesh is replicated for many geoms. -9. Fixed a bug where capsules might be ignored in broadphase colliision checking. +10. Fixed a bug where capsules might be ignored in broadphase colliision checking. Bug fixes ^^^^^^^^^ -10. Defaults of lights were not being saved, now fixed. -11. Prevent overwriting of frame names by body names when saving an XML. Introduced in 3.1.4. -12. Fixed bug in Python binding of :ref:`mj_saveModel`: ``buffer`` argument was documented as optional but was actually +11. Defaults of lights were not being saved, now fixed. +12. Prevent overwriting of frame names by body names when saving an XML. Introduced in 3.1.4. +13. Fixed bug in Python binding of :ref:`mj_saveModel`: ``buffer`` argument was documented as optional but was actually not optional. diff --git a/mjx/mujoco/mjx/__init__.py b/mjx/mujoco/mjx/__init__.py index 01939ab7..120db360 100644 --- a/mjx/mujoco/mjx/__init__.py +++ b/mjx/mujoco/mjx/__init__.py @@ -17,8 +17,6 @@ # pylint:disable=g-importing-member from mujoco.mjx._src.collision_driver import collision from mujoco.mjx._src.constraint import make_constraint -from mujoco.mjx._src.device import device_get_into -from mujoco.mjx._src.device import device_put from mujoco.mjx._src.forward import euler from mujoco.mjx._src.forward import forward from mujoco.mjx._src.forward import fwd_acceleration diff --git a/mjx/mujoco/mjx/_src/constraint_test.py b/mjx/mujoco/mjx/_src/constraint_test.py index 133640e3..99346828 100644 --- a/mjx/mujoco/mjx/_src/constraint_test.py +++ b/mjx/mujoco/mjx/_src/constraint_test.py @@ -65,7 +65,7 @@ class ConstraintTest(absltest.TestCase): pos = jp.ones(3) m.opt.disableflags = m.opt.disableflags | mjx.DisableBit.REFSAFE - mx = mjx.device_put(m) + mx = mjx.put_model(m) k, *_ = constraint._kbi(mx, solimp, solref, pos) self.assertEqual(k, 1 / (0.99**2 * timeconst**2)) diff --git a/mjx/mujoco/mjx/_src/device.py b/mjx/mujoco/mjx/_src/device.py deleted file mode 100644 index 43b3d256..00000000 --- a/mjx/mujoco/mjx/_src/device.py +++ /dev/null @@ -1,312 +0,0 @@ -# Copyright 2023 DeepMind Technologies Limited -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -# ============================================================================== -"""Get and put mujoco data on/off device.""" - -import copy -import dataclasses -from typing import Any, Dict, Iterable, List, Union, overload -import warnings - -import jax -from jax import numpy as jp -import mujoco -from mujoco.mjx._src import collision_driver -from mujoco.mjx._src import types -import numpy as np - -_MJ_TYPE_ATTR = { - mujoco.mjtBias: (mujoco.MjModel.actuator_biastype,), - mujoco.mjtDyn: (mujoco.MjModel.actuator_dyntype,), - mujoco.mjtEq: (mujoco.MjModel.eq_type,), - mujoco.mjtGain: (mujoco.MjModel.actuator_gaintype,), - mujoco.mjtTrn: (mujoco.MjModel.actuator_trntype,), - mujoco.mjtCone: ( - mujoco.MjModel.opt, - mujoco.MjOption.cone, - ), - mujoco.mjtIntegrator: ( - mujoco.MjModel.opt, - mujoco.MjOption.integrator, - ), - mujoco.mjtSolver: ( - mujoco.MjModel.opt, - mujoco.MjOption.solver, - ), -} - -_TYPE_MAP = { - mujoco._structs._MjContactList: types.Contact, # pylint: disable=protected-access - mujoco.MjData: types.Data, - mujoco.MjModel: types.Model, - mujoco.MjOption: types.Option, - mujoco.MjStatistic: types.Statistic, - mujoco.mjtBias: types.BiasType, - mujoco.mjtCone: types.ConeType, - mujoco.mjtDisableBit: types.DisableBit, - mujoco.mjtDyn: types.DynType, - mujoco.mjtEq: types.EqType, - mujoco.mjtGain: types.GainType, - mujoco.mjtIntegrator: types.IntegratorType, - mujoco.mjtSolver: types.SolverType, - mujoco.mjtTrn: types.TrnType, -} - -_TRANSFORMS = { - (types.Data, 'ximat'): lambda x: x.reshape(x.shape[:-1] + (3, 3)), - (types.Data, 'xmat'): lambda x: x.reshape(x.shape[:-1] + (3, 3)), - (types.Data, 'geom_xmat'): lambda x: x.reshape(x.shape[:-1] + (3, 3)), - (types.Data, 'site_xmat'): lambda x: x.reshape(x.shape[:-1] + (3, 3)), - (types.Data, 'cam_xmat'): lambda x: x.reshape(x.shape[:-1] + (3, 3)), - (types.Model, 'cam_mat0'): lambda x: x.reshape(x.shape[:-1] + (3, 3)), - (types.Contact, 'frame'): ( - lambda x: x.reshape(x.shape[:-1] + (3, 3)) # pylint: disable=g-long-lambda - if x is not None and x.shape[0] - else jp.zeros((0, 3, 3)) - ), -} - -_INVERSE_TRANSFORMS = { - (types.Data, 'ximat'): lambda x: x.reshape(x.shape[:-2] + (9,)), - (types.Data, 'xmat'): lambda x: x.reshape(x.shape[:-2] + (9,)), - (types.Data, 'geom_xmat'): lambda x: x.reshape(x.shape[:-2] + (9,)), - (types.Data, 'site_xmat'): lambda x: x.reshape(x.shape[:-2] + (9,)), - (types.Data, 'cam_xmat'): lambda x: x.reshape(x.shape[:-2] + (9,)), - (types.Model, 'cam_mat0'): lambda x: x.reshape(x.shape[:-2] + (9,)), - (types.Contact, 'frame'): ( - lambda x: x.reshape(x.shape[:-2] + (9,)) # pylint: disable=g-long-lambda - if x is not None and x.shape[0] - else jp.zeros((0, 9)) - ), -} - -_DERIVED = { - # efc_J is dense in MJX, sparse in MJ. ignore for now. - (types.Data, 'efc_J'), (types.Option, 'has_fluid_params') -} - - -def _data_derived(value: mujoco.MjData) -> Dict[str, Any]: - return {'efc_J': jax.device_put(value.efc_J)} - - -def _option_derived(value: types.Option) -> Dict[str, Any]: - has_fluid = ( - value.density > 0 or value.viscosity > 0 or (value.wind != 0.0).any() - ) - return {'has_fluid_params': has_fluid} - - -def _validate(m: mujoco.MjModel): - """Validates that an mjModel is compatible with MJX.""" - - # check enum types - for mj_type, attrs in _MJ_TYPE_ATTR.items(): - val = m - for attr in attrs: - val = attr.fget(val) # pytype: disable=attribute-error - - typs = set(val) if isinstance(val, Iterable) else {val} - unsupported_typs = typs - set(_TYPE_MAP[mj_type]) - unsupported = [mj_type(t) for t in unsupported_typs] # pylint: disable=too-many-function-args - if unsupported: - raise NotImplementedError(f'{unsupported} not implemented.') - - if m.ntendon: - raise NotImplementedError('Tendons are not supported.') - - # check condim - if (m.geom_condim != 3).any() or (m.pair_dim != 3).any(): - raise NotImplementedError('Only condim=3 is supported.') - - if m.body_gravcomp.any(): - raise NotImplementedError('gravcomp is not supported') - - 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): - t1, t2 = mujoco.mjtGeom(t1), mujoco.mjtGeom(t2) - raise NotImplementedError(f'({t1}, {t2}) collisions not implemented.') - # margin/gap not supported for geoms - if mujoco.mjtGeom.mjGEOM_MESH in (t1, t2): - if ip != -1: - margin = m.pair_margin[ip] - else: - margin = m.geom_margin[g1] + m.geom_margin[g2] - if margin.any(): - t1, t2 = mujoco.mjtGeom(t1), mujoco.mjtGeom(t2) - raise NotImplementedError(f'({t1}, {t2}) margin/gap not implemented.') - - # TODO(erikfrey): warn for high solver iterations, nefc, etc. - - # mjNDISABLE is not a DisableBit flag, so must be explicitly ignored - disablebit_members = set(mujoco.mjtDisableBit.__members__.values()) - { - mujoco.mjtDisableBit.mjNDISABLE} - unsupported_disable = disablebit_members - { - mujoco.mjtDisableBit(t.value) for t in types.DisableBit - } - for f in unsupported_disable: - if f & m.opt.disableflags: - warnings.warn(f'Ignoring disable flag {f.name}.') - - # mjNENABLE is not an EnableBit flag, so must be explicitly ignored - unsupported_enable = set(mujoco.mjtEnableBit.__members__.values()) - { - mujoco.mjtEnableBit.mjNENABLE - } - for f in unsupported_enable: - if f & m.opt.enableflags: - warnings.warn(f'Ignoring enable flag {f.name}.') - - if not np.allclose(m.dof_frictionloss, 0): - raise NotImplementedError('dof_frictionloss is not implemented.') - - -@overload -def device_put(value: mujoco.MjData) -> types.Data: - ... - - -@overload -def device_put(value: mujoco.MjModel) -> types.Model: - ... - - -def device_put(value): - """Places mujoco data onto a device. - - Args: - value: a mujoco struct to transfer - - Returns: - on-device MJX struct reflecting the input value - """ - warnings.warn( - 'device_put is deprecated, use put_model and put_data instead', - category=DeprecationWarning, - ) - - clz = _TYPE_MAP.get(type(value)) - if clz is None: - raise NotImplementedError(f'{type(value)} is not supported for device_put.') - - if isinstance(value, mujoco.MjModel): - _validate(value) # type: ignore - - init_kwargs = {} - for f in dataclasses.fields(clz): # type: ignore - if (clz, f.name) in _DERIVED: - continue - - field_value = getattr(value, f.name) - if (clz, f.name) in _TRANSFORMS: - field_value = _TRANSFORMS[(clz, f.name)](field_value) - - if f.type is jax.Array: - field_value = jax.device_put(field_value) - elif type(field_value) in _TYPE_MAP.keys(): - field_value = device_put(field_value) - - init_kwargs[f.name] = copy.copy(field_value) - - derived_kwargs = {} - if isinstance(value, mujoco.MjModel): - derived_kwargs = {} - elif isinstance(value, mujoco.MjData): - derived_kwargs = _data_derived(value) - elif isinstance(value, mujoco.MjOption): - derived_kwargs = _option_derived(value) - - return clz(**init_kwargs, **derived_kwargs) # type: ignore - - -@overload -def device_get_into( - result: Union[mujoco.MjData, List[mujoco.MjData]], value: types.Data -): - ... - - -def device_get_into(result, value): - """Transfers data off device into a mujoco MjData. - - Data on device often has a batch dimension which adds (N,) to the beginning - of each array shape where N = batch size. - - If result is a single MjData, arrays are copied over with the batch dimension - intact. If result is a list, the list must be length N and will be populated - with distinct MjData structs where the batch dimension is stripped. - - Args: - result: struct (or list of structs) to transfer into - value: device value to transfer - - Raises: - RuntimeError: if result length doesn't match data batch size - """ - warnings.warn( - 'device_get_into is deprecated, use get_data instead', - category=DeprecationWarning, - ) - - value = jax.device_get(value) - - if isinstance(result, list): - array_shapes = [s.shape for s in jax.tree_util.tree_flatten(value)[0]] - - if any(len(s) < 1 or s[0] != array_shapes[0][0] for s in array_shapes): - raise ValueError('unrecognizable batch dimension in value') - - batch_size = array_shapes[0][0] - - if len(result) != batch_size: - raise ValueError( - f"result length ({len(result)}) doesn't match value batch size" - f' ({batch_size})' - ) - - for i in range(batch_size): - value_i = jax.tree_util.tree_map(lambda x, i=i: x[i], value) - device_get_into(result[i], value_i) - - else: - if isinstance(result, mujoco.MjData): - ncon = value.contact.dist.shape[0] - nefc = value.efc_J.shape[0] - mujoco._functions._realloc_con_efc( # pylint: disable=protected-access - result, ncon=ncon, nefc=nefc - ) - result.ncon = ncon - result.nefc = nefc - efc_start = nefc - ncon * 4 - result.contact.efc_address[:] = np.arange(efc_start, nefc, 4) - result.contact.dim[:] = 3 - - for f in dataclasses.fields(value): # type: ignore - if (type(value), f.name) in _DERIVED: - continue - - field_value = getattr(value, f.name) - - if (type(value), f.name) in _INVERSE_TRANSFORMS: - field_value = _INVERSE_TRANSFORMS[(type(value), f.name)](field_value) - - if type(field_value) in _TYPE_MAP.values(): - device_get_into(getattr(result, f.name), field_value) - continue - - try: - setattr(result, f.name, field_value) - except AttributeError: - getattr(result, f.name)[:] = field_value diff --git a/mjx/mujoco/mjx/_src/device_test.py b/mjx/mujoco/mjx/_src/device_test.py deleted file mode 100644 index 530a4052..00000000 --- a/mjx/mujoco/mjx/_src/device_test.py +++ /dev/null @@ -1,203 +0,0 @@ -# Copyright 2023 DeepMind Technologies Limited -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -# ============================================================================== -"""Tests for moving mujoco structs on and off device.""" - -import dataclasses - -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 device -from mujoco.mjx._src import test_util -from mujoco.mjx._src import types -# pylint: disable=g-importing-member -from mujoco.mjx._src.dataclasses import PyTreeNode -# pylint: enable=g-importing-member -import numpy as np - - -def _assert_eq(testcase, a, b, attr=None, name=None): - if (type(a), attr) in device._DERIVED: - return - - if attr: - a, b = getattr(a, attr), getattr(b, attr) - - if isinstance(a, PyTreeNode): - for field in dataclasses.fields(a): - _assert_eq(testcase, a, b, field.name, type(a).__name__) - return - - typ = {'Model': types.Model, 'Data': types.Data, - 'Contact': types.Contact}.get(name) - if (typ, attr) in device._TRANSFORMS: - b = device._TRANSFORMS[(typ, attr)](b) - - err_msg = f'mismatch: {attr} in {name}' - if not hasattr(b, 'shape') or not b.shape: - testcase.assertEqual(a, b, err_msg) - return - - a, b = np.array(a), np.array(b) - np.testing.assert_allclose(a, b, err_msg=err_msg, atol=1e-8) - - -class DeviceTest(parameterized.TestCase): - - @parameterized.parameters('constraints.xml', 'pendula.xml') - def testdevice_put(self, fname): - """Test putting MjData and MjModel on device.""" - m = test_util.load_test_file(fname) - # advance state to ensure non-zero fields - d = mujoco.MjData(m) - for _ in range(10): - mujoco.mj_step(m, d) - - _assert_eq(self, mjx.device_put(d), d) - _assert_eq(self, mjx.device_put(m), m) - - @parameterized.parameters('constraints.xml', 'pendula.xml') - def testdevice_get(self, fname): - """Test getting MjData from a device.""" - m = test_util.load_test_file(fname) - m.opt.jacobian = mujoco.mjtJacobian.mjJAC_SPARSE # force sparse for testing - mx = device.device_put(m) - dx = mjx.make_data(mx) - d = mujoco.MjData(m) - device.device_get_into(d, dx) - _assert_eq(self, dx, d) - - @parameterized.parameters('constraints.xml', 'pendula.xml') - def testdevice_get_batched(self, fname): - """Test getting MjData from a device.""" - m = test_util.load_test_file(fname) - m.opt.jacobian = mujoco.mjtJacobian.mjJAC_SPARSE # force sparse for testing - mx = device.device_put(m) - batch_size = 32 - - # create mjx_data and batch it - dx = mjx.make_data(mx) - dx = jax.tree_util.tree_map( - lambda x: jp.repeat(x, batch_size).reshape((batch_size,) + x.shape), - dx, - ) - ds = [mujoco.MjData(m) for _ in range(batch_size - 1)] - - with self.assertRaises(ValueError): - device.device_get_into(ds, dx) - - ds = [mujoco.MjData(m) for _ in range(batch_size)] - device.device_get_into(ds, dx) - dx = jax.device_get(dx) # faster indexing for testing - for i in range(batch_size): - _assert_eq(self, jax.tree_util.tree_map(lambda x, i=i: x[i], dx), ds[i]) - - -class ValidateInputTest(absltest.TestCase): - - def test_solver(self): - m = mujoco.MjModel.from_xml_string( - '' - ) - with self.assertRaises(NotImplementedError): - mjx.device_put(m) - - def test_integrator(self): - m = mujoco.MjModel.from_xml_string( - '' - ) - with self.assertRaises(NotImplementedError): - mjx.device_put(m) - - def test_cone(self): - m = mujoco.MjModel.from_xml_string( - '' - ) - with self.assertRaises(NotImplementedError): - mjx.device_put(m) - - def test_dyn(self): - m = test_util.load_test_file('pendula.xml') - m.actuator_dyntype[0] = mujoco.mjtDyn.mjDYN_MUSCLE - with self.assertRaises(NotImplementedError): - mjx.device_put(m) - - def test_gain(self): - m = test_util.load_test_file('pendula.xml') - m.actuator_gaintype[0] = mujoco.mjtGain.mjGAIN_MUSCLE - with self.assertRaises(NotImplementedError): - mjx.device_put(m) - - def test_bias(self): - m = test_util.load_test_file('pendula.xml') - m.actuator_gaintype[0] = mujoco.mjtGain.mjGAIN_MUSCLE - with self.assertRaises(NotImplementedError): - mjx.device_put(m) - - def test_condim(self): - m = test_util.load_test_file('constraints.xml') - for i in [1, 4, 6]: - m.geom_condim[0] = i - with self.assertRaises(NotImplementedError): - mjx.device_put(m) - - def test_geoms(self): - m = mujoco.MjModel.from_xml_string(""" - - - - - - - - - - - - - """) - with self.assertRaises(NotImplementedError): - mjx.device_put(m) - - def test_tendon(self): - m = mujoco.MjModel.from_xml_string(""" - - - - - - - - - - - - - - - - - - - """) - with self.assertRaises(NotImplementedError): - mjx.device_put(m) - - -if __name__ == '__main__': - absltest.main() diff --git a/mjx/mujoco/mjx/_src/scan_test.py b/mjx/mujoco/mjx/_src/scan_test.py index 456067d6..edb57339 100644 --- a/mjx/mujoco/mjx/_src/scan_test.py +++ b/mjx/mujoco/mjx/_src/scan_test.py @@ -55,7 +55,7 @@ class ScanTest(absltest.TestCase): """) - m = mjx.device_put(m) + m = mjx.put_model(m) def fn(body_id): return body_id + 1 @@ -69,7 +69,7 @@ class ScanTest(absltest.TestCase): def test_flat_joints(self): """Tests scanning over bodies with joints of different types.""" m = mujoco.MjModel.from_xml_string(self._MULTI_DOF_XML) - m = mjx.device_put(m) + m = mjx.put_model(m) # we will test two functions: # 1) j_fn receives jnt_types as a jp array @@ -105,7 +105,7 @@ class ScanTest(absltest.TestCase): def test_body_tree(self): """Tests tree scanning over bodies with different joint counts.""" m = mujoco.MjModel.from_xml_string(self._MULTI_DOF_XML) - m = mjx.device_put(m) + m = mjx.put_model(m) # we will test two functions: # 1) j_fn receives jnt_pos which is a jp array @@ -196,7 +196,7 @@ class ScanTest(absltest.TestCase): def test_scan_actuators(self): """Tests scanning over actuators.""" m = mujoco.MjModel.from_xml_string(self._MULTI_ACT_XML) - m = mjx.device_put(m) + m = mjx.put_model(m) fn = lambda *args: args args = (