diff --git a/doc/changelog.rst b/doc/changelog.rst index ab19efef..67a6748e 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -2,19 +2,6 @@ Changelog ========= -Upcoming version (not yet released) ------------------------------------ - -MJX -^^^ -- Promote ``ten_length`` to the public MJX API. Add Warp support for ``mjx.tendon``. - -.. admonition:: Breaking API changes - :class: attention - - - ``ten_length`` was moved from ``mjx.Data._impl.ten_length`` to a public field ``mjx.Data.ten_length``. - - Version 3.3.5 (August 8, 2025) ----------------------------------- diff --git a/mjx/mujoco/mjx/_src/constraint.py b/mjx/mujoco/mjx/_src/constraint.py index 042d2b72..410783f6 100644 --- a/mjx/mujoco/mjx/_src/constraint.py +++ b/mjx/mujoco/mjx/_src/constraint.py @@ -322,8 +322,8 @@ def _efc_equality_tendon(m: Model, d: Data) -> Optional[_Efc]: inv1, inv2 = m.tendon_invweight0[obj1id], m.tendon_invweight0[obj2id] jac1, jac2 = d._impl.ten_J[obj1id], d._impl.ten_J[obj2id] - pos1 = d.ten_length[obj1id] - m.tendon_length0[obj1id] - pos2 = d.ten_length[obj2id] - m.tendon_length0[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( @@ -436,7 +436,7 @@ 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._impl.ten_length, d._impl.ten_J, m.tendon_range, m.tendon_margin, diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index d42333d0..c3600c2c 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -576,7 +576,6 @@ def _make_data_public_fields(m: types.Model) -> Dict[str, Any]: 'qfrc_constraint': (m.nv, float_), 'qfrc_inverse': (m.nv, float_), 'cvel': (m.nbody, 6, float_), - 'ten_length': (m.ntendon, float_), } zero_fields = { k: np.zeros(v[:-1], dtype=v[-1]) for k, v in zero_fields.items() @@ -638,6 +637,7 @@ def _make_data_jax( '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_), @@ -742,6 +742,7 @@ def _make_data_c( '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), diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 8c0954b3..dba9c429 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -478,7 +478,7 @@ 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_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))) diff --git a/mjx/mujoco/mjx/_src/passive.py b/mjx/mujoco/mjx/_src/passive.py index 156ef4e9..f7acdcd2 100644 --- a/mjx/mujoco/mjx/_src/passive.py +++ b/mjx/mujoco/mjx/_src/passive.py @@ -75,7 +75,7 @@ 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._impl.ten_velocity diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py index 7d7fc6ed..31f7f030 100644 --- a/mjx/mujoco/mjx/_src/sensor.py +++ b/mjx/mujoco/mjx/_src/sensor.py @@ -22,6 +22,7 @@ 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 Impl from mujoco.mjx._src.types import Data from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import DisableBit @@ -173,7 +174,7 @@ def sensor_pos(m: Model, d: Data) -> Data: 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._impl.actuator_length[objid] elif sensor_type == SensorType.BALLQUAT: diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py index 1303f357..8f19f6f1 100644 --- a/mjx/mujoco/mjx/_src/smooth.py +++ b/mjx/mujoco/mjx/_src/smooth.py @@ -26,7 +26,6 @@ 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 Impl from mujoco.mjx._src.types import JointType from mujoco.mjx._src.types import Model from mujoco.mjx._src.types import ModelJAX @@ -34,16 +33,11 @@ from mujoco.mjx._src.types import ObjType from mujoco.mjx._src.types import TrnType from mujoco.mjx._src.types import WrapType # pylint: enable=g-importing-member -import mujoco.mjx.warp as mjxw import numpy as np def kinematics(m: Model, d: Data) -> Data: """Converts position/velocity from generalized coordinates to maximal.""" - if m.impl == Impl.WARP and d.impl == Impl.WARP and mjxw.WARP_INSTALLED: - 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 @@ -850,10 +844,6 @@ def rne_postconstraint(m: Model, d: Data) -> Data: def tendon(m: Model, d: Data) -> Data: """Computes tendon lengths and moments.""" - if m.impl == Impl.WARP and d.impl == Impl.WARP and mjxw.WARP_INSTALLED: - from mujoco.mjx.warp import smooth as mjxw_smooth # pylint: disable=g-import-not-at-top # pytype: disable=import-error - return mjxw_smooth.tendon(m, d) - if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): raise ValueError('tendon requires JAX backend implementation.') @@ -1101,7 +1091,7 @@ 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) + 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) @@ -1171,7 +1161,7 @@ def tendon(m: Model, d: Data) -> Data: ).reshape((m.nwrap, 2)) return d.tree_replace({ - 'ten_length': ten_length, + '_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), @@ -1273,7 +1263,7 @@ 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] + 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)}') diff --git a/mjx/mujoco/mjx/_src/smooth_test.py b/mjx/mujoco/mjx/_src/smooth_test.py index 4f9621a4..061087f9 100644 --- a/mjx/mujoco/mjx/_src/smooth_test.py +++ b/mjx/mujoco/mjx/_src/smooth_test.py @@ -122,7 +122,7 @@ class SmoothTest(absltest.TestCase): # tendon dx = jax.jit(mjx.tendon)(mx, mjx.put_data(m, d)) _assert_attr_eq(d, dx._impl, 'ten_J') - _assert_attr_eq(d, dx, 'ten_length') + _assert_attr_eq(d, dx._impl, 'ten_length') # transmission dx = jax.jit(mjx.transmission)(mx, dx) _assert_attr_eq(d, dx._impl, 'actuator_length') @@ -394,7 +394,7 @@ 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_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') diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 0425477f..a71f934c 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -998,6 +998,7 @@ class DataC(PyTreeNode): 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 @@ -1069,6 +1070,7 @@ class DataJAX(PyTreeNode): 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 @@ -1131,7 +1133,6 @@ class Data(PyTreeNode): ximat: jax.Array xanchor: jax.Array xaxis: jax.Array - ten_length: jax.Array geom_xpos: jax.Array geom_xmat: jax.Array site_xpos: jax.Array diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index c1ec298e..1daac020 100644 --- a/mjx/mujoco/mjx/warp/forward.py +++ b/mjx/mujoco/mjx/warp/forward.py @@ -1113,7 +1113,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'ten_Jdot': d._impl.ten_Jdot.shape, 'ten_actfrc': d._impl.ten_actfrc.shape, 'ten_bias_coef': d._impl.ten_bias_coef.shape, - 'ten_length': d.ten_length.shape, + 'ten_length': d._impl.ten_length.shape, 'ten_velocity': d._impl.ten_velocity.shape, 'ten_wrapadr': d._impl.ten_wrapadr.shape, 'ten_wrapnum': d._impl.ten_wrapnum.shape, @@ -1777,7 +1777,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.ten_Jdot, d._impl.ten_actfrc, d._impl.ten_bias_coef, - d.ten_length, + d._impl.ten_length, d._impl.ten_velocity, d._impl.ten_wrapadr, d._impl.ten_wrapnum, @@ -1956,7 +1956,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): '_impl.ten_Jdot': out[99], '_impl.ten_actfrc': out[100], '_impl.ten_bias_coef': out[101], - 'ten_length': out[102], + '_impl.ten_length': out[102], '_impl.ten_velocity': out[103], '_impl.ten_wrapadr': out[104], '_impl.ten_wrapnum': out[105], @@ -3176,7 +3176,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'ten_Jdot': d._impl.ten_Jdot.shape, 'ten_actfrc': d._impl.ten_actfrc.shape, 'ten_bias_coef': d._impl.ten_bias_coef.shape, - 'ten_length': d.ten_length.shape, + 'ten_length': d._impl.ten_length.shape, 'ten_velocity': d._impl.ten_velocity.shape, 'ten_wrapadr': d._impl.ten_wrapadr.shape, 'ten_wrapnum': d._impl.ten_wrapnum.shape, @@ -3866,7 +3866,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.ten_Jdot, d._impl.ten_actfrc, d._impl.ten_bias_coef, - d.ten_length, + d._impl.ten_length, d._impl.ten_velocity, d._impl.ten_wrapadr, d._impl.ten_wrapnum, @@ -4057,7 +4057,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): '_impl.ten_Jdot': out[111], '_impl.ten_actfrc': out[112], '_impl.ten_bias_coef': out[113], - 'ten_length': out[114], + '_impl.ten_length': out[114], '_impl.ten_velocity': out[115], '_impl.ten_wrapadr': out[116], '_impl.ten_wrapnum': out[117], diff --git a/mjx/mujoco/mjx/warp/forward_test.py b/mjx/mujoco/mjx/warp/forward_test.py index ff890658..5fae3d22 100644 --- a/mjx/mujoco/mjx/warp/forward_test.py +++ b/mjx/mujoco/mjx/warp/forward_test.py @@ -142,7 +142,7 @@ class ForwardTest(parameterized.TestCase): if m.ncam: tu.assert_attr_eq(dx, d, 'cam_xpos') tu.assert_eq(dx.cam_xmat, d.cam_xmat.reshape((-1, 3, 3)), 'cam_xmat') - tu.assert_attr_eq(dx, d, 'ten_length') + tu.assert_attr_eq(dx._impl, d, 'ten_length') tu.assert_attr_eq(dx._impl, d, 'ten_J') tu.assert_attr_eq(dx._impl, d, 'ten_wrapadr') tu.assert_attr_eq(dx._impl, d, 'ten_wrapnum') diff --git a/mjx/mujoco/mjx/warp/smooth.py b/mjx/mujoco/mjx/warp/smooth.py index 3b881891..cd0e4138 100644 --- a/mjx/mujoco/mjx/warp/smooth.py +++ b/mjx/mujoco/mjx/warp/smooth.py @@ -285,213 +285,3 @@ def kinematics(m: types.Model, d: types.Data): def kinematics_vmap(unused_axis_size, is_batched, m, d): d = kinematics(m, d) return d, is_batched[1] - - -_m = mjwarp.Model( - **{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init} -) -_d = mjwarp.Data( - **{f.name: None for f in dataclasses.fields(mjwarp.Data) if f.init} -) -_o = mjwarp.Option( - **{f.name: None for f in dataclasses.fields(mjwarp.Option) if f.init} -) -_s = mjwarp.Statistic( - **{f.name: None for f in dataclasses.fields(mjwarp.Statistic) if f.init} -) -_c = mjwarp.Contact( - **{f.name: None for f in dataclasses.fields(mjwarp.Contact) if f.init} -) -_e = mjwarp.Constraint( - **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} -) - - -@ffi.format_args_for_warp -def _tendon_shim( - # Model - nworld: int, - body_parentid: wp.array(dtype=int), - body_rootid: wp.array(dtype=int), - dof_bodyid: wp.array(dtype=int), - geom_bodyid: wp.array(dtype=int), - geom_size: wp.array2d(dtype=wp.vec3), - jnt_dofadr: wp.array(dtype=int), - jnt_qposadr: wp.array(dtype=int), - ntendon: int, - nv: int, - site_bodyid: wp.array(dtype=int), - tendon_adr: wp.array(dtype=int), - tendon_geom_adr: wp.array(dtype=int), - tendon_jnt_adr: wp.array(dtype=int), - tendon_num: wp.array(dtype=int), - tendon_site_pair_adr: wp.array(dtype=int), - wrap_geom_adr: wp.array(dtype=int), - wrap_jnt_adr: wp.array(dtype=int), - wrap_objid: wp.array(dtype=int), - wrap_prm: wp.array(dtype=float), - wrap_pulley_scale: wp.array(dtype=float), - wrap_site_pair_adr: wp.array(dtype=int), - wrap_type: wp.array(dtype=int), - # Data - cdof: wp.array2d(dtype=wp.spatial_vector), - geom_xmat: wp.array2d(dtype=wp.mat33), - geom_xpos: wp.array2d(dtype=wp.vec3), - qpos: wp.array2d(dtype=float), - site_xpos: wp.array2d(dtype=wp.vec3), - subtree_com: wp.array2d(dtype=wp.vec3), - ten_J: wp.array3d(dtype=float), - ten_length: wp.array2d(dtype=float), - ten_wrapadr: wp.array2d(dtype=int), - ten_wrapnum: wp.array2d(dtype=int), - wrap_geom_xpos: wp.array2d(dtype=wp.spatial_vector), - wrap_obj: wp.array2d(dtype=wp.vec2i), - wrap_xpos: wp.array2d(dtype=wp.spatial_vector), -): - _m.stat = _s - _m.opt = _o - _d.efc = _e - _d.contact = _c - _m.body_parentid = body_parentid - _m.body_rootid = body_rootid - _m.dof_bodyid = dof_bodyid - _m.geom_bodyid = geom_bodyid - _m.geom_size = geom_size - _m.jnt_dofadr = jnt_dofadr - _m.jnt_qposadr = jnt_qposadr - _m.ntendon = ntendon - _m.nv = nv - _m.site_bodyid = site_bodyid - _m.tendon_adr = tendon_adr - _m.tendon_geom_adr = tendon_geom_adr - _m.tendon_jnt_adr = tendon_jnt_adr - _m.tendon_num = tendon_num - _m.tendon_site_pair_adr = tendon_site_pair_adr - _m.wrap_geom_adr = wrap_geom_adr - _m.wrap_jnt_adr = wrap_jnt_adr - _m.wrap_objid = wrap_objid - _m.wrap_prm = wrap_prm - _m.wrap_pulley_scale = wrap_pulley_scale - _m.wrap_site_pair_adr = wrap_site_pair_adr - _m.wrap_type = wrap_type - _d.cdof = cdof - _d.geom_xmat = geom_xmat - _d.geom_xpos = geom_xpos - _d.qpos = qpos - _d.site_xpos = site_xpos - _d.subtree_com = subtree_com - _d.ten_J = ten_J - _d.ten_length = ten_length - _d.ten_wrapadr = ten_wrapadr - _d.ten_wrapnum = ten_wrapnum - _d.wrap_geom_xpos = wrap_geom_xpos - _d.wrap_obj = wrap_obj - _d.wrap_xpos = wrap_xpos - _d.nworld = nworld - mjwarp.tendon(_m, _d) - - -def _tendon_jax_impl(m: types.Model, d: types.Data): - output_dims = { - 'cdof': d._impl.cdof.shape, - 'geom_xmat': d.geom_xmat.shape, - 'geom_xpos': d.geom_xpos.shape, - 'qpos': d.qpos.shape, - 'site_xpos': d.site_xpos.shape, - 'subtree_com': d.subtree_com.shape, - 'ten_J': d._impl.ten_J.shape, - 'ten_length': d.ten_length.shape, - 'ten_wrapadr': d._impl.ten_wrapadr.shape, - 'ten_wrapnum': d._impl.ten_wrapnum.shape, - 'wrap_geom_xpos': d._impl.wrap_geom_xpos.shape, - 'wrap_obj': d._impl.wrap_obj.shape, - 'wrap_xpos': d._impl.wrap_xpos.shape, - } - jf = ffi.jax_callable_variadic_tuple( - _tendon_shim, - num_outputs=13, - output_dims=output_dims, - vmap_method=None, - in_out_argnames={ - 'cdof', - 'geom_xmat', - 'geom_xpos', - 'qpos', - 'site_xpos', - 'subtree_com', - 'ten_J', - 'ten_length', - 'ten_wrapadr', - 'ten_wrapnum', - 'wrap_geom_xpos', - 'wrap_obj', - 'wrap_xpos', - }, - ) - out = jf( - d.qpos.shape[0], - m.body_parentid, - m.body_rootid, - m.dof_bodyid, - m.geom_bodyid, - m.geom_size, - m.jnt_dofadr, - m.jnt_qposadr, - m.ntendon, - m.nv, - m.site_bodyid, - m.tendon_adr, - m._impl.tendon_geom_adr, - m._impl.tendon_jnt_adr, - m.tendon_num, - m._impl.tendon_site_pair_adr, - m._impl.wrap_geom_adr, - m._impl.wrap_jnt_adr, - m.wrap_objid, - m.wrap_prm, - m._impl.wrap_pulley_scale, - m._impl.wrap_site_pair_adr, - m.wrap_type, - d._impl.cdof, - d.geom_xmat, - d.geom_xpos, - d.qpos, - d.site_xpos, - d.subtree_com, - d._impl.ten_J, - d.ten_length, - d._impl.ten_wrapadr, - d._impl.ten_wrapnum, - d._impl.wrap_geom_xpos, - d._impl.wrap_obj, - d._impl.wrap_xpos, - ) - d = d.tree_replace({ - '_impl.cdof': out[0], - 'geom_xmat': out[1], - 'geom_xpos': out[2], - 'qpos': out[3], - 'site_xpos': out[4], - 'subtree_com': out[5], - '_impl.ten_J': out[6], - 'ten_length': out[7], - '_impl.ten_wrapadr': out[8], - '_impl.ten_wrapnum': out[9], - '_impl.wrap_geom_xpos': out[10], - '_impl.wrap_obj': out[11], - '_impl.wrap_xpos': out[12], - }) - return d - - -@jax.custom_batching.custom_vmap -@ffi.marshal_jax_warp_callable -def tendon(m: types.Model, d: types.Data): - return _tendon_jax_impl(m, d) - - -@tendon.def_vmap -@ffi.marshal_custom_vmap -def tendon_vmap(unused_axis_size, is_batched, m, d): - d = tendon(m, d) - return d, is_batched[1] diff --git a/mjx/mujoco/mjx/warp/smooth_test.py b/mjx/mujoco/mjx/warp/smooth_test.py index c83574fe..6d9bc349 100644 --- a/mjx/mujoco/mjx/warp/smooth_test.py +++ b/mjx/mujoco/mjx/warp/smooth_test.py @@ -19,7 +19,6 @@ import os import tempfile from absl.testing import absltest -from absl.testing import parameterized import jax from jax import numpy as jp import mujoco @@ -39,7 +38,7 @@ except ImportError: _FORCE_TEST = os.environ.get('MJX_WARP_FORCE_TEST', '0') == '1' -class SmoothTest(parameterized.TestCase): +class SmoothTest(absltest.TestCase): def setUp(self): super().setUp() @@ -238,5 +237,6 @@ class SmoothTest(parameterized.TestCase): tu.assert_attr_eq(d, dx, 'site_xpos') tu.assert_eq(d.site_xmat.reshape((-1, 3, 3)), dx.site_xmat, 'site_xmat') + if __name__ == '__main__': absltest.main() diff --git a/mjx/mujoco/mjx/warp/test_util.py b/mjx/mujoco/mjx/warp/test_util.py index 7521e094..b5c25f3c 100644 --- a/mjx/mujoco/mjx/warp/test_util.py +++ b/mjx/mujoco/mjx/warp/test_util.py @@ -152,8 +152,7 @@ def _mjx_efc(dx, worldid: int): return 0, empty, empty, np.zeros((0, dx.qvel.shape[0])), empty, empty efc_pos = select(dx._impl.efc__pos[:nefc]) efc_type = select(dx._impl.efc__type[:nefc]) - efc_d = select(dx._impl.efc__D[:nefc]) - keys_sorted = np.lexsort((-efc_pos, efc_type, efc_d)) + keys_sorted = np.lexsort((-efc_pos, efc_type)) keys = keys[keys_sorted] nefc = len(keys) @@ -179,7 +178,7 @@ def _mj_efc(d): else: efc_j = d.efc_J.reshape((-1, d.qvel.shape[0])) - keys = np.lexsort((-d.efc_pos, d.efc_type, d.efc_D)) + keys = np.lexsort((-d.efc_pos, d.efc_type)) type_ = d.efc_type[keys] pos = d.efc_pos[keys] efc_j = efc_j[keys] diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index 2d70fb60..e86b691d 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -361,6 +361,7 @@ class DataWarp(PyTreeNode): ten_Jdot: jax.Array ten_actfrc: jax.Array ten_bias_coef: jax.Array + ten_length: jax.Array ten_velocity: jax.Array ten_wrapadr: jax.Array ten_wrapnum: jax.Array