diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml
index b2daab71..ac1a88e5 100644
--- a/.github/workflows/build.yml
+++ b/.github/workflows/build.yml
@@ -281,11 +281,12 @@ jobs:
working-directory: mjx
run:
source ${{ matrix.tmpdir }}/venv/bin/activate &&
- pip install --require-hashes -r requirements.txt &&
+ pip install --require-hashes -r requirements.txt
pip install --no-index dist/mujoco_mjx-*.whl
- name: Test MJX
if: ${{ runner.os != 'Windows' }}
shell: bash
+ working-directory: mjx
run:
source ${{ matrix.tmpdir }}/venv/bin/activate &&
pytest -n auto -v -k 'not IntegrationTest' --pyargs mujoco.mjx
diff --git a/doc/changelog.rst b/doc/changelog.rst
index 2f492926..4c82b923 100644
--- a/doc/changelog.rst
+++ b/doc/changelog.rst
@@ -47,6 +47,9 @@ Version 3.3.4 (July 8, 2025)
3. In the mjSpec C API, directly setting an element's name using :ref:`mjs_setString` has been replaced with a new
function :ref:`mjs_setName` which allows checking for naming collisions at set-time rather than compile-time, for
earlier catching of errors. Relatedly, the ``name`` attribute has been removed from all mjs elements.
+ 4. For MJX, the ``mjx.Option`` dataclass now has private and public fields similar to ``mjx.Model`` and
+ ``mjx.Data``. Some fields are no longer publicly available due to differences in the
+ underlying implementations of this data structure.
General
^^^^^^^
@@ -66,6 +69,13 @@ Documentation
8. Added missing item documentation and clarified the nature of breaking changes in the 3.3.3 changelog.
See items 3 and 4 below.
+MJX
+^^^
+- Add Warp as a backend implementation for MJX. The implementation can be specified via
+ ``mjx.put_model(m, impl='warp')`` and ``mjx.make_data(m, impl='warp')``. The warp implementation requires
+ a CUDA device, Python 3.12, and `warp-lang` to be installed. This feature is available in "beta" and
+ some bugs are expected.
+
Version 3.3.3 (June 10, 2025)
-----------------------------
diff --git a/mjx/cuda_requirements.txt b/mjx/cuda_requirements.txt
index 3791dff1..4ea1c845 100644
--- a/mjx/cuda_requirements.txt
+++ b/mjx/cuda_requirements.txt
@@ -1,15 +1,18 @@
-jax[cuda12_local]==0.4.34; python_version >= '3.10' \
- --hash=sha256:b957ca1fc91f7343f91a186af9f19c7f342c946f95a8c11c7f1e5cdfe2e58d9e
-jax[cuda12_local]==0.4.30; python_version == '3.9' \
- --hash=sha256:289b30ae03b52f7f4baf6ef082a9f4e3e29c1080e22d13512c5ecf02d5f1a55b
-jax-cuda12-plugin==0.4.34; python_version >= '3.10' \
- --hash=sha256:d035ea72bd9b8a65a6ea621bca1affdd33127fa3a52e7bded7692670d360adab \
- --hash=sha256:db988b7ba5063483a936ddbf162f04d1b4412e0d64340f11788c7bbc877e8a43 \
- --hash=sha256:e23721d1654b311b47cd6b35768520284bd036f8c7e6b11600143258b4a0409a \
- --hash=sha256:b2099a4407225122ff76f6dcdc8dbdae47e6f29343bdfd21460ad337dc34a209
+jax-cuda12-plugin==0.5.3; python_version >= '3.10' \
+ --hash=sha256:6171aed2f4b3bdd5fc13782de1072c6a634fce13731b75d0cb0a6ab8f4e6e650 \
+ --hash=sha256:ba2555967f9b6c381c8b4ef9fb03d05bc55ec25ecfee5cfe45c5ace34f7d4152 \
+ --hash=sha256:298d2d768f1029b74a0b1d01270e549349d2c37dc07658796542cda967eb7bd3 \
+ --hash=sha256:aaa704a5ef547595d022db1c1e4878a0677116412a9360c115d67ff4b64e1596 \
+ --hash=sha256:c2517a7c2186f8708894696e26cf96ebd60b7879ceca398b2c46abb28d2c96c8 \
+ --hash=sha256:2030cf1208ce4ea70ee56cac61ddd239f9798695fc39bb7739c50a25d6e9da44 \
+ --hash=sha256:21fec1b56c98783ea0569b747a56751f1f9ff2187b48acc11c700d3bfc5e1a31 \
+ --hash=sha256:1862595b2b6d815679d11e0e889e523185ee54a46d46e022689f70fc4554dd91 \
+ --hash=sha256:6d43677f22f3be9544a205216cd6dac591335b1d9bbbed018cd17dbb1f3f4def \
+ --hash=sha256:5bb9ea0e68d72d44e57e4cb6a58a1a729fe3fe32e964f71e398d8a25c2103b19
jax-cuda12-plugin==0.4.30; python_version == '3.9' \
--hash=sha256:d8d196241b9253ecb1144a4409b5deacbb9771624f097b2bbf025da3c7d8f4f8
-jax-cuda12-pjrt==0.4.34; python_version >= '3.10' \
- --hash=sha256:0c7cc98f962cc7fc8e0a5ea6331b42a0cee516f202f1c3019f6aa5cd9530cca0
+jax-cuda12-pjrt==0.5.3; python_version >= '3.10' \
+ --hash=sha256:04ee111eaf5fc2692978ad4a5c84d5925e42eb05c1701849ba3a53f6515400cc \
+ --hash=sha256:c5378306568ba0c81b230a779dd3194c9dd10339ab6360ae80928108d37e7f75
jax-cuda12-pjrt==0.4.30; python_version == '3.9' \
--hash=sha256:895d0198ad99638fcaf976c47592e2a543eef79ea15fabd24a402d055390c328
diff --git a/mjx/mujoco/mjx/_src/collision_driver.py b/mjx/mujoco/mjx/_src/collision_driver.py
index 2f45bca8..158be7a2 100644
--- a/mjx/mujoco/mjx/_src/collision_driver.py
+++ b/mjx/mujoco/mjx/_src/collision_driver.py
@@ -76,6 +76,7 @@ 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
+from mujoco.mjx._src.types import OptionJAX
# pylint: enable=g-importing-member
import numpy as np
@@ -342,6 +343,16 @@ def _numeric(m: Union[Model, mujoco.MjModel], name: str) -> int:
def make_condim(m: Union[Model, mujoco.MjModel]) -> np.ndarray:
"""Returns the dims of the contacts for a Model."""
+ if isinstance(m, mujoco.MjModel):
+ sdf_initpoints = m.opt.sdf_initpoints
+ elif isinstance(m.opt._impl, OptionJAX):
+ sdf_initpoints = m.opt._impl.sdf_initpoints
+ else:
+ raise ValueError(
+ 'make_condim requires mujoco.MjModel or mjx.Model with JAX backend'
+ ' implementation.'
+ )
+
if m.opt.disableflags & DisableBit.CONTACT:
return np.empty(0, dtype=int)
@@ -364,7 +375,7 @@ def make_condim(m: Union[Model, mujoco.MjModel]) -> np.ndarray:
condim_counts = {}
for k, v in group_counts.items():
if k.types[1] == mujoco.mjtGeom.mjGEOM_SDF:
- ncon = m.opt.sdf_initpoints
+ ncon = sdf_initpoints
else:
func = _COLLISION_FUNC[k.types]
ncon = func.ncon # pytype: disable=attribute-error
diff --git a/mjx/mujoco/mjx/_src/constraint.py b/mjx/mujoco/mjx/_src/constraint.py
index 9b41daa7..410783f6 100644
--- a/mjx/mujoco/mjx/_src/constraint.py
+++ b/mjx/mujoco/mjx/_src/constraint.py
@@ -35,6 +35,7 @@ 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
+from mujoco.mjx._src.types import OptionJAX
# pylint: enable=g-importing-member
import numpy as np
@@ -496,7 +497,11 @@ def _efc_contact_frictionless(m: Model, d: Data) -> Optional[_Efc]:
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):
+ if (
+ not isinstance(m._impl, ModelJAX)
+ or not isinstance(d._impl, DataJAX)
+ or not isinstance(m.opt._impl, OptionJAX)
+ ):
raise ValueError(
'_efc_contact_pyramidal requires JAX backend implementation.'
)
@@ -545,7 +550,11 @@ def _efc_contact_pyramidal(m: Model, d: Data, condim: int) -> Optional[_Efc]:
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):
+ if (
+ not isinstance(m._impl, ModelJAX)
+ or not isinstance(d._impl, DataJAX)
+ or not isinstance(m.opt._impl, OptionJAX)
+ ):
raise ValueError(
'_efc_contact_elliptic requires JAX backend implementation.'
)
diff --git a/mjx/mujoco/mjx/_src/dataclasses.py b/mjx/mujoco/mjx/_src/dataclasses.py
index b8bdb758..e82d6047 100644
--- a/mjx/mujoco/mjx/_src/dataclasses.py
+++ b/mjx/mujoco/mjx/_src/dataclasses.py
@@ -62,11 +62,19 @@ def dataclass(clz: _T, register_as_pytree: bool) -> _T:
def iterate_clz_with_keys(x):
def to_meta(field, obj):
val = getattr(obj, field.name)
- # numpy arrays are not hashable so return raw bytes instead
if isinstance(val, np.ndarray):
+ # numpy arrays are not hashable so return raw bytes instead
return (val.tobytes(), val.dtype, val.shape)
- else:
- return val
+ if typing.get_origin(field.type) == tuple:
+ # variadic tuples of numpy arrays
+ type_args = typing.get_args(field.type)
+ if (
+ len(type_args) == 2
+ and type_args[0] == np.ndarray
+ and type_args[1] == ...
+ ):
+ return tuple((v.tobytes(), v.dtype, v.shape) for v in val)
+ return val
def to_data(field, obj):
return (jax.tree_util.GetAttrKey(field.name), getattr(obj, field.name))
@@ -81,8 +89,20 @@ def dataclass(clz: _T, register_as_pytree: bool) -> _T:
if field.type is np.ndarray:
arr = np.frombuffer(meta[0], dtype=meta[1]).reshape(meta[2])
return (field.name, arr)
- else:
- return (field.name, meta)
+ if typing.get_origin(field.type) == tuple:
+ type_args = typing.get_args(field.type)
+ if (
+ len(type_args) == 2
+ and type_args[0] == np.ndarray
+ and type_args[1] == ...
+ ):
+ return (
+ field.name,
+ tuple(
+ np.frombuffer(m[0], dtype=m[1]).reshape(m[2]) for m in meta
+ ),
+ )
+ return (field.name, meta)
from_data = lambda field, meta: (field.name, meta)
diff --git a/mjx/mujoco/mjx/_src/dataclasses_test.py b/mjx/mujoco/mjx/_src/dataclasses_test.py
new file mode 100644
index 00000000..b5727941
--- /dev/null
+++ b/mjx/mujoco/mjx/_src/dataclasses_test.py
@@ -0,0 +1,69 @@
+# Copyright 2025 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 custom PyTreeNode object."""
+
+from absl.testing import absltest
+import jax
+from jax import numpy as jp
+from mujoco.mjx._src import dataclasses
+import numpy as np
+
+
+class Obj(dataclasses.PyTreeNode):
+ a: int
+ b: np.ndarray
+ c: tuple[int, ...]
+ d: tuple[np.ndarray, ...]
+ e: jax.Array
+ f: tuple[jax.Array, ...]
+
+
+class DataclassesTest(absltest.TestCase):
+
+ def test_pytree_structure(self):
+ obj = Obj(
+ a=1,
+ b=np.array([1, 2, 3]),
+ c=(4, 5, 6),
+ d=(np.array([7, 8]), np.array([9, 10])),
+ e=jax.numpy.array([11, 12]),
+ f=(jax.numpy.array([13, 14]), jax.numpy.array([15, 16])),
+ )
+
+ data, meta = jax.tree_util.tree_flatten_with_path(obj)
+
+ # data fields
+ self.assertLen(data, 3)
+ self.assertEqual(data[0][0][0].name, 'e')
+ np.testing.assert_array_equal(data[0][1], jp.array([11, 12]))
+ self.assertEqual(data[1][0][0].name, 'f')
+ np.testing.assert_array_equal(data[1][1], jp.array([13, 14]))
+ self.assertEqual(data[2][0][0].name, 'f')
+ np.testing.assert_array_equal(data[2][1], jp.array([15, 16]))
+
+ # meta fields
+ unflattened_meta = meta.unflatten([x[1] for x in data])
+ self.assertEqual(unflattened_meta.a, 1)
+ np.testing.assert_array_equal(unflattened_meta.b, np.array([1, 2, 3]))
+ self.assertEqual(unflattened_meta.c, (4, 5, 6))
+ np.testing.assert_array_equal(unflattened_meta.d[0], np.array([7, 8]))
+ np.testing.assert_array_equal(unflattened_meta.d[1], np.array([9, 10]))
+
+ # ensure hashable meta
+ hash(meta)
+
+
+if __name__ == '__main__':
+ absltest.main()
diff --git a/mjx/mujoco/mjx/_src/derivative.py b/mjx/mujoco/mjx/_src/derivative.py
index 4e3d6d68..cd56ad06 100644
--- a/mjx/mujoco/mjx/_src/derivative.py
+++ b/mjx/mujoco/mjx/_src/derivative.py
@@ -21,14 +21,24 @@ from jax import numpy as jp
# 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 Model
+from mujoco.mjx._src.types import ModelJAX
+from mujoco.mjx._src.types import OptionJAX
+# pylint: enable=g-importing-member
def deriv_smooth_vel(m: Model, d: Data) -> Optional[jax.Array]:
"""Analytical derivative of smooth forces w.r.t. velocities."""
+ if (
+ not isinstance(m._impl, ModelJAX)
+ or not isinstance(d._impl, DataJAX)
+ or not isinstance(m.opt._impl, OptionJAX)
+ ):
+ raise ValueError('deriv_smooth_vel requires JAX MJX implementation.')
qderiv = None
@@ -53,7 +63,7 @@ def deriv_smooth_vel(m: Model, d: Data) -> Optional[jax.Array]:
if m.ntendon:
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: # pytype: disable=attribute-error
+ if m.opt._impl.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 c760b0ff..bb01727d 100644
--- a/mjx/mujoco/mjx/_src/forward.py
+++ b/mjx/mujoco/mjx/_src/forward.py
@@ -37,12 +37,14 @@ 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 Impl
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 mujoco.mjx.warp as mjxw
import numpy as np
# RK4 tableau
@@ -423,6 +425,10 @@ def implicit(m: Model, d: Data) -> Data:
@named_scope
def forward(m: Model, d: Data) -> Data:
"""Forward dynamics."""
+ if m.impl == Impl.WARP and d.impl == Impl.WARP and mjxw.WARP_INSTALLED:
+ from mujoco.mjx.warp import forward as mjxw_forward # pylint: disable=g-import-not-at-top # pytype: disable=import-error
+ return mjxw_forward.forward(m, d)
+
if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX):
raise ValueError('forward requires JAX backend implementation.')
@@ -446,6 +452,10 @@ def forward(m: Model, d: Data) -> Data:
@named_scope
def step(m: Model, d: Data) -> Data:
"""Advance simulation."""
+ if m.impl == Impl.WARP and d.impl == Impl.WARP and mjxw.WARP_INSTALLED:
+ from mujoco.mjx.warp import forward as mjxw_forward # pylint: disable=g-import-not-at-top # pytype: disable=import-error
+ return mjxw_forward.step(m, d)
+
d = forward(m, d)
if m.opt.integrator == IntegratorType.EULER:
diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py
index c690011c..fe2dabd6 100644
--- a/mjx/mujoco/mjx/_src/io.py
+++ b/mjx/mujoco/mjx/_src/io.py
@@ -22,23 +22,31 @@ import warnings
import jax
from jax import numpy as jp
+from jax.extend import backend
import mujoco
from mujoco.mjx._src import collision_driver
from mujoco.mjx._src import constraint
from mujoco.mjx._src import mesh
from mujoco.mjx._src import support
from mujoco.mjx._src import types
+import mujoco.mjx.warp as mjxw
+# pylint: disable=g-importing-member
+from mujoco.mjx.warp import mjwp_types
+from mujoco.mjx.warp import mujoco_warp as mjwp
+from mujoco.mjx.warp import warp as wp
+# pylint: enable=g-importing-member
import numpy as np
import scipy
+def has_cuda_gpu_device() -> bool:
+ return 'cuda' in backend.backends()
+
+
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").')
+ if not has_cuda_gpu_device():
return False
- return device in cuda_devices
+ return device in jax.devices('cuda')
def _resolve_impl(
@@ -46,24 +54,19 @@ def _resolve_impl(
) -> types.Impl:
"""Pick a default implementation 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:
+ # TODO(btaba): Remove flag once Warp is ready for GPU default.
+ mjx_gpu_default_warp = (
+ os.environ.get('MJX_GPU_DEFAULT_WARP', 'f').lower() == 'true'
+ )
+ if mjx_gpu_default_warp and mjxw.WARP_INSTALLED:
logging.debug('Picking default implementation: Warp.')
return types.Impl.WARP
- logging.info('MJX Warp is disabled via MJX_WARP_ENABLED=false.')
if device.platform in ('gpu', 'tpu'):
logging.debug('Picking default implementation: JAX.')
return types.Impl.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 implementation: C.')
- return types.Impl.C
return types.Impl.JAX
raise ValueError(f'Unsupported device: {device}')
@@ -115,11 +118,9 @@ def _check_impl_device_compatibility(
'Warp 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 implementation is disabled via MJX_WARP_ENABLED=false.'
+ if not mjxw.WARP_INSTALLED:
+ raise RuntimeError(
+ 'Warp is not installed. Cannot use Warp implementation of MJX.'
)
is_cpu_device = device.platform == 'cpu'
@@ -165,10 +166,53 @@ def _strip_weak_type(tree):
return jax.tree_util.tree_map(f, tree)
+def _wp_to_np_type(wp_field: Any, name: str = '') -> Any:
+ """Converts a warp type to an MJX compatible numpy type."""
+ if hasattr(wp_field, '_is_batched'):
+ wp_field.strides = wp_field.strides[1:]
+ wp_field.shape = wp_field.shape[1:]
+ # warp scalars
+ wp_dtype = type(wp_field)
+ if wp_dtype in wp.types.warp_type_to_np_dtype:
+ return wp.types.warp_type_to_np_dtype[wp_dtype](wp_field)
+
+ # warp arrays
+ if isinstance(wp_field, wp.types.array):
+ return wp_field.numpy()
+
+ # static
+ static_types = (bool, int, float, np.bool, np.int32, np.int64,
+ np.float32, np.float64) # fmt: skip
+ is_static = lambda x: isinstance(x, static_types)
+ if is_static(wp_field):
+ return wp_field
+
+ # tuples
+ if isinstance(wp_field, tuple) and len(wp_field) == 0:
+ return ()
+ if isinstance(wp_field, tuple) and isinstance(wp_field[0], wp.types.array):
+ return tuple(f.numpy() for f in wp_field)
+ if isinstance(wp_field, tuple) and isinstance(
+ wp_field[0], mjwp_types.TileSet
+ ):
+ return tuple(
+ mjxw.types.TileSet(wp_field[i].adr.numpy(), wp_field[i].size)
+ for i in range(len(wp_field))
+ )
+ if isinstance(wp_field, mjwp_types.BlockDim):
+ return mjxw.types.BlockDim(**wp_field.__dict__)
+ if isinstance(wp_field, tuple) and is_static(wp_field[0]):
+ return wp_field
+
+ raise NotImplementedError(
+ f'Field {name} has unsupported type {type(wp_field)}.'
+ )
+
+
def _put_option(
o: mujoco.MjOption,
impl: types.Impl,
- impl_fields: Optional[dict[str, Any]] = None,
+ impl_fields: Optional[Dict[str, Any]] = None,
) -> types.Option:
"""Returns mjx.Option given mujoco.MjOption."""
if o.integrator not in set(types.IntegratorType):
@@ -187,32 +231,56 @@ def _put_option(
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 = {
+ f.name: getattr(o, f.name, None)
+ for f in types.Option.fields()
+ if f.name != '_impl'
+ }
fields['integrator'] = types.IntegratorType(o.integrator)
fields['cone'] = types.ConeType(o.cone)
- fields['jacobian'] = types.JacobianType(o.jacobian)
fields['solver'] = types.SolverType(o.solver)
fields['disableflags'] = types.DisableBit(o.disableflags)
fields['enableflags'] = types.EnableBit(o.enableflags)
+ fields['jacobian'] = types.JacobianType(o.jacobian)
+
+ option_obj = {
+ types.Impl.C: types.OptionC,
+ types.Impl.JAX: types.OptionJAX,
+ types.Impl.WARP: mjxw.types.OptionWarp,
+ }[impl]
+ private_fields = {
+ f.name: getattr(o, f.name, None) for f in option_obj.fields()
+ }
+ impl_fields = impl_fields or {}
+ impl_fields = {**private_fields, **impl_fields}
if impl == types.Impl.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 {}))
+ impl_fields['has_fluid_params'] = has_fluid_params
+ return types.Option(**fields, _impl=types.OptionJAX(**impl_fields))
if impl == types.Impl.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 {}))
+ return types.Option(**fields, _impl=types.OptionC(**impl_fields))
+
+ if impl == types.Impl.WARP:
+ impl_fields = {k: _wp_to_np_type(v) for k, v in impl_fields.items()}
+ return types.Option(**fields, _impl=mjxw.types.OptionWarp(**impl_fields))
raise NotImplementedError(f'Unsupported implementation: {impl}')
-def _put_statistic(s: mujoco.MjStatistic) -> types.Statistic:
+def _put_statistic(
+ s: mujoco.MjStatistic, impl: types.Impl
+) -> Union[types.Statistic, types.StatisticWarp]:
"""Puts mujoco.MjStatistic onto a device, resulting in mjx.Statistic."""
+ if impl == types.Impl.WARP:
+ fields = {
+ f.name: getattr(s, f.name, None) for f in types.StatisticWarp.fields()
+ }
+ return types.StatisticWarp(**fields)
return types.Statistic(
meaninertia=s.meaninertia,
meanmass=s.meanmass,
@@ -305,7 +373,7 @@ def _put_model_jax(
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, types.Impl.JAX)
- fields['stat'] = _put_statistic(m.stat)
+ fields['stat'] = _put_statistic(m.stat, types.Impl.JAX)
fields_jax = {}
fields_jax['dof_hasfrictionloss'] = fields['dof_frictionloss'] > 0
@@ -362,7 +430,7 @@ def _put_model_c(
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)
+ fields['stat'] = _put_statistic(m.stat, impl=types.Impl.C)
c_impl_keys = (
types.ModelC.__annotations__.keys() - types.Model.__annotations__.keys()
@@ -377,6 +445,51 @@ def _put_model_c(
return _strip_weak_type(model)
+def _put_model_warp(
+ m: mujoco.MjModel,
+ device: Optional[jax.Device] = None,
+) -> types.Model:
+ """Puts mujoco.MjModel onto a device, resulting in mjx.Model."""
+ if not mjxw.WARP_INSTALLED:
+ raise RuntimeError('Warp not installed.')
+
+ with wp.ScopedDevice('cpu'): # pylint: disable=undefined-variable
+ mw = mjwp.put_model(m) # pylint: disable=undefined-variable
+ mw.opt.graph_conditional = False
+
+ fields = {f.name for f in types.Model.fields() if f.name != '_impl'}
+ fields = {f: getattr(m, f) for f in fields}
+ # Grab MJW private Option fields, and assume that public MjOption fields are
+ # directly compatible with MJXW.
+ option_keys = {f.name for f in mjxw.types.OptionWarp.fields()} - {
+ f.name for f in types.Option.fields()
+ }
+ private_options = {k: getattr(mw.opt, k) for k in option_keys}
+ fields['opt'] = _put_option(m.opt, types.Impl.WARP, private_options)
+ fields['stat'] = _put_statistic(m.stat, types.Impl.WARP)
+
+ # Use MJW fields directly instead of MjModel, so that shape and dtype are
+ # always compatible with MJXW (e.g. cam_mat0/geom_aabb).
+ for k in fields:
+ if not hasattr(mw, k) or k in ('stat', 'opt'):
+ continue
+ field = _wp_to_np_type(getattr(mw, k), k)
+ fields[k] = field
+
+ impl_fields = {}
+ for k in mjxw.types.ModelWarp.__annotations__.keys():
+ field = _wp_to_np_type(getattr(mw, k), k)
+ impl_fields[k] = field
+
+ model = types.Model(
+ **fields,
+ _impl=mjxw.types.ModelWarp(**impl_fields),
+ )
+
+ model = jax.device_put(model, device=device)
+ return _strip_weak_type(model)
+
+
def put_model(
m: mujoco.MjModel,
device: Optional[jax.Device] = None,
@@ -416,7 +529,7 @@ def put_model(
elif impl == types.Impl.C:
return _put_model_c(m, device)
elif impl == types.Impl.WARP:
- raise NotImplementedError('Warp implementation not implemented yet.')
+ return _put_model_warp(m, device)
else:
raise ValueError(f'Unsupported implementation: {impl}')
@@ -719,11 +832,81 @@ def _make_data_c(
return d
+def _get_nested_attr(obj: Any, attr_name: str, split: str) -> Any:
+ """Returns the nested attribute from an object."""
+ for part in attr_name.split(split):
+ obj = getattr(obj, part)
+ return obj
+
+
+def _make_data_warp(
+ m: Union[types.Model, mujoco.MjModel],
+ device: Optional[jax.Device] = None,
+ nconmax: int = -1,
+ njmax: int = -1,
+) -> types.Data:
+ """Allocate and initialize Data for the Warp implementation."""
+ if not isinstance(m, mujoco.MjModel):
+ raise ValueError(
+ 'make_data for warp, only supports a mujoco.MjModel input, got'
+ f' {type(m)}.'
+ )
+
+ if not mjxw.WARP_INSTALLED:
+ raise RuntimeError('Warp is not installed.')
+
+ with wp.ScopedDevice('cpu'): # pylint: disable=undefined-variable
+ dw = mjwp.make_data(m, nworld=1, nconmax=nconmax, njmax=njmax) # pylint: disable=undefined-variable
+
+ fields = _make_data_public_fields(m)
+ for k in fields:
+ if k in {'userdata', 'plugin_state'}:
+ continue
+ if not hasattr(dw, k):
+ raise ValueError(f'Public data field {k} not found in Warp data.')
+ field = _wp_to_np_type(getattr(dw, k))
+ if mjxw.types.BATCH_DIM['Data'][k]:
+ field = field.reshape(field.shape[1:])
+ fields[k] = field
+
+ impl_fields = {}
+ for k in mjxw.types.DataWarp.__annotations__.keys():
+ field = _get_nested_attr(dw, k, split='__')
+ field = _wp_to_np_type(field)
+ if mjxw.types.BATCH_DIM['Data'][k]:
+ field = field.reshape(field.shape[1:])
+ impl_fields[k] = field
+
+ data = types.Data(
+ qpos=m.qpos0.astype(np.float32),
+ eq_active=m.eq_active0.astype(bool),
+ **fields,
+ _impl=mjxw.types.DataWarp(**impl_fields),
+ )
+
+ data = jax.device_put(data, device=device)
+
+ with wp.ScopedDevice('cuda:0'): # pylint: disable=undefined-variable
+ # Warm-up the warp kernel cache.
+ # TODO(robotics-simulation): remove this warmup compilation once warp
+ # stops unloading modules during XLA graph capture for tile kernels.
+ # pylint: disable=undefined-variable
+ dw = mjwp.make_data(m, nworld=1)
+ mw = mjwp.put_model(m)
+ _ = mjwp.step(mw, dw)
+ # pylint: enable=undefined-variable
+ del dw, mw
+
+ return data
+
+
def make_data(
m: Union[types.Model, mujoco.MjModel],
device: Optional[jax.Device] = None,
impl: Optional[Union[str, types.Impl]] = None,
_full_compat: bool = False, # pylint: disable=invalid-name
+ nconmax: int = -1,
+ njmax: int = -1,
) -> types.Data:
"""Allocate and initialize Data.
@@ -734,6 +917,8 @@ def make_data(
_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.
+ nconmax: maximum number of contacts to allocate for warp
+ njmax: maximum number of constraints to allocate for warp
Returns:
an initialized mjx.Data placed on device
@@ -764,6 +949,8 @@ def make_data(
return _make_data_jax(m, device)
elif impl == types.Impl.C:
return _make_data_c(m, device)
+ elif impl == types.Impl.WARP:
+ return _make_data_warp(m, device, nconmax, njmax)
raise NotImplementedError(
f'make_data for implementation "{impl}" not implemented yet.'
@@ -1053,6 +1240,8 @@ def put_data(
d: mujoco.MjData,
device: Optional[jax.Device] = None,
impl: Optional[Union[str, types.Impl]] = None,
+ nconmax: int = -1,
+ njmax: int = -1,
_full_compat: bool = False, # pylint: disable=invalid-name
) -> types.Data:
"""Puts mujoco.MjData onto a device, resulting in mjx.Data.
@@ -1062,6 +1251,8 @@ def put_data(
d: the data to put on device
device: which device to use - if unspecified picks the default device
impl: implementation to use ('jax', 'warp')
+ nconmax: maximum number of contacts to allocate for warp
+ njmax: maximum number of constraints to allocate for warp
_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.
@@ -1069,6 +1260,7 @@ def put_data(
Returns:
an mjx.Data placed on device
"""
+ del nconmax, njmax
if _full_compat:
warnings.warn(
'mjx.put_data(..., _full_compat=True) is deprecated. Use'
@@ -1084,6 +1276,8 @@ def put_data(
elif impl == types.Impl.C:
return _put_data_c(m, d, device)
+ # TODO(robotics-team): implement put_data_warp
+
raise NotImplementedError(
f'put_data for implementation "{impl}" not implemented yet.'
)
@@ -1099,6 +1293,86 @@ def _get_contact(c: mujoco._structs._MjContactList, cx: types.Contact):
getattr(c, field.name)[:] = value
+def _get_data_into_warp(
+ 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
+
+ for i in range(batch_size):
+ d_i = (
+ jax.tree.map_with_path(
+ lambda path, x, i=i: x[i]
+ if path[-1].name not in mjxw.types.DATA_NON_VMAP
+ else x,
+ d,
+ )
+ if batched
+ else d
+ )
+ result_i = result[i] if batched else result
+ ncon = d_i._impl.ncon[0]
+ nefc = int(d_i._impl.nefc[0])
+ # nj = int(d_i._impl.nj[0])
+ nj = 0 # TODO(btaba): add nj back
+
+ if ncon != result_i.ncon or nefc != result_i.nefc or nj != result_i.nJ:
+ mujoco._functions._realloc_con_efc(result_i, ncon=ncon, nefc=nefc, nJ=nj) # pylint: disable=protected-access
+
+ all_fields = types.Data.fields() + mjxw.types.DataWarp.fields()
+ for field in all_fields:
+ if field.name not in mujoco.MjData.__dict__.keys():
+ continue
+
+ # TODO(btaba): contact
+ # TODO(btaba): actuator_moment
+
+ 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 ('ne', 'nl', 'nf'):
+ value = value[0]
+ elif 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 == 'efc_J': # TODO(btaba): add this back
+ # elif field.name.startswith('efc_'): # TODO(btaba): add this back
+ # TODO(btaba): qM, qLD, qLDiagInv
+
+ if field.name in (
+ 'actuator_moment',
+ 'contact',
+ 'efc_J',
+ 'qM',
+ 'qLD',
+ 'qLDiagInv',
+ ):
+ continue
+ if field.name.startswith('efc_'):
+ continue
+
+ 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)
+
+ # TODO(btaba): add M back
+ # mujoco.mj_factorM(m, result_i)
+
+
def _get_data_into(
result: Union[mujoco.MjData, List[mujoco.MjData]],
m: mujoco.MjModel,
@@ -1250,6 +1524,9 @@ def get_data_into(
# TODO(stunya): Split out _get_data_into once codepaths diverge enough.
return _get_data_into(result, m, d)
+ if d.impl == types.Impl.WARP:
+ return _get_data_into_warp(result, m, d)
+
raise NotImplementedError(
f'get_data_into for implementation "{d.impl}" not implemented yet.'
)
diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py
index 3edef119..b2fc6e11 100644
--- a/mjx/mujoco/mjx/_src/io_test.py
+++ b/mjx/mujoco/mjx/_src/io_test.py
@@ -24,12 +24,13 @@ 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 ConeType
from mujoco.mjx._src.types import Impl
from mujoco.mjx._src.types import JacobianType
# pylint: enable=g-importing-member
+import mujoco.mjx.warp as mjxw
+from mujoco.mjx.warp import types as mjxw_types
import numpy as np
@@ -110,14 +111,31 @@ _SIMPLE_BODY = """
"""
+def _get_name_from_path(path: jax.tree_util.KeyPath) -> str:
+ """Returns a flattened name from a jax.tree_util.KeyPath."""
+ if any(isinstance(p, jax.tree_util.SequenceKey) for p in path):
+ is_seq_key = [isinstance(p, jax.tree_util.SequenceKey) for p in path]
+ path = path[: is_seq_key.index(True)]
+ assert all(isinstance(p, jax.tree_util.GetAttrKey) for p in path)
+ path = [p for p in path if p.name != '_impl']
+ attr = '__'.join(p.name for p in path)
+ return attr
+
+
class ModelIOTest(parameterized.TestCase):
"""IO tests for mjx.Model."""
@parameterized.product(
xml=(_MULTIPLE_CONVEX_OBJECTS, _MULTIPLE_CONSTRAINTS),
- impl=('jax', 'c'),
+ impl=('jax', 'c', 'warp'),
)
+ @mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'true'})
def test_put_model(self, xml, impl):
+ if impl == 'warp' and not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp not installed.')
+ if impl == 'warp' and not mjx_io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device available.')
+
m = mujoco.MjModel.from_xml_string(xml)
mx = mjx.put_model(m, impl=impl)
@@ -146,9 +164,14 @@ class ModelIOTest(parameterized.TestCase):
self.assertFalse(hasattr(mx, 'bvh_aabb'))
elif impl == 'c':
# Options specific to C are populated.
- self.assertEqual(mx.opt.apirate, m.opt.apirate)
+ self.assertEqual(mx.opt._impl.apirate, m.opt.apirate)
# Fields private to C backend impl are populated.
self.assertTrue(hasattr(mx._impl, 'bvh_aabb'))
+ elif impl == 'warp':
+ # Options specific to Warp are populated.
+ self.assertTrue(hasattr(mx.opt._impl, 'ls_parallel'))
+ # Fields private to Warp backend impl are populated.
+ self.assertTrue(hasattr(mx._impl, 'nxn_geom_pair'))
np.testing.assert_allclose(mx.body_parentid, m.body_parentid)
np.testing.assert_allclose(mx.geom_type, m.geom_type)
@@ -170,7 +193,7 @@ class ModelIOTest(parameterized.TestCase):
np.testing.assert_equal(mx.wrap_type, m.wrap_type)
np.testing.assert_equal(mx.wrap_objid, m.wrap_objid)
- np.testing.assert_equal(mx.wrap_prm, m.wrap_prm)
+ np.testing.assert_almost_equal(mx.wrap_prm, m.wrap_prm)
def test_fluid_params(self):
"""Test that has_fluid_params is set when fluid params are present."""
@@ -180,7 +203,7 @@ class ModelIOTest(parameterized.TestCase):
),
impl='jax',
)
- self.assertTrue(m.opt.has_fluid_params)
+ self.assertTrue(m.opt._impl.has_fluid_params)
def test_implicit_not_implemented(self):
"""Test that MJX guards against models with unimplemented features."""
@@ -276,6 +299,29 @@ class ModelIOTest(parameterized.TestCase):
with self.assertRaises(NotImplementedError):
mjx.put_model(m, impl='jax')
+ def test_put_model_warp_has_expected_shapes(self):
+ """Tests that put_model produces expected shapes for MuJoCo Warp."""
+ if not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp not installed.')
+ if not mjx_io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device available.')
+
+ m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS)
+ mx = mjx.put_model(m, impl='warp')
+
+ def check_ndim(path, x):
+ k = _get_name_from_path(path)
+ if k not in mjxw_types.NDIM['Model']:
+ return
+ is_batched = mjxw_types.BATCH_DIM['Model'][k]
+ expected_ndim = mjxw_types.NDIM['Model'][k] - is_batched
+ if not hasattr(x, 'ndim'):
+ return
+ msg = f'Field {k} has ndim {x.ndim} but expected {expected_ndim}'
+ self.assertEqual(x.ndim, expected_ndim, msg)
+
+ _ = jax.tree.map_with_path(check_ndim, mx)
+
class DataIOTest(parameterized.TestCase):
"""IO tests for mjx.Data."""
@@ -366,6 +412,17 @@ class DataIOTest(parameterized.TestCase):
self.assertEqual(d._impl.light_xpos.shape, (m.nlight, 3))
self.assertEqual(d._impl.bvh_active.shape, (m.nbvh,))
+ @mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'true'})
+ def test_make_data_warp(self):
+ if not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp is not installed.')
+ if not mjx_io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device.')
+ m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONVEX_OBJECTS)
+ d = mjx.make_data(m, impl='warp', nconmax=9, njmax=11)
+ self.assertEqual(d._impl.contact__dist.shape[0], 9)
+ self.assertEqual(d._impl.efc__J.shape[0], 11)
+
@parameterized.parameters('jax', 'c')
def test_put_data(self, impl: str):
"""Test that put_data puts the correct data for dense and sparse."""
@@ -676,6 +733,47 @@ class DataIOTest(parameterized.TestCase):
np.testing.assert_allclose(res_mj[0], res, rtol=1e-3, atol=1e-3)
+ def test_make_data_warp_has_expected_shapes(self):
+ """Tests that make_data produces expected shapes for MuJoCo Warp."""
+ if not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp is not installed.')
+ if not mjx_io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device.')
+
+ m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS)
+ dx = mjx.make_data(m, impl='warp')
+
+ def check_ndim(path, x):
+ k = _get_name_from_path(path)
+ if k not in mjxw_types.NDIM['Data']:
+ return
+ is_batched = mjxw_types.BATCH_DIM['Data'][k]
+ expected_ndim = mjxw_types.NDIM['Data'][k] - is_batched
+ if not hasattr(x, 'ndim'):
+ return
+ msg = f'Field {k} has ndim {x.ndim} but expected {expected_ndim}'
+ self.assertEqual(x.ndim, expected_ndim, msg)
+
+ _ = jax.tree.map_with_path(check_ndim, dx)
+
+ @parameterized.parameters('jax', 'warp')
+ def test_data_slice(self, impl):
+ """Tests that slice on Data works as expected."""
+ if impl == 'warp' and not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp is not installed.')
+ if impl == 'warp' and not mjx_io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device.')
+
+ m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS)
+ dx = jax.vmap(lambda x: mjx.make_data(m, impl=impl))(jp.arange(10))
+
+ self.assertEqual(dx.qpos.shape, (10, m.nq))
+ self.assertEqual(dx[0].qpos.shape, (m.nq,))
+
+ if impl == 'warp':
+ self.assertEqual(dx._impl.contact__dist.shape, (dx._impl.nconmax,))
+ self.assertEqual(dx[0]._impl.contact__dist.shape, (dx._impl.nconmax,))
+
class FullCompatTest(parameterized.TestCase):
"""Tests for the _full_compat flag."""
@@ -711,7 +809,7 @@ _DEVICE_TEST_CASES = [
# (device_type_str, impl_str,
# (expected_device, expected_impl)))
# No backend specified.
- ('cpu', None, ('cpu', Impl.C)),
+ ('cpu', None, ('cpu', Impl.JAX)),
('gpu-notnvidia', None, ('gpu', Impl.JAX)),
('gpu-nvidia', None, ('gpu', Impl.WARP)),
('tpu', None, ('tpu', Impl.JAX)),
@@ -739,7 +837,7 @@ _DEFAULT_DEVICE_TEST_CASES = [
# (jax.default_device, impl_str,
# (expected_device, expected_impl))
# No backend impl specified.
- ('cpu', None, ('cpu', Impl.C)),
+ ('cpu', None, ('cpu', Impl.JAX)),
('gpu-notnvidia', None, ('gpu', Impl.JAX)),
('gpu-nvidia', None, ('gpu', Impl.WARP)),
('tpu', None, ('tpu', Impl.JAX)),
@@ -790,6 +888,9 @@ class ResolveImplAndDeviceTest(parameterized.TestCase):
# Patch jax.devices for the entire test class using enter_context
self.mock_jax_devices = self.enter_context(mock.patch('jax.devices'))
+ self.mock_jax_backends = self.enter_context(
+ mock.patch('jax.extend.backend.backends')
+ )
self.mock_default_backend = self.enter_context(
mock.patch('jax.default_backend')
)
@@ -797,9 +898,7 @@ class ResolveImplAndDeviceTest(parameterized.TestCase):
@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'}
- )
+ @mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'true'})
def test_resolve_with_device(
self,
device_type_str,
@@ -819,7 +918,7 @@ class ResolveImplAndDeviceTest(parameterized.TestCase):
if backend == 'cpu':
return [self.mock_cpu]
elif backend == 'gpu':
- if 'nvidia' in device_type_str:
+ if device_type_str == 'gpu-nvidia':
return [self.mock_nvidia_gpu]
return [self.mock_other_gpu]
elif backend == 'tpu':
@@ -828,9 +927,19 @@ class ResolveImplAndDeviceTest(parameterized.TestCase):
return [self.mock_nvidia_gpu]
raise AssertionError('Should not be called.')
-
self.mock_jax_devices.side_effect = devices_side_effect
+ def backends_side_effect():
+ if device_type_str == 'gpu-nvidia':
+ return ['cuda', 'cpu']
+ if 'tpu' in device_type_str:
+ return ['tpu', 'cpu']
+ if 'gpu' in device_type_str:
+ return ['gpu', 'cpu']
+ return ['cpu']
+
+ self.mock_jax_backends.side_effect = backends_side_effect
+
expected_device, expected_impl = expected
if expected_impl == 'error':
with self.assertRaises(AssertionError):
@@ -839,6 +948,14 @@ class ResolveImplAndDeviceTest(parameterized.TestCase):
)
return
+ if impl_str == 'warp' and not mjxw.WARP_INSTALLED:
+ with self.assertRaisesRegex(RuntimeError, 'not installed'):
+ mjx_io._resolve_impl_and_device(impl=impl_str, device=input_device)
+ return
+
+ if impl_str is None and not mjxw.WARP_INSTALLED:
+ expected_impl = Impl.JAX
+
actual_impl, actual_device = (
mjx_io._resolve_impl_and_device(
impl=impl_str, device=input_device
@@ -853,9 +970,7 @@ class ResolveImplAndDeviceTest(parameterized.TestCase):
(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'}
- )
+ @mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'true'})
def test_resolve_without_device(
self,
default_device_str,
@@ -886,8 +1001,8 @@ class ResolveImplAndDeviceTest(parameterized.TestCase):
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',
@@ -898,6 +1013,17 @@ class ResolveImplAndDeviceTest(parameterized.TestCase):
lambda: default_device_side_effect_str
)
+ def backends_side_effect():
+ if default_device_str == 'gpu-nvidia':
+ return ['cuda', 'cpu']
+ if 'tpu' in default_device_str:
+ return ['tpu', 'cpu']
+ if 'gpu' in default_device_str:
+ return ['gpu', 'cpu']
+ return ['cpu']
+
+ self.mock_jax_backends.side_effect = backends_side_effect
+
expected_device, expected_impl = expected
if (
expected_impl == 'error'
@@ -905,31 +1031,36 @@ class ResolveImplAndDeviceTest(parameterized.TestCase):
and impl_str == 'warp'
):
with self.assertRaisesRegex(RuntimeError, 'cuda backend not supported'):
- mjx_io._resolve_impl_and_device(
- impl=impl_str, device=None
- )
+ mjx_io._resolve_impl_and_device(impl=impl_str, device=None)
return
if expected_impl == 'error':
with self.assertRaises(AssertionError):
- mjx_io._resolve_impl_and_device(
- impl=impl_str, device=None
- )
+ mjx_io._resolve_impl_and_device(impl=impl_str, device=None)
return
- actual_impl, actual_device = (
- mjx_io._resolve_impl_and_device(
- impl=impl_str, device=None
- )
+ if impl_str == 'warp' and not mjxw.WARP_INSTALLED:
+ with self.assertRaises(RuntimeError):
+ mjx_io._resolve_impl_and_device(impl=impl_str, device=None)
+ return
+
+ if impl_str is None and not mjxw.WARP_INSTALLED:
+ expected_impl = Impl.JAX
+
+ actual_impl, actual_device = mjx_io._resolve_impl_and_device(
+ impl=impl_str, device=None
)
self.assertEqual(actual_impl, expected_impl)
self.assertIsNotNone(actual_device)
self.assertEqual(actual_device.platform, expected_device)
- @mock.patch.dict(os.environ, {'MJX_WARP_ENABLED': 'false'})
+ @mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'false'})
def test_resolve_warp_disabled(self):
- """Tests behavior when MJX_WARP_ENABLED is false."""
+ """Tests behavior when MJX_GPU_DEFAULT_WARP is false."""
+ if not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp is not installed.')
+
self.mock_jax_devices.side_effect = lambda backend=None: (
[self.mock_nvidia_gpu, self.mock_cpu]
if backend is None
@@ -951,64 +1082,6 @@ class ResolveImplAndDeviceTest(parameterized.TestCase):
self.assertEqual(impl, Impl.JAX)
self.assertEqual(device.platform, 'gpu')
- # Requesting warp explicitly should fail since it is disabled.
- with self.assertRaises(AssertionError):
- mjx_io._resolve_impl_and_device(
- impl='warp', device=self.mock_nvidia_gpu
- )
- with self.assertRaises(AssertionError):
- mjx_io._resolve_impl_and_device(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.
- impl, device = mjx_io._resolve_impl_and_device(
- impl=None, device=None
- )
- self.assertEqual(impl, Impl.JAX)
- self.assertEqual(device.platform, 'cpu')
-
- # Specifing CPU should still choose JAX.
- impl, device = mjx_io._resolve_impl_and_device(
- impl=None, device=self.mock_cpu
- )
- self.assertEqual(impl, Impl.JAX)
- self.assertEqual(device.platform, 'cpu')
-
- # Specifying C should choose C!
- impl, device = mjx_io._resolve_impl_and_device(
- impl='c', device=None
- )
- self.assertEqual(impl, Impl.C)
- self.assertEqual(device.platform, 'cpu')
-
- impl, device = mjx_io._resolve_impl_and_device(
- impl='c', device=self.mock_cpu
- )
- self.assertEqual(impl, Impl.C)
- self.assertEqual(device.platform, 'cpu')
-
- def test_flex_jax(self):
- with self.assertRaises(NotImplementedError):
- m = mujoco.MjModel.from_xml_string("""
-
-
-
-
-
-
-
- """)
- mjx.put_model(m, impl='jax')
-
if __name__ == '__main__':
absltest.main()
diff --git a/mjx/mujoco/mjx/_src/passive.py b/mjx/mujoco/mjx/_src/passive.py
index b4a8b674..f7acdcd2 100644
--- a/mjx/mujoco/mjx/_src/passive.py
+++ b/mjx/mujoco/mjx/_src/passive.py
@@ -116,7 +116,11 @@ 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):
+ if (
+ not isinstance(m._impl, ModelJAX)
+ or not isinstance(d._impl, DataJAX)
+ or not isinstance(m.opt._impl, OptionJAX)
+ ):
raise ValueError('passive requires JAX backend implementation.')
if m.opt.disableflags & DisableBit.PASSIVE:
@@ -130,7 +134,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: # pytype: disable=attribute-error
+ if m.opt._impl.has_fluid_params: # pytype: disable=attribute-error
qfrc_passive += _fluid(m, d)
d = d.replace(qfrc_passive=qfrc_passive, qfrc_gravcomp=qfrc_gravcomp)
@@ -147,6 +151,11 @@ def _inertia_box_fluid_model(
cvel: jax.Array,
) -> Tuple[jax.Array, jax.Array]:
"""Fluid forces based on inertia-box approximation."""
+ if not isinstance(m.opt._impl, OptionJAX):
+ raise ValueError(
+ '_inertia_box_fluid_model requires JAX backend implementation.'
+ )
+
box = jp.repeat(inertia[None, :], 3, axis=0)
box *= jp.ones((3, 3)) - 2 * jp.eye(3)
box = 6.0 * jp.clip(jp.sum(box, axis=-1), a_min=1e-12)
diff --git a/mjx/mujoco/mjx/_src/solver.py b/mjx/mujoco/mjx/_src/solver.py
index 90c1add8..f57dbd1c 100644
--- a/mjx/mujoco/mjx/_src/solver.py
+++ b/mjx/mujoco/mjx/_src/solver.py
@@ -28,6 +28,7 @@ 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 OptionJAX
from mujoco.mjx._src.types import SolverType
# pylint: enable=g-importing-member
@@ -75,7 +76,9 @@ class Context(PyTreeNode):
@classmethod
def create(cls, m: Model, d: Data, grad: bool = True) -> 'Context':
- if not isinstance(d._impl, DataJAX):
+ if not isinstance(d._impl, DataJAX) or not isinstance(
+ m.opt._impl, OptionJAX
+ ):
raise ValueError(
'Constraint context requires JAX backend implementation.'
)
@@ -430,7 +433,11 @@ 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):
+ if (
+ not isinstance(m._impl, ModelJAX)
+ or not isinstance(d._impl, DataJAX)
+ or not isinstance(m.opt._impl, OptionJAX)
+ ):
raise ValueError('_lineasearch requires JAX backend implementation.')
smag = math.norm(ctx.search) * m.stat.meaninertia * max(1, m.nv)
@@ -549,6 +556,8 @@ def _linesearch(m: Model, d: Data, ctx: Context) -> Context:
def solve(m: Model, d: Data) -> Data:
"""Finds forces that satisfy constraints using conjugate gradient descent."""
+ if not isinstance(m.opt._impl, OptionJAX):
+ raise ValueError('solve requires JAX backend implementation.')
def cond(ctx: Context) -> jax.Array:
improvement = _rescale(m, ctx.prev_cost - ctx.cost)
diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py
index ae21fa2e..a71f934c 100644
--- a/mjx/mujoco/mjx/_src/types.py
+++ b/mjx/mujoco/mjx/_src/types.py
@@ -21,6 +21,7 @@ import warnings
import jax
import mujoco
from mujoco.mjx._src.dataclasses import PyTreeNode # pylint: disable=g-importing-member
+from mujoco.mjx.warp import types as mjxw_types
import numpy as np
@@ -465,48 +466,65 @@ class Statistic(PyTreeNode):
center: jax.Array
-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
+class StatisticWarp(mjxw_types.StatisticWarp, Statistic):
+ """Warp-specific model statistics."""
+
+ # NB: StatisticWarp type annotations may not match those on Statistic.
+ pass
+
+
+class OptionJAX(PyTreeNode):
+ """JAX-specific option."""
+
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
sdf_initpoints: int
+ has_fluid_params: bool
-class OptionC(Option):
+class OptionC(PyTreeNode):
"""C-specific option."""
+ o_margin: jax.Array
+ o_solref: jax.Array
+ o_solimp: jax.Array
+ o_friction: jax.Array
+ disableactuator: int
+ sdf_initpoints: int
+ has_fluid_params: bool
apirate: jax.Array
noslip_tolerance: jax.Array
ccd_tolerance: jax.Array
noslip_iterations: int
ccd_iterations: int
sdf_iterations: int
+ sdf_initpoints: int
-class OptionJAX(Option):
- """JAX-specific option."""
+class Option(PyTreeNode):
+ """Physics options."""
- has_fluid_params: bool
+ iterations: int
+ ls_iterations: int
+ tolerance: jax.Array
+ ls_tolerance: jax.Array
+ impratio: jax.Array
+ gravity: jax.Array
+ density: jax.Array
+ viscosity: jax.Array
+ magnetic: jax.Array
+ wind: jax.Array
+ jacobian: JacobianType
+ cone: ConeType
+ disableflags: DisableBit
+ enableflags: int
+ integrator: IntegratorType
+ solver: SolverType
+ timestep: jax.Array
+ _impl: Union[OptionJAX, OptionC, mjxw_types.OptionWarp]
class ModelC(PyTreeNode):
@@ -640,7 +658,7 @@ class Model(PyTreeNode):
nsensordata: int
npluginstate: int
opt: Option
- stat: Statistic
+ stat: Union[Statistic, StatisticWarp]
qpos0: jax.Array
qpos_spring: jax.Array
body_parentid: np.ndarray
@@ -743,8 +761,8 @@ class Model(PyTreeNode):
light_pos: jax.Array
light_dir: jax.Array
light_poscom0: jax.Array
- light_pos0: np.ndarray
- light_dir0: np.ndarray
+ light_pos0: jax.Array
+ light_dir0: jax.Array
light_cutoff: jax.Array
mesh_vertadr: np.ndarray
mesh_vertnum: np.ndarray
@@ -883,13 +901,14 @@ class Model(PyTreeNode):
names: bytes
signature: np.uint64
_sizes: jax.Array
- _impl: Union[ModelC, ModelJAX]
+ _impl: Union[ModelC, ModelJAX, mjxw_types.ModelWarp]
@property
def impl(self) -> Impl:
return {
ModelC: Impl.C,
ModelJAX: Impl.JAX,
+ mjxw_types.ModelWarp: Impl.WARP,
}[type(self._impl)]
def __getattr__(self, name: str):
@@ -1132,13 +1151,14 @@ class Data(PyTreeNode):
qacc_smooth: jax.Array
qfrc_constraint: jax.Array
qfrc_inverse: jax.Array
- _impl: Union[DataC, DataJAX]
+ _impl: Union[DataC, DataJAX, mjxw_types.DataWarp]
@property
def impl(self) -> Impl:
return {
DataC: Impl.C,
DataJAX: Impl.JAX,
+ mjxw_types.DataWarp: Impl.WARP,
}[type(self._impl)]
def __getattr__(self, name: str):
@@ -1157,3 +1177,23 @@ class Data(PyTreeNode):
f"'{type(self).__name__}' object has no attribute '{name}'"
)
return val
+
+ def __getitem__(self, key):
+ def get_name_from_path(path: jax.tree_util.KeyPath) -> str:
+ if any(isinstance(p, jax.tree_util.SequenceKey) for p in path):
+ is_seq_key = [isinstance(p, jax.tree_util.SequenceKey) for p in path]
+ path = path[: is_seq_key.index(True)]
+ assert all(isinstance(p, jax.tree_util.GetAttrKey) for p in path)
+ path = [p for p in path if p.name != '_impl']
+ attr = '__'.join(p.name for p in path)
+ return attr
+
+ if self.impl == Impl.WARP:
+ return jax.tree.map_with_path(
+ lambda path, x, k=key: x[k]
+ if get_name_from_path(path) not in mjxw_types.DATA_NON_VMAP
+ else x,
+ self,
+ )
+
+ return jax.tree.map(lambda x: x[key], self)
diff --git a/mjx/mujoco/mjx/test_data/humanoid/humanoid.xml b/mjx/mujoco/mjx/test_data/humanoid/humanoid.xml
index 0851b8ec..168310bc 100644
--- a/mjx/mujoco/mjx/test_data/humanoid/humanoid.xml
+++ b/mjx/mujoco/mjx/test_data/humanoid/humanoid.xml
@@ -252,5 +252,11 @@
-0.24 -0.007 -0.34 -1.76 -0.466 -0.0415
-0.08 -0.01 -0.37 -0.685 -0.35 -0.09
0.109 -0.067 -0.7 -0.05 0.12 0.16"/>
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py
new file mode 100644
index 00000000..1b76ba88
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py
@@ -0,0 +1,79 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+"""Public API for MJWarp."""
+
+# isort: off
+from mujoco.mjx.third_party.mujoco_warp._src.forward import step as step
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model as Model
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data as Data
+# isort: on
+
+from mujoco.mjx.third_party.mujoco_warp._src.collision_driver import collision as collision
+from mujoco.mjx.third_party.mujoco_warp._src.collision_driver import nxn_broadphase as nxn_broadphase
+from mujoco.mjx.third_party.mujoco_warp._src.collision_driver import sap_broadphase as sap_broadphase
+from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import primitive_narrowphase as primitive_narrowphase
+from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import sdf_narrowphase as sdf_narrowphase
+from mujoco.mjx.third_party.mujoco_warp._src.constraint import make_constraint as make_constraint
+from mujoco.mjx.third_party.mujoco_warp._src.derivative import deriv_smooth_vel as deriv_smooth_vel
+from mujoco.mjx.third_party.mujoco_warp._src.forward import euler as euler
+from mujoco.mjx.third_party.mujoco_warp._src.forward import forward as forward
+from mujoco.mjx.third_party.mujoco_warp._src.forward import fwd_acceleration as fwd_acceleration
+from mujoco.mjx.third_party.mujoco_warp._src.forward import fwd_actuation as fwd_actuation
+from mujoco.mjx.third_party.mujoco_warp._src.forward import fwd_position as fwd_position
+from mujoco.mjx.third_party.mujoco_warp._src.forward import fwd_velocity as fwd_velocity
+from mujoco.mjx.third_party.mujoco_warp._src.forward import implicit as implicit
+from mujoco.mjx.third_party.mujoco_warp._src.forward import rungekutta4 as rungekutta4
+from mujoco.mjx.third_party.mujoco_warp._src.inverse import inverse as inverse
+from mujoco.mjx.third_party.mujoco_warp._src.io import get_data_into as get_data_into
+from mujoco.mjx.third_party.mujoco_warp._src.io import make_data as make_data
+from mujoco.mjx.third_party.mujoco_warp._src.io import put_data as put_data
+from mujoco.mjx.third_party.mujoco_warp._src.io import put_model as put_model
+from mujoco.mjx.third_party.mujoco_warp._src.passive import passive as passive
+from mujoco.mjx.third_party.mujoco_warp._src.ray import ray as ray
+from mujoco.mjx.third_party.mujoco_warp._src.sensor import energy_pos as energy_pos
+from mujoco.mjx.third_party.mujoco_warp._src.sensor import energy_vel as energy_vel
+from mujoco.mjx.third_party.mujoco_warp._src.sensor import sensor_acc as sensor_acc
+from mujoco.mjx.third_party.mujoco_warp._src.sensor import sensor_pos as sensor_pos
+from mujoco.mjx.third_party.mujoco_warp._src.sensor import sensor_vel as sensor_vel
+from mujoco.mjx.third_party.mujoco_warp._src.smooth import camlight as camlight
+from mujoco.mjx.third_party.mujoco_warp._src.smooth import com_pos as com_pos
+from mujoco.mjx.third_party.mujoco_warp._src.smooth import com_vel as com_vel
+from mujoco.mjx.third_party.mujoco_warp._src.smooth import crb as crb
+from mujoco.mjx.third_party.mujoco_warp._src.smooth import factor_m as factor_m
+from mujoco.mjx.third_party.mujoco_warp._src.smooth import kinematics as kinematics
+from mujoco.mjx.third_party.mujoco_warp._src.smooth import rne as rne
+from mujoco.mjx.third_party.mujoco_warp._src.smooth import rne_postconstraint as rne_postconstraint
+from mujoco.mjx.third_party.mujoco_warp._src.smooth import solve_m as solve_m
+from mujoco.mjx.third_party.mujoco_warp._src.smooth import subtree_vel as subtree_vel
+from mujoco.mjx.third_party.mujoco_warp._src.smooth import tendon as tendon
+from mujoco.mjx.third_party.mujoco_warp._src.smooth import transmission as transmission
+from mujoco.mjx.third_party.mujoco_warp._src.solver import solve as solve
+from mujoco.mjx.third_party.mujoco_warp._src.support import contact_force as contact_force
+from mujoco.mjx.third_party.mujoco_warp._src.support import mul_m as mul_m
+from mujoco.mjx.third_party.mujoco_warp._src.support import xfrc_accumulate as xfrc_accumulate
+from mujoco.mjx.third_party.mujoco_warp._src.test_util import benchmark as benchmark
+from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseType as BroadphaseType
+from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType as ConeType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Constraint as Constraint
+from mujoco.mjx.third_party.mujoco_warp._src.types import Contact as Contact
+from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit as DisableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import DynType as DynType
+from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit as EnableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import JointType as JointType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Option as Option
+from mujoco.mjx.third_party.mujoco_warp._src.types import SolverType as SolverType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Statistic as Statistic
+from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType as TrnType
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/__init__.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/__init__.py
new file mode 100644
index 00000000..2b276beb
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/__init__.py
@@ -0,0 +1,14 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py
new file mode 100644
index 00000000..a080679b
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py
@@ -0,0 +1,218 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+from functools import lru_cache
+
+import warp as wp
+
+
+@lru_cache(maxsize=None)
+def create_blocked_cholesky_func(block_size: int):
+ @wp.func
+ def blocked_cholesky_func(
+ # In:
+ tid_block: int,
+ A: wp.array(dtype=float, ndim=2),
+ active_matrix_size: int,
+ # Out:
+ L: wp.array(dtype=float, ndim=2),
+ ):
+ """
+ Computes the Cholesky factorization of a symmetric positive definite matrix A in blocks.
+ It returns a lower-triangular matrix L such that A = L L^T.
+ """
+
+ num_threads_per_block = wp.block_dim()
+
+ # Round up active_matrix_size to next multiple of block_size
+ n = ((active_matrix_size + block_size - 1) // block_size) * block_size
+
+ # Process the matrix in blocks along its leading dimension.
+ for k in range(0, n, block_size):
+ end = k + block_size
+
+ # Load current diagonal block A[k:end, k:end]
+ # and update with contributions from previously computed blocks.
+ A_kk_tile = wp.tile_load(A, shape=(block_size, block_size), offset=(k, k), storage="shared")
+ # The following if pads the matrix if it is not divisible by block_size
+ if k + block_size > active_matrix_size or k + block_size > active_matrix_size:
+ num_tile_elements = block_size * block_size
+ num_iterations = (num_tile_elements + num_threads_per_block - 1) // num_threads_per_block
+
+ for i in range(num_iterations):
+ linear_index = tid_block + i * num_threads_per_block
+ linear_index = linear_index % num_tile_elements
+ row = linear_index // block_size
+ col = linear_index % block_size
+ value = A_kk_tile[row, col]
+ if k + row >= active_matrix_size or k + col >= active_matrix_size:
+ value = wp.where(row == col, float(1), float(0))
+ A_kk_tile[row, col] = value
+
+ if k > 0:
+ for j in range(0, k, block_size):
+ L_block = wp.tile_load(L, shape=(block_size, block_size), offset=(k, j))
+ L_block_T = wp.tile_transpose(L_block)
+ L_L_T_block = wp.tile_matmul(L_block, L_block_T)
+ A_kk_tile -= L_L_T_block
+
+ # Compute the Cholesky factorization for the block
+ L_kk_tile = wp.tile_cholesky(A_kk_tile)
+ wp.tile_store(L, L_kk_tile, offset=(k, k))
+
+ # Process the blocks below the current block
+ for i in range(end, n, block_size):
+ A_ik_tile = wp.tile_load(A, shape=(block_size, block_size), offset=(i, k), storage="shared")
+ # The following if pads the matrix if it is not divisible by block_size
+ if i + block_size > active_matrix_size or k + block_size > active_matrix_size:
+ num_tile_elements = block_size * block_size
+ num_iterations = (num_tile_elements + num_threads_per_block - 1) // num_threads_per_block
+
+ for ii in range(num_iterations):
+ linear_index = tid_block + ii * num_threads_per_block
+ linear_index = linear_index % num_tile_elements
+ row = linear_index // block_size
+ col = linear_index % block_size
+ value = A_ik_tile[row, col]
+ if i + row >= active_matrix_size or k + col >= active_matrix_size:
+ value = wp.where(i + row == k + col, float(1), float(0))
+ A_ik_tile[row, col] = value
+
+ if k > 0:
+ for j in range(0, k, block_size):
+ L_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(i, j))
+ L_2_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(k, j))
+ L_T_tile = wp.tile_transpose(L_2_tile)
+ L_L_T_tile = wp.tile_matmul(L_tile, L_T_tile)
+ A_ik_tile -= L_L_T_tile
+
+ t = wp.tile_transpose(A_ik_tile)
+ tmp = wp.tile_lower_solve(L_kk_tile, t)
+ sol_tile = wp.tile_transpose(tmp)
+
+ wp.tile_store(L, sol_tile, offset=(i, k))
+
+ return blocked_cholesky_func
+
+
+@lru_cache(maxsize=None)
+def create_blocked_cholesky_solve_func(block_size: int):
+ @wp.func
+ def blocked_cholesky_solve_func(
+ # In:
+ tid_block: int,
+ L: wp.array(dtype=float, ndim=2),
+ b: wp.array(dtype=float, ndim=2),
+ tmp: wp.array(dtype=float, ndim=2),
+ active_matrix_size: int,
+ # Out:
+ x: wp.array(dtype=float, ndim=2),
+ ):
+ """
+ Solves A x = b given the Cholesky factor L (A = L L^T) using
+ blocked forward and backward substitution.
+ """
+
+ num_threads_per_block = wp.block_dim()
+
+ # Round up active_matrix_size to next multiple of block_size
+ n = ((active_matrix_size + block_size - 1) // block_size) * block_size
+
+ # Forward substitution: solve L y = b
+ for i in range(0, n, block_size):
+ i_end = i + block_size
+ rhs_tile = wp.tile_load(b, shape=(block_size, 1), offset=(i, 0))
+ if i > 0:
+ for j in range(0, i, block_size):
+ L_block = wp.tile_load(L, shape=(block_size, block_size), offset=(i, j))
+ y_block = wp.tile_load(tmp, shape=(block_size, 1), offset=(j, 0))
+ Ly_block = wp.tile_matmul(L_block, y_block)
+ rhs_tile -= Ly_block
+ L_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(i, i))
+
+ # The following if pads the matrix if it is not divisible by block_size
+ if i + block_size > active_matrix_size:
+ num_tile_elements = block_size * block_size
+ num_iterations = (num_tile_elements + num_threads_per_block - 1) // num_threads_per_block
+
+ for ii in range(num_iterations):
+ linear_index = tid_block + ii * num_threads_per_block
+ linear_index = linear_index % num_tile_elements
+ row = linear_index // block_size
+ col = linear_index % block_size
+ value = L_tile[row, col]
+ if i + row >= active_matrix_size or i + col >= active_matrix_size:
+ value = wp.where(row == col, float(1), float(0))
+ L_tile[row, col] = value
+
+ # Handle rhs
+ num_tile_elements = block_size
+ num_iterations = (num_tile_elements + num_threads_per_block - 1) // num_threads_per_block
+
+ for ii in range(num_iterations):
+ linear_index = tid_block + ii * num_threads_per_block
+ linear_index = linear_index % num_tile_elements
+ value = rhs_tile[linear_index, 0]
+ if i + linear_index >= active_matrix_size:
+ value = float(0)
+ rhs_tile[linear_index, 0] = value
+
+ y_tile = wp.tile_lower_solve(L_tile, rhs_tile)
+ wp.tile_store(tmp, y_tile, offset=(i, 0))
+
+ # Backward substitution: solve L^T x = y
+ for i in range(n - block_size, -1, -block_size):
+ i_end = i + block_size
+ rhs_tile = wp.tile_load(tmp, shape=(block_size, 1), offset=(i, 0))
+ if i_end < n:
+ for j in range(i_end, n, block_size):
+ L_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(j, i))
+ L_T_tile = wp.tile_transpose(L_tile)
+ x_tile = wp.tile_load(x, shape=(block_size, 1), offset=(j, 0))
+ L_T_x_tile = wp.tile_matmul(L_T_tile, x_tile)
+ rhs_tile -= L_T_x_tile
+ L_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(i, i))
+
+ # The following if pads the matrix if it is not divisible by block_size
+ if i + block_size > active_matrix_size:
+ num_tile_elements = block_size * block_size
+ num_iterations = (num_tile_elements + num_threads_per_block - 1) // num_threads_per_block
+
+ for ii in range(num_iterations):
+ linear_index = tid_block + ii * num_threads_per_block
+ linear_index = linear_index % num_tile_elements
+ row = linear_index // block_size
+ col = linear_index % block_size
+ value = L_tile[row, col]
+ if i + row >= active_matrix_size or i + col >= active_matrix_size:
+ value = wp.where(row == col, float(1), float(0))
+ L_tile[row, col] = value
+
+ # Handle rhs
+ num_tile_elements = block_size
+ num_iterations = (num_tile_elements + num_threads_per_block - 1) // num_threads_per_block
+
+ for ii in range(num_iterations):
+ linear_index = tid_block + ii * num_threads_per_block
+ linear_index = linear_index % num_tile_elements
+ value = rhs_tile[linear_index, 0]
+ if i + linear_index >= active_matrix_size:
+ value = float(0)
+ rhs_tile[linear_index, 0] = value
+
+ x_tile = wp.tile_upper_solve(wp.tile_transpose(L_tile), rhs_tile)
+ wp.tile_store(x, x_tile, offset=(i, 0))
+
+ return blocked_cholesky_solve_func
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/broadphase_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/broadphase_test.py
new file mode 100644
index 00000000..57a7b439
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/broadphase_test.py
@@ -0,0 +1,335 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 broadphase functions."""
+
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+from absl.testing import parameterized
+
+import mujoco_warp as mjwarp
+
+from mujoco.mjx.third_party.mujoco_warp._src import collision_driver
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseFilter
+from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseType
+
+
+def broadphase_caller(m, d):
+ if m.opt.broadphase == int(BroadphaseType.NXN):
+ collision_driver.nxn_broadphase(m, d)
+ else:
+ collision_driver.sap_broadphase(m, d)
+
+
+class BroadphaseTest(parameterized.TestCase):
+ # filter combinations
+ plane_sphere = BroadphaseFilter.PLANE.value | BroadphaseFilter.SPHERE.value
+ plane_aabb = BroadphaseFilter.PLANE.value | BroadphaseFilter.AABB.value
+ plane_obb = BroadphaseFilter.PLANE | BroadphaseFilter.OBB.value
+ plane_sphere_aabb = plane_sphere | BroadphaseFilter.AABB.value
+ plane_sphere_obb = plane_sphere | BroadphaseFilter.OBB.value
+ plane_sphere_aabb_obb = plane_sphere_aabb | BroadphaseFilter.OBB.value
+
+ @parameterized.product(
+ broadphase=list(BroadphaseType),
+ filter=[plane_sphere, plane_aabb, plane_obb, plane_sphere_aabb, plane_sphere_obb, plane_sphere_aabb_obb],
+ )
+ def test_broadphase(self, broadphase, filter):
+ """Tests collision broadphase algorithms."""
+
+ _XML = """
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """
+
+ # one world and zero collisions
+ mjm, _, m, d0 = test_util.fixture(xml=_XML, keyframe=0)
+
+ m.opt.broadphase = broadphase
+ m.opt.broadphase_filter = filter
+
+ broadphase_caller(m, d0)
+ np.testing.assert_allclose(d0.ncollision.numpy()[0], 0)
+
+ # one world and one collision
+ _, mjd1, _, d1 = test_util.fixture(xml=_XML, keyframe=1)
+ broadphase_caller(m, d1)
+
+ np.testing.assert_allclose(d1.ncollision.numpy()[0], 1)
+ np.testing.assert_allclose(d1.collision_pair.numpy()[0][0], 0)
+ np.testing.assert_allclose(d1.collision_pair.numpy()[0][1], 1)
+
+ # one world and three collisions
+ _, mjd2, _, d2 = test_util.fixture(xml=_XML, keyframe=2)
+ broadphase_caller(m, d2)
+
+ ncollision = d2.ncollision.numpy()[0]
+ np.testing.assert_allclose(ncollision, 3)
+
+ collision_pairs = [[0, 1], [0, 2], [1, 2]]
+ for i in range(ncollision):
+ self.assertTrue([d2.collision_pair.numpy()[i][0], d2.collision_pair.numpy()[i][1]] in collision_pairs)
+
+ # two worlds and four collisions
+ d3 = mjwarp.make_data(mjm, nworld=2, nconmax=512, njmax=512)
+ d3.geom_xpos = wp.array(
+ np.vstack([np.expand_dims(mjd1.geom_xpos, axis=0), np.expand_dims(mjd2.geom_xpos, axis=0)]),
+ dtype=wp.vec3,
+ )
+ d3.geom_xmat = wp.array(
+ np.vstack([np.expand_dims(mjd1.geom_xmat, axis=0), np.expand_dims(mjd2.geom_xmat, axis=0)]),
+ dtype=wp.mat33,
+ )
+ broadphase_caller(m, d3)
+
+ ncollision = d3.ncollision.numpy()[0]
+ np.testing.assert_allclose(ncollision, 4)
+
+ collision_pairs = [[[0, 1]], [[0, 1], [0, 2], [1, 2]]]
+ worldids = [0, 1, 1, 1]
+ for i in range(ncollision):
+ worldid = d3.collision_worldid.numpy()[i]
+ self.assertTrue(worldid == worldids[i])
+ self.assertTrue([d3.collision_pair.numpy()[i][0], d3.collision_pair.numpy()[i][1]] in collision_pairs[worldid])
+
+ # one world and zero collisions: contype and conaffinity incompatibility
+ mjm4, _, m4, d4 = test_util.fixture(xml=_XML, keyframe=1)
+ mjm4.geom_contype[:3] = 0
+ m4 = mjwarp.put_model(mjm4)
+
+ broadphase_caller(m4, d4)
+ np.testing.assert_allclose(d4.ncollision.numpy()[0], 0)
+
+ # one world and one collision: geomtype ordering
+ _, _, _, d5 = test_util.fixture(xml=_XML, keyframe=3)
+ broadphase_caller(m, d5)
+ np.testing.assert_allclose(d5.ncollision.numpy()[0], 1)
+ np.testing.assert_allclose(d5.collision_pair.numpy()[0][0], 3)
+ np.testing.assert_allclose(d5.collision_pair.numpy()[0][1], 2)
+
+ @parameterized.parameters((0, 0, 0), (0, 0.011, 1), (0.011, 0, 1), (0.00999, 0, 0), (0, 0.00999, 0), (0.00999, 0.00999, 0))
+ def test_broadphase_margin(self, margin1, margin2, ncollision):
+ _MJCF = f"""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """
+ _, _, m, d = test_util.fixture(xml=_MJCF, keyframe=0)
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], ncollision)
+
+ @parameterized.parameters(True, False)
+ def test_broadphase_filterparent(self, filterparent):
+ _MJCF = """
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """
+ _, _, m, d = test_util.fixture(xml=_MJCF, filterparent=filterparent, keyframe=0)
+
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 0 if filterparent else 1)
+
+ def test_broadphase_filter(self):
+ plane = BroadphaseFilter.PLANE.value
+ sphere = BroadphaseFilter.SPHERE.value
+ aabb = BroadphaseFilter.AABB.value
+ obb = BroadphaseFilter.OBB.value
+ plane_sphere = plane | sphere
+ plane_aabb = plane | aabb
+ plane_obb = plane | obb
+
+ _PLANE_CAPSULE_CAPSULE = """
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """
+
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=0)
+ m.opt.broadphase_filter = plane_sphere
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 0)
+
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=0)
+ m.opt.broadphase_filter = plane_aabb
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 0)
+
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=0)
+ m.opt.broadphase_filter = plane_obb
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 0)
+
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=1)
+ m.opt.broadphase_filter = plane_sphere
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 1)
+
+ # note: collision_driver._plane_filter checks bounding sphere
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=1)
+ m.opt.broadphase_filter = plane
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 2)
+
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=1)
+ m.opt.broadphase_filter = plane_sphere
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 1)
+
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=1)
+ m.opt.broadphase_filter = plane_obb
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 1)
+
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=2)
+ m.opt.broadphase_filter = plane_sphere
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 1)
+
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=2)
+ m.opt.broadphase_filter = plane_aabb
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 0)
+
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=2)
+ m.opt.broadphase_filter = plane_obb
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 0)
+
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=3)
+ m.opt.broadphase_filter = plane_sphere
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 1)
+
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=3)
+ m.opt.broadphase_filter = plane_aabb
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 1)
+
+ _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=3)
+ m.opt.broadphase_filter = plane_obb
+ broadphase_caller(m, d)
+ self.assertEqual(d.ncollision.numpy()[0], 0)
+
+ # TODO(team): test margin
+ # TODO(team): test DisableBit.FILTERPARENT
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py
new file mode 100644
index 00000000..8ddb801e
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py
@@ -0,0 +1,495 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import ccd
+from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk_legacy import epa_legacy
+from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk_legacy import gjk_legacy
+from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk_legacy import multicontact_legacy
+from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_prism_vertex
+from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import _geom
+from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import contact_params
+from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import write_contact
+from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame
+from mujoco.mjx.third_party.mujoco_warp._src.math import upper_tri_index
+from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index
+from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXCONPAIR
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec5
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel
+
+# TODO(team): improve compile time to enable backward pass
+wp.config.enable_backward = False
+
+MULTI_CONTACT_COUNT = 4
+mat3c = wp.types.matrix(shape=(MULTI_CONTACT_COUNT, 3), dtype=float)
+
+_CONVEX_COLLISION_PAIRS = [
+ (GeomType.HFIELD.value, GeomType.SPHERE.value),
+ (GeomType.HFIELD.value, GeomType.CAPSULE.value),
+ (GeomType.HFIELD.value, GeomType.ELLIPSOID.value),
+ (GeomType.HFIELD.value, GeomType.CYLINDER.value),
+ (GeomType.HFIELD.value, GeomType.BOX.value),
+ (GeomType.HFIELD.value, GeomType.MESH.value),
+ (GeomType.SPHERE.value, GeomType.ELLIPSOID.value),
+ (GeomType.SPHERE.value, GeomType.MESH.value),
+ (GeomType.CAPSULE.value, GeomType.ELLIPSOID.value),
+ (GeomType.CAPSULE.value, GeomType.CYLINDER.value),
+ (GeomType.CAPSULE.value, GeomType.MESH.value),
+ (GeomType.ELLIPSOID.value, GeomType.ELLIPSOID.value),
+ (GeomType.ELLIPSOID.value, GeomType.CYLINDER.value),
+ (GeomType.ELLIPSOID.value, GeomType.BOX.value),
+ (GeomType.ELLIPSOID.value, GeomType.MESH.value),
+ (GeomType.CYLINDER.value, GeomType.CYLINDER.value),
+ (GeomType.CYLINDER.value, GeomType.BOX.value),
+ (GeomType.CYLINDER.value, GeomType.MESH.value),
+ (GeomType.BOX.value, GeomType.MESH.value),
+ (GeomType.MESH.value, GeomType.MESH.value),
+]
+
+
+def _check_convex_collision_pairs():
+ prev_idx = -1
+ for pair in _CONVEX_COLLISION_PAIRS:
+ idx = upper_trid_index(len(GeomType), pair[0], pair[1])
+ if pair[1] < pair[0] or idx <= prev_idx:
+ return False
+ prev_idx = idx
+ return True
+
+
+assert _check_convex_collision_pairs(), "_CONVEX_COLLISION_PAIRS is in invalid order."
+
+
+@wp.func
+def _max_contacts_height_field(
+ # Model:
+ ngeom: int,
+ geom_type: wp.array(dtype=int),
+ geompair2hfgeompair: wp.array(dtype=int),
+ # In:
+ g1: int,
+ g2: int,
+ worldid: int,
+ # Data out:
+ ncon_hfield_out: wp.array2d(dtype=int),
+):
+ hfield = int(GeomType.HFIELD.value)
+ if geom_type[g1] == hfield or (geom_type[g2] == hfield):
+ geompairid = upper_tri_index(ngeom, g1, g2)
+ hfgeompairid = geompair2hfgeompair[geompairid]
+ hfncon = wp.atomic_add(ncon_hfield_out[worldid], hfgeompairid, 1)
+ if hfncon >= MJ_MAXCONPAIR:
+ return True
+
+ return False
+
+
+@cache_kernel
+def ccd_kernel_builder(
+ default_gjk: bool,
+ geomtype1: int,
+ geomtype2: int,
+ gjk_iterations: int,
+ epa_iterations: int,
+ epa_exact_neg_distance: bool,
+ depth_extension: float,
+):
+ # runs convex collision on a set of geom pairs to recover contact info
+ @nested_kernel
+ def ccd_kernel(
+ # Model:
+ ngeom: int,
+ geom_type: wp.array(dtype=int),
+ geom_condim: wp.array(dtype=int),
+ geom_dataid: wp.array(dtype=int),
+ geom_priority: wp.array(dtype=int),
+ geom_solmix: wp.array2d(dtype=float),
+ geom_solref: wp.array2d(dtype=wp.vec2),
+ geom_solimp: wp.array2d(dtype=vec5),
+ geom_size: wp.array2d(dtype=wp.vec3),
+ geom_friction: wp.array2d(dtype=wp.vec3),
+ geom_margin: wp.array2d(dtype=float),
+ geom_gap: wp.array2d(dtype=float),
+ hfield_adr: wp.array(dtype=int),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ hfield_data: wp.array(dtype=float),
+ mesh_vertadr: wp.array(dtype=int),
+ mesh_vertnum: wp.array(dtype=int),
+ mesh_vert: wp.array(dtype=wp.vec3),
+ mesh_graphadr: wp.array(dtype=int),
+ mesh_graph: wp.array(dtype=int),
+ mesh_polynum: wp.array(dtype=int),
+ mesh_polyadr: wp.array(dtype=int),
+ mesh_polynormal: wp.array(dtype=wp.vec3),
+ mesh_polyvertadr: wp.array(dtype=int),
+ mesh_polyvertnum: wp.array(dtype=int),
+ mesh_polyvert: wp.array(dtype=int),
+ mesh_polymapadr: wp.array(dtype=int),
+ mesh_polymapnum: wp.array(dtype=int),
+ mesh_polymap: wp.array(dtype=int),
+ pair_dim: wp.array(dtype=int),
+ pair_solref: wp.array2d(dtype=wp.vec2),
+ pair_solreffriction: wp.array2d(dtype=wp.vec2),
+ pair_solimp: wp.array2d(dtype=vec5),
+ pair_margin: wp.array2d(dtype=float),
+ pair_gap: wp.array2d(dtype=float),
+ pair_friction: wp.array2d(dtype=vec5),
+ geompair2hfgeompair: wp.array(dtype=int),
+ # Data in:
+ nconmax_in: int,
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ collision_pair_in: wp.array(dtype=wp.vec2i),
+ collision_hftri_index_in: wp.array(dtype=int),
+ collision_pairid_in: wp.array(dtype=int),
+ collision_worldid_in: wp.array(dtype=int),
+ ncollision_in: wp.array(dtype=int),
+ epa_vert_in: wp.array2d(dtype=wp.vec3),
+ epa_vert1_in: wp.array2d(dtype=wp.vec3),
+ epa_vert2_in: wp.array2d(dtype=wp.vec3),
+ epa_vert_index1_in: wp.array2d(dtype=int),
+ epa_vert_index2_in: wp.array2d(dtype=int),
+ epa_face_in: wp.array2d(dtype=wp.vec3i),
+ epa_pr_in: wp.array2d(dtype=wp.vec3),
+ epa_norm2_in: wp.array2d(dtype=float),
+ epa_index_in: wp.array2d(dtype=int),
+ epa_map_in: wp.array2d(dtype=int),
+ epa_horizon_in: wp.array2d(dtype=int),
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ ncon_hfield_out: wp.array2d(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+ ):
+ tid = wp.tid()
+ if tid >= ncollision_in[0]:
+ return
+
+ geoms = collision_pair_in[tid]
+ g1 = geoms[0]
+ g2 = geoms[1]
+
+ if geom_type[g1] != geomtype1 or geom_type[g2] != geomtype2:
+ return
+
+ worldid = collision_worldid_in[tid]
+ _, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params(
+ geom_condim,
+ geom_priority,
+ geom_solmix,
+ geom_solref,
+ geom_solimp,
+ geom_friction,
+ geom_margin,
+ geom_gap,
+ pair_dim,
+ pair_solref,
+ pair_solreffriction,
+ pair_solimp,
+ pair_margin,
+ pair_gap,
+ pair_friction,
+ collision_pair_in,
+ collision_pairid_in,
+ tid,
+ worldid,
+ )
+
+ hftri_index = collision_hftri_index_in[tid]
+
+ geom1 = _geom(
+ geom_type,
+ geom_dataid,
+ geom_size,
+ hfield_adr,
+ hfield_nrow,
+ hfield_ncol,
+ hfield_size,
+ hfield_data,
+ mesh_vertadr,
+ mesh_vertnum,
+ mesh_vert,
+ mesh_graphadr,
+ mesh_graph,
+ mesh_polynum,
+ mesh_polyadr,
+ mesh_polynormal,
+ mesh_polyvertadr,
+ mesh_polyvertnum,
+ mesh_polyvert,
+ mesh_polymapadr,
+ mesh_polymapnum,
+ mesh_polymap,
+ geom_xpos_in,
+ geom_xmat_in,
+ worldid,
+ g1,
+ hftri_index,
+ )
+
+ geom2 = _geom(
+ geom_type,
+ geom_dataid,
+ geom_size,
+ hfield_adr,
+ hfield_nrow,
+ hfield_ncol,
+ hfield_size,
+ hfield_data,
+ mesh_vertadr,
+ mesh_vertnum,
+ mesh_vert,
+ mesh_graphadr,
+ mesh_graph,
+ mesh_polynum,
+ mesh_polyadr,
+ mesh_polynormal,
+ mesh_polyvertadr,
+ mesh_polyvertnum,
+ mesh_polyvert,
+ mesh_polymapadr,
+ mesh_polymapnum,
+ mesh_polymap,
+ geom_xpos_in,
+ geom_xmat_in,
+ worldid,
+ g2,
+ hftri_index,
+ )
+
+ points = mat3c()
+
+ if default_gjk:
+ simplex, normal = gjk_legacy(
+ gjk_iterations,
+ geom1,
+ geom2,
+ geomtype1,
+ geomtype2,
+ )
+
+ depth, normal = epa_legacy(
+ epa_iterations, geom1, geom2, geomtype1, geomtype2, depth_extension, epa_exact_neg_distance, simplex, normal
+ )
+ dist = -depth
+
+ if (dist - margin) >= 0.0 or depth != depth:
+ return
+ sphere = int(GeomType.SPHERE.value)
+ ellipsoid = int(GeomType.ELLIPSOID.value)
+ if geom_type[g1] == sphere or geom_type[g1] == ellipsoid or geom_type[g2] == sphere or geom_type[g2] == ellipsoid:
+ count, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 1, 2, 1.0e-5)
+ else:
+ count, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 4, 8, 1.0e-3)
+ else:
+ x1 = geom1.pos
+ x2 = geom2.pos
+
+ # find prism center for height field
+ if geomtype1 == int(GeomType.HFIELD.value):
+ x1 = wp.vec3(0.0, 0.0, 0.0)
+ for i in range(6):
+ x1 += hfield_prism_vertex(geom1.hfprism, i)
+ x1 = x1 / 6.0
+
+ dist, x1, x2 = ccd(
+ 1e-6,
+ 0.0,
+ gjk_iterations,
+ epa_iterations,
+ geom1,
+ geom2,
+ geomtype1,
+ geomtype2,
+ x1,
+ x2,
+ epa_vert_in[tid],
+ epa_vert1_in[tid],
+ epa_vert2_in[tid],
+ epa_vert_index1_in[tid],
+ epa_vert_index2_in[tid],
+ epa_face_in[tid],
+ epa_pr_in[tid],
+ epa_norm2_in[tid],
+ epa_index_in[tid],
+ epa_map_in[tid],
+ epa_horizon_in[tid],
+ )
+ count = 0
+ if dist < 0.0:
+ count = 1
+
+ points[0] = 0.5 * (x1 + x2)
+ normal = x1 - x2
+
+ frame = make_frame(normal)
+ for i in range(count):
+ # limit maximum number of contacts with height field
+ if _max_contacts_height_field(ngeom, geom_type, geompair2hfgeompair, g1, g2, worldid, ncon_hfield_out):
+ return
+
+ write_contact(
+ nconmax_in,
+ dist,
+ points[i],
+ frame,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+ return ccd_kernel
+
+
+@event_scope
+def convex_narrowphase(m: Model, d: Data):
+ """Runs narrowphase collision detection for convex geom pairs.
+
+ This function handles collision detection for pairs of convex geometries that were
+ identified during the broadphase. It uses the Gilbert-Johnson-Keerthi (GJK) algorithm to
+ determine the distance between shapes and the Expanding Polytope Algorithm (EPA) to find
+ the penetration depth and contact normal for colliding pairs.
+
+ The convex geom types handled by this function are SPHERE, CAPSULE, ELLIPSOID, CYLINDER,
+ BOX, MESH, HFIELD.
+
+ To optimize performance, this function dynamically builds and launches a specialized
+ kernel for each type of convex collision pair present in the model, avoiding unnecessary
+ computations for non-existent pair types.
+ """
+ for geom_pair in _CONVEX_COLLISION_PAIRS:
+ if m.geom_pair_type_count[upper_trid_index(len(GeomType), geom_pair[0], geom_pair[1])]:
+ wp.launch(
+ ccd_kernel_builder(
+ False,
+ geom_pair[0],
+ geom_pair[1],
+ m.opt.gjk_iterations,
+ m.opt.epa_iterations,
+ False,
+ 0.1,
+ ),
+ dim=d.nconmax,
+ inputs=[
+ m.ngeom,
+ m.geom_type,
+ m.geom_condim,
+ m.geom_dataid,
+ m.geom_priority,
+ m.geom_solmix,
+ m.geom_solref,
+ m.geom_solimp,
+ m.geom_size,
+ m.geom_friction,
+ m.geom_margin,
+ m.geom_gap,
+ m.hfield_adr,
+ m.hfield_nrow,
+ m.hfield_ncol,
+ m.hfield_size,
+ m.hfield_data,
+ m.mesh_vertadr,
+ m.mesh_vertnum,
+ m.mesh_vert,
+ m.mesh_graphadr,
+ m.mesh_graph,
+ m.mesh_polynum,
+ m.mesh_polyadr,
+ m.mesh_polynormal,
+ m.mesh_polyvertadr,
+ m.mesh_polyvertnum,
+ m.mesh_polyvert,
+ m.mesh_polymapadr,
+ m.mesh_polymapnum,
+ m.mesh_polymap,
+ m.pair_dim,
+ m.pair_solref,
+ m.pair_solreffriction,
+ m.pair_solimp,
+ m.pair_margin,
+ m.pair_gap,
+ m.pair_friction,
+ m.geompair2hfgeompair,
+ d.nconmax,
+ d.geom_xpos,
+ d.geom_xmat,
+ d.collision_pair,
+ d.collision_hftri_index,
+ d.collision_pairid,
+ d.collision_worldid,
+ d.ncollision,
+ d.epa_vert,
+ d.epa_vert1,
+ d.epa_vert2,
+ d.epa_vert_index1,
+ d.epa_vert_index2,
+ d.epa_face,
+ d.epa_pr,
+ d.epa_norm2,
+ d.epa_index,
+ d.epa_map,
+ d.epa_horizon,
+ ],
+ outputs=[
+ d.ncon,
+ d.ncon_hfield,
+ d.contact.dist,
+ d.contact.pos,
+ d.contact.frame,
+ d.contact.includemargin,
+ d.contact.friction,
+ d.contact.solref,
+ d.contact.solreffriction,
+ d.contact.solimp,
+ d.contact.dim,
+ d.contact.geom,
+ d.contact.worldid,
+ ],
+ )
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py
new file mode 100644
index 00000000..4cdd85c2
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py
@@ -0,0 +1,767 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+from typing import Any
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src.collision_convex import convex_narrowphase
+from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_midphase
+from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import primitive_narrowphase
+from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import sdf_narrowphase
+from mujoco.mjx.third_party.mujoco_warp._src.math import upper_tri_index
+from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL
+from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseFilter
+from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
+
+wp.set_module_options({"enable_backward": False})
+
+
+@wp.kernel
+def _zero_collision_arrays(
+ # Data in:
+ nworld_in: int,
+ # In:
+ hfield_geom_pair_in: int,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ ncon_hfield_out: wp.array(dtype=int), # kernel_analyzer: ignore
+ collision_hftri_index_out: wp.array(dtype=int),
+ ncollision_out: wp.array(dtype=int),
+):
+ tid = wp.tid()
+
+ if tid == 0:
+ # Zero the single collision counter
+ ncollision_out[0] = 0
+ ncon_out[0] = 0
+
+ if tid < hfield_geom_pair_in * nworld_in:
+ ncon_hfield_out[tid] = 0
+
+ # Zero collision pair indices
+ collision_hftri_index_out[tid] = 0
+
+
+@wp.func
+def _plane_filter(
+ size1: float, size2: float, margin1: float, margin2: float, xpos1: wp.vec3, xpos2: wp.vec3, xmat1: wp.mat33, xmat2: wp.mat33
+) -> bool:
+ if size1 == 0.0:
+ # geom1 is a plane
+ dist = wp.dot(xpos2 - xpos1, wp.vec3(xmat1[0, 2], xmat1[1, 2], xmat1[2, 2]))
+ return dist <= size2 + wp.max(margin1, margin2)
+ elif size2 == 0.0:
+ # geom2 is a plane
+ dist = wp.dot(xpos1 - xpos2, wp.vec3(xmat2[0, 2], xmat2[1, 2], xmat2[2, 2]))
+ return dist <= size1 + wp.max(margin1, margin2)
+
+ return True
+
+
+@wp.func
+def _sphere_filter(size1: float, size2: float, margin1: float, margin2: float, xpos1: wp.vec3, xpos2: wp.vec3) -> bool:
+ bound = size1 + size2 + wp.max(margin1, margin2)
+ dif = xpos2 - xpos1
+ dist_sq = wp.dot(dif, dif)
+ return dist_sq <= bound * bound
+
+
+# TODO(team): improve performance by precomputing bounding box
+@wp.func
+def _aabb_filter(
+ # In:
+ center1: wp.vec3,
+ center2: wp.vec3,
+ size1: wp.vec3,
+ size2: wp.vec3,
+ margin1: float,
+ margin2: float,
+ xpos1: wp.vec3,
+ xpos2: wp.vec3,
+ xmat1: wp.mat33,
+ xmat2: wp.mat33,
+) -> bool:
+ """Axis aligned boxes collision.
+
+ references: see Ericson, Real-time Collision Detection section 4.2.
+ filterBox: filter contact based on global AABBs.
+ """
+ center1 = xmat1 @ center1 + xpos1
+ center2 = xmat2 @ center2 + xpos2
+
+ margin = wp.max(margin1, margin2)
+
+ max_x1 = -MJ_MAXVAL
+ max_y1 = -MJ_MAXVAL
+ max_z1 = -MJ_MAXVAL
+ min_x1 = MJ_MAXVAL
+ min_y1 = MJ_MAXVAL
+ min_z1 = MJ_MAXVAL
+
+ max_x2 = -MJ_MAXVAL
+ max_y2 = -MJ_MAXVAL
+ max_z2 = -MJ_MAXVAL
+ min_x2 = MJ_MAXVAL
+ min_y2 = MJ_MAXVAL
+ min_z2 = MJ_MAXVAL
+
+ sign = wp.vec2(-1.0, 1.0)
+
+ for i in range(2):
+ for j in range(2):
+ for k in range(2):
+ corner1 = wp.vec3(sign[i] * size1[0], sign[j] * size1[1], sign[k] * size1[2])
+ pos1 = xmat1 @ corner1
+
+ corner2 = wp.vec3(sign[i] * size2[0], sign[j] * size2[1], sign[k] * size2[2])
+ pos2 = xmat2 @ corner2
+
+ if pos1[0] > max_x1:
+ max_x1 = pos1[0]
+
+ if pos1[1] > max_y1:
+ max_y1 = pos1[1]
+
+ if pos1[2] > max_z1:
+ max_z1 = pos1[2]
+
+ if pos1[0] < min_x1:
+ min_x1 = pos1[0]
+
+ if pos1[1] < min_y1:
+ min_y1 = pos1[1]
+
+ if pos1[2] < min_z1:
+ min_z1 = pos1[2]
+
+ if pos2[0] > max_x2:
+ max_x2 = pos2[0]
+
+ if pos2[1] > max_y2:
+ max_y2 = pos2[1]
+
+ if pos2[2] > max_z2:
+ max_z2 = pos2[2]
+
+ if pos2[0] < min_x2:
+ min_x2 = pos2[0]
+
+ if pos2[1] < min_y2:
+ min_y2 = pos2[1]
+
+ if pos2[2] < min_z2:
+ min_z2 = pos2[2]
+
+ if center1[0] + max_x1 + margin < center2[0] + min_x2:
+ return False
+ if center1[1] + max_y1 + margin < center2[1] + min_y2:
+ return False
+ if center1[2] + max_z1 + margin < center2[2] + min_z2:
+ return False
+ if center2[0] + max_x2 + margin < center1[0] + min_x1:
+ return False
+ if center2[1] + max_y2 + margin < center1[1] + min_y1:
+ return False
+ if center2[2] + max_z2 + margin < center1[2] + min_z1:
+ return False
+
+ return True
+
+
+mat23 = wp.types.matrix(shape=(2, 3), dtype=float)
+mat63 = wp.types.matrix(shape=(6, 3), dtype=float)
+
+
+# TODO(team): improve performance by precomputing bounding box
+@wp.func
+def _obb_filter(
+ # In:
+ center1: wp.vec3,
+ center2: wp.vec3,
+ size1: wp.vec3,
+ size2: wp.vec3,
+ margin1: float,
+ margin2: float,
+ xpos1: wp.vec3,
+ xpos2: wp.vec3,
+ xmat1: wp.mat33,
+ xmat2: wp.mat33,
+) -> bool:
+ """Oriented bounding boxes collision (see Gottschalk et al.), see mj_collideOBB."""
+ margin = wp.max(margin1, margin2)
+
+ xcenter = mat23()
+ normal = mat63()
+ proj = wp.vec2()
+ radius = wp.vec2()
+
+ # compute centers in local coordinates
+ xcenter[0] = xmat1 @ center1 + xpos1
+ xcenter[1] = xmat2 @ center2 + xpos2
+
+ # compute normals in global coordinates
+ normal[0] = wp.vec3(xmat1[0, 0], xmat1[1, 0], xmat1[2, 0])
+ normal[1] = wp.vec3(xmat1[0, 1], xmat1[1, 1], xmat1[2, 1])
+ normal[2] = wp.vec3(xmat1[0, 2], xmat1[1, 2], xmat1[2, 2])
+ normal[3] = wp.vec3(xmat2[0, 0], xmat2[1, 0], xmat2[2, 0])
+ normal[4] = wp.vec3(xmat2[0, 1], xmat2[1, 1], xmat2[2, 1])
+ normal[5] = wp.vec3(xmat2[0, 2], xmat2[1, 2], xmat2[2, 2])
+
+ # check intersections
+ for j in range(2):
+ for k in range(3):
+ for i in range(2):
+ proj[i] = wp.dot(xcenter[i], normal[3 * j + k])
+ if i == 0:
+ size = size1
+ else:
+ size = size2
+
+ # fmt: off
+ radius[i] = (
+ wp.abs(size[0] * wp.dot(normal[3 * i + 0], normal[3 * j + k]))
+ + wp.abs(size[1] * wp.dot(normal[3 * i + 1], normal[3 * j + k]))
+ + wp.abs(size[2] * wp.dot(normal[3 * i + 2], normal[3 * j + k]))
+ )
+ # fmt: on
+ if radius[0] + radius[1] + margin < wp.abs(proj[1] - proj[0]):
+ return False
+
+ return True
+
+
+@wp.func
+def _broadphase_filter(
+ # Model:
+ opt_broadphase_filter: int,
+ geom_aabb: wp.array2d(dtype=wp.vec3),
+ geom_rbound: wp.array2d(dtype=float),
+ geom_margin: wp.array2d(dtype=float),
+ # Data in:
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ # In:
+ geom1: int,
+ geom2: int,
+ worldid: int,
+) -> bool:
+ # 1: plane
+ # 2: sphere
+ # 4: aabb
+ # 8: obb
+
+ center1 = geom_aabb[geom1, 0]
+ center2 = geom_aabb[geom2, 0]
+ size1 = geom_aabb[geom1, 1]
+ size2 = geom_aabb[geom2, 1]
+ rbound1 = geom_rbound[worldid, geom1]
+ rbound2 = geom_rbound[worldid, geom2]
+ margin1 = geom_margin[worldid, geom1]
+ margin2 = geom_margin[worldid, geom2]
+ xpos1 = geom_xpos_in[worldid, geom1]
+ xpos2 = geom_xpos_in[worldid, geom2]
+ xmat1 = geom_xmat_in[worldid, geom1]
+ xmat2 = geom_xmat_in[worldid, geom2]
+
+ if rbound1 == 0.0 or rbound2 == 0.0:
+ if opt_broadphase_filter & int(BroadphaseFilter.PLANE.value):
+ return _plane_filter(rbound1, rbound2, margin1, margin2, xpos1, xpos2, xmat1, xmat2)
+ else:
+ if opt_broadphase_filter & int(BroadphaseFilter.SPHERE.value):
+ if not _sphere_filter(rbound1, rbound2, margin1, margin2, xpos1, xpos2):
+ return False
+ if opt_broadphase_filter & int(BroadphaseFilter.AABB.value):
+ if not _aabb_filter(center1, center2, size1, size2, margin1, margin2, xpos1, xpos2, xmat1, xmat2):
+ return False
+ if opt_broadphase_filter & int(BroadphaseFilter.OBB.value):
+ if not _obb_filter(center1, center2, size1, size2, margin1, margin2, xpos1, xpos2, xmat1, xmat2):
+ return False
+
+ return True
+
+
+@wp.func
+def _add_geom_pair(
+ # Model:
+ geom_type: wp.array(dtype=int),
+ nxn_pairid: wp.array(dtype=int),
+ # Data in:
+ nconmax_in: int,
+ # In:
+ geom1: int,
+ geom2: int,
+ worldid: int,
+ nxnid: int,
+ # Data out:
+ collision_pair_out: wp.array(dtype=wp.vec2i),
+ collision_hftri_index_out: wp.array(dtype=int),
+ collision_pairid_out: wp.array(dtype=int),
+ collision_worldid_out: wp.array(dtype=int),
+ ncollision_out: wp.array(dtype=int),
+):
+ pairid = wp.atomic_add(ncollision_out, 0, 1)
+
+ if pairid >= nconmax_in:
+ return
+
+ type1 = geom_type[geom1]
+ type2 = geom_type[geom2]
+
+ if type1 > type2:
+ pair = wp.vec2i(geom2, geom1)
+ else:
+ pair = wp.vec2i(geom1, geom2)
+
+ collision_pair_out[pairid] = pair
+ collision_pairid_out[pairid] = nxn_pairid[nxnid]
+ collision_worldid_out[pairid] = worldid
+
+ # Writing -1 to collision_hftri_index_out[pairid] signals
+ # hfield_midphase to generate a collision pair for every
+ # potentially colliding triangle
+ if type1 == int(GeomType.HFIELD.value) or type2 == int(GeomType.HFIELD.value):
+ collision_hftri_index_out[pairid] = -1
+
+
+@wp.func
+def _binary_search(values: wp.array(dtype=Any), value: Any, lower: int, upper: int) -> int:
+ while lower < upper:
+ mid = (lower + upper) >> 1
+ if values[mid] > value:
+ upper = mid
+ else:
+ lower = mid + 1
+
+ return upper
+
+
+@wp.kernel
+def _sap_project(
+ # Model:
+ geom_rbound: wp.array2d(dtype=float),
+ geom_margin: wp.array2d(dtype=float),
+ # Data in:
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ # In:
+ direction_in: wp.vec3,
+ # Data out:
+ sap_projection_lower_out: wp.array2d(dtype=float), # kernel_analyzer: ignore
+ sap_projection_upper_out: wp.array2d(dtype=float),
+ sap_sort_index_out: wp.array2d(dtype=int), # kernel_analyzer: ignore
+):
+ worldid, geomid = wp.tid()
+
+ xpos = geom_xpos_in[worldid, geomid]
+ rbound = geom_rbound[worldid, geomid]
+
+ if rbound == 0.0:
+ # geom is a plane
+ rbound = MJ_MAXVAL
+
+ radius = rbound + geom_margin[worldid, geomid]
+ center = wp.dot(direction_in, xpos)
+
+ sap_sort_index_out[worldid, geomid] = geomid
+ if not wp.isnan(center):
+ sap_projection_lower_out[worldid, geomid] = center - radius
+ sap_projection_upper_out[worldid, geomid] = center + radius
+ else:
+ sap_projection_lower_out[worldid, geomid] = MJ_MAXVAL
+ sap_projection_upper_out[worldid, geomid] = MJ_MAXVAL
+
+
+@wp.kernel
+def _sap_range(
+ # Model:
+ ngeom: int,
+ # Data in:
+ sap_projection_lower_in: wp.array2d(dtype=float), # kernel_analyzer: ignore
+ sap_projection_upper_in: wp.array2d(dtype=float),
+ sap_sort_index_in: wp.array2d(dtype=int), # kernel_analyzer: ignore
+ # Data out:
+ sap_range_out: wp.array2d(dtype=int),
+):
+ worldid, geomid = wp.tid()
+
+ # current bounding geom
+ idx = sap_sort_index_in[worldid, geomid]
+
+ upper = sap_projection_upper_in[worldid, idx]
+
+ limit = _binary_search(sap_projection_lower_in[worldid], upper, geomid + 1, ngeom)
+ limit = wp.min(ngeom - 1, limit)
+
+ # range of geoms for the sweep and prune process
+ sap_range_out[worldid, geomid] = limit - geomid
+
+
+@wp.kernel
+def _sap_broadphase(
+ # Model:
+ ngeom: int,
+ opt_broadphase_filter: int,
+ geom_type: wp.array(dtype=int),
+ geom_aabb: wp.array2d(dtype=wp.vec3),
+ geom_rbound: wp.array2d(dtype=float),
+ geom_margin: wp.array2d(dtype=float),
+ nxn_pairid: wp.array(dtype=int),
+ # Data in:
+ nworld_in: int,
+ nconmax_in: int,
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ sap_sort_index_in: wp.array2d(dtype=int), # kernel_analyzer: ignore
+ sap_cumulative_sum_in: wp.array(dtype=int), # kernel_analyzer: ignore
+ # In:
+ nsweep_in: int,
+ # Data out:
+ collision_pair_out: wp.array(dtype=wp.vec2i),
+ collision_hftri_index_out: wp.array(dtype=int),
+ collision_pairid_out: wp.array(dtype=int),
+ collision_worldid_out: wp.array(dtype=int),
+ ncollision_out: wp.array(dtype=int),
+):
+ worldgeomid = wp.tid()
+
+ nworldgeom = nworld_in * ngeom
+ nworkpackages = sap_cumulative_sum_in[nworldgeom - 1]
+
+ while worldgeomid < nworkpackages:
+ # binary search to find current and next geom pair indices
+ i = _binary_search(sap_cumulative_sum_in, worldgeomid, 0, nworldgeom)
+ j = i + worldgeomid + 1
+
+ if i > 0:
+ j -= sap_cumulative_sum_in[i - 1]
+
+ worldid = i // ngeom
+ i = i % ngeom
+ j = j % ngeom
+
+ # get geom indices and swap if necessary
+ geom1 = sap_sort_index_in[worldid, i]
+ geom2 = sap_sort_index_in[worldid, j]
+
+ # find linear index of (geom1, geom2) in upper triangular nxn_pairid
+ if geom2 < geom1:
+ idx = upper_tri_index(ngeom, geom2, geom1)
+ else:
+ idx = upper_tri_index(ngeom, geom1, geom2)
+
+ if nxn_pairid[idx] < -1:
+ worldgeomid += nsweep_in
+ continue
+
+ if _broadphase_filter(
+ opt_broadphase_filter, geom_aabb, geom_rbound, geom_margin, geom_xpos_in, geom_xmat_in, geom1, geom2, worldid
+ ):
+ _add_geom_pair(
+ geom_type,
+ nxn_pairid,
+ nconmax_in,
+ geom1,
+ geom2,
+ worldid,
+ idx,
+ collision_pair_out,
+ collision_hftri_index_out,
+ collision_pairid_out,
+ collision_worldid_out,
+ ncollision_out,
+ )
+
+ worldgeomid += nsweep_in
+
+
+def _segmented_sort(tile_size: int):
+ @wp.kernel
+ def segmented_sort(
+ # Data in:
+ sap_projection_lower_in: wp.array2d(dtype=float), # kernel_analyzer: ignore
+ sap_sort_index_in: wp.array2d(dtype=int), # kernel_analyzer: ignore
+ ):
+ worldid = wp.tid()
+
+ # Load input into shared memory
+ keys = wp.tile_load(sap_projection_lower_in[worldid], shape=tile_size, storage="shared")
+ values = wp.tile_load(sap_sort_index_in[worldid], shape=tile_size, storage="shared")
+
+ # Perform in-place sorting
+ wp.tile_sort(keys, values)
+
+ # Store sorted shared memory into output arrays
+ wp.tile_store(sap_projection_lower_in[worldid], keys)
+ wp.tile_store(sap_sort_index_in[worldid], values)
+
+ return segmented_sort
+
+
+@event_scope
+def sap_broadphase(m: Model, d: Data):
+ """Runs broadphase collision detection using a sweep-and-prune (SAP) algorithm.
+
+ This method is more efficient than the N-squared approach for large numbers of
+ objects. It works by projecting the bounding spheres of all geoms onto a
+ single axis and sorting them. It then sweeps along the axis, only checking
+ for overlaps between geoms whose projections are close to each other.
+
+ For each potentially colliding pair identified by the sweep, a more precise
+ bounding sphere check is performed. If this check passes, the pair is added
+ to the collision arrays in `d` for the narrowphase stage.
+
+ Two sorting strategies are supported, controlled by `m.opt.broadphase`:
+ - `SAP_TILE`: Uses a tile-based sort.
+ - `SAP_SEGMENTED`: Uses a segmented sort.
+ """
+
+ nworldgeom = d.nworld * m.ngeom
+
+ # TODO(team): direction
+
+ # random fixed direction
+ direction = wp.vec3(0.5935, 0.7790, 0.1235)
+ direction = wp.normalize(direction)
+
+ wp.launch(
+ kernel=_sap_project,
+ dim=(d.nworld, m.ngeom),
+ inputs=[
+ m.geom_rbound,
+ m.geom_margin,
+ d.geom_xpos,
+ direction,
+ ],
+ outputs=[
+ d.sap_projection_lower.reshape((-1, m.ngeom)),
+ d.sap_projection_upper,
+ d.sap_sort_index.reshape((-1, m.ngeom)),
+ ],
+ )
+
+ if m.opt.broadphase == int(BroadphaseType.SAP_TILE):
+ wp.launch_tiled(
+ kernel=_segmented_sort(m.ngeom),
+ dim=(d.nworld),
+ inputs=[d.sap_projection_lower.reshape((-1, m.ngeom)), d.sap_sort_index.reshape((-1, m.ngeom))],
+ block_dim=m.block_dim.segmented_sort,
+ )
+ else:
+ wp.utils.segmented_sort_pairs(
+ d.sap_projection_lower.reshape((-1, m.ngeom)),
+ d.sap_sort_index.reshape((-1, m.ngeom)),
+ nworldgeom,
+ d.sap_segment_index.reshape(-1),
+ )
+
+ wp.launch(
+ kernel=_sap_range,
+ dim=(d.nworld, m.ngeom),
+ inputs=[
+ m.ngeom,
+ d.sap_projection_lower.reshape((-1, m.ngeom)),
+ d.sap_projection_upper,
+ d.sap_sort_index.reshape((-1, m.ngeom)),
+ ],
+ outputs=[
+ d.sap_range,
+ ],
+ )
+
+ # scan is used for load balancing among the threads
+ wp.utils.array_scan(d.sap_range.reshape(-1), d.sap_cumulative_sum.reshape(-1), True)
+
+ # estimate number of overlap checks
+ # assumes each geom has 5 other geoms (batched over all worlds)
+ nsweep = 5 * nworldgeom
+ wp.launch(
+ kernel=_sap_broadphase,
+ dim=nsweep,
+ inputs=[
+ m.ngeom,
+ m.opt.broadphase_filter,
+ m.geom_type,
+ m.geom_aabb,
+ m.geom_rbound,
+ m.geom_margin,
+ m.nxn_pairid,
+ d.nworld,
+ d.nconmax,
+ d.geom_xpos,
+ d.geom_xmat,
+ d.sap_sort_index.reshape((-1, m.ngeom)),
+ d.sap_cumulative_sum.reshape(-1),
+ nsweep,
+ ],
+ outputs=[
+ d.collision_pair,
+ d.collision_hftri_index,
+ d.collision_pairid,
+ d.collision_worldid,
+ d.ncollision,
+ ],
+ )
+
+
+@wp.kernel
+def _nxn_broadphase(
+ # Model:
+ opt_broadphase_filter: int,
+ geom_type: wp.array(dtype=int),
+ geom_aabb: wp.array2d(dtype=wp.vec3),
+ geom_rbound: wp.array2d(dtype=float),
+ geom_margin: wp.array2d(dtype=float),
+ nxn_geom_pair: wp.array(dtype=wp.vec2i),
+ nxn_pairid: wp.array(dtype=int),
+ # Data in:
+ nconmax_in: int,
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ # Data out:
+ collision_pair_out: wp.array(dtype=wp.vec2i),
+ collision_hftri_index_out: wp.array(dtype=int),
+ collision_pairid_out: wp.array(dtype=int),
+ collision_worldid_out: wp.array(dtype=int),
+ ncollision_out: wp.array(dtype=int),
+):
+ worldid, elementid = wp.tid()
+
+ geom = nxn_geom_pair[elementid]
+ geom1 = geom[0]
+ geom2 = geom[1]
+
+ if _broadphase_filter(
+ opt_broadphase_filter, geom_aabb, geom_rbound, geom_margin, geom_xpos_in, geom_xmat_in, geom1, geom2, worldid
+ ):
+ _add_geom_pair(
+ geom_type,
+ nxn_pairid,
+ nconmax_in,
+ geom1,
+ geom2,
+ worldid,
+ elementid,
+ collision_pair_out,
+ collision_hftri_index_out,
+ collision_pairid_out,
+ collision_worldid_out,
+ ncollision_out,
+ )
+
+
+@event_scope
+def nxn_broadphase(m: Model, d: Data):
+ """Runs broadphase collision detection using a brute-force N-squared approach.
+
+ This function iterates through a pre-filtered list of all possible geometry pairs and
+ performs a quick bounding sphere check to identify potential collisions.
+
+ For each pair that passes the sphere check, it populates the collision arrays in `d`
+ (`d.collision_pair`, `d.collision_pairid`, etc.), which are then consumed by the
+ narrowphase.
+
+ The initial list of pairs is filtered at model creation time to exclude pairs based on
+ `contype`/`conaffinity`, parent-child relationships, and explicit `` tags.
+ """
+
+ wp.launch(
+ _nxn_broadphase,
+ dim=(d.nworld, m.nxn_geom_pair_filtered.shape[0]),
+ inputs=[
+ m.opt.broadphase_filter,
+ m.geom_type,
+ m.geom_aabb,
+ m.geom_rbound,
+ m.geom_margin,
+ m.nxn_geom_pair_filtered,
+ m.nxn_pairid_filtered,
+ d.nconmax,
+ d.geom_xpos,
+ d.geom_xmat,
+ ],
+ outputs=[
+ d.collision_pair,
+ d.collision_hftri_index,
+ d.collision_pairid,
+ d.collision_worldid,
+ d.ncollision,
+ ],
+ )
+
+
+def _narrowphase(m, d):
+ # Process heightfield collisions
+ if m.nhfield > 0:
+ hfield_midphase(m, d)
+
+ # TODO(team): we should reject far-away contacts in the narrowphase instead of constraint
+ # partitioning because we can move some pressure of the atomics
+ convex_narrowphase(m, d)
+ primitive_narrowphase(m, d)
+
+ if m.has_sdf_geom:
+ sdf_narrowphase(m, d)
+
+
+@event_scope
+def collision(m: Model, d: Data):
+ """Runs the full collision detection pipeline.
+
+ This function orchestrates the broadphase and narrowphase collision detection stages. It
+ first identifies potential collision pairs using a broadphase algorithm (either N-squared
+ or Sweep-and-Prune, based on `m.opt.broadphase`). Then, for each potential pair, it
+ performs narrowphase collision detection to compute detailed contact information like
+ distance, position, and frame.
+
+ The results are used to populate the `d.contact` array, and the total number of contacts
+ is stored in `d.ncon`. If `d.ncon` is larger than `d.nconmax` then an overflow has
+ occurred and the remaining contacts will be skipped. If this happens, raise the `nconmax`
+ parameter in `io.make_data` or `io.put_data`.
+
+ This function will do nothing except zero out arrays if collision detection is disabled
+ via `m.opt.disableflags` or if `d.nconmax` is 0.
+ """
+
+ # zero collision-related arrays
+ wp.launch(
+ _zero_collision_arrays,
+ dim=d.nconmax,
+ inputs=[
+ d.nworld,
+ d.ncon_hfield.shape[1],
+ d.ncon,
+ d.ncon_hfield.reshape(-1),
+ d.collision_hftri_index,
+ d.ncollision,
+ ],
+ )
+
+ if d.nconmax == 0 or m.opt.disableflags & (DisableBit.CONSTRAINT | DisableBit.CONTACT):
+ return
+
+ if m.opt.broadphase == int(BroadphaseType.NXN):
+ nxn_broadphase(m, d)
+ else:
+ sap_broadphase(m, d)
+
+ if m.opt.graph_conditional:
+ wp.capture_if(condition=d.ncollision, on_true=_narrowphase, m=m, d=d)
+ else:
+ _narrowphase(m, d)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py
new file mode 100644
index 00000000..130121d4
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py
@@ -0,0 +1,870 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 the collision driver."""
+
+import mujoco
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+from absl.testing import parameterized
+
+import mujoco_warp as mjwarp
+from mujoco.mjx.third_party.mujoco_warp.test_data.collision_sdf.utils import register_sdf_plugins
+
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+from mujoco.mjx.third_party.mujoco_warp._src import types
+
+
+class CollisionTest(parameterized.TestCase):
+ """Tests the collision contact functions."""
+
+ _SDF_SDF = {
+ "_NUT_NUT": """
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+""",
+ "NUT_BOLT": """
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+""",
+ }
+ _FIXTURES = {
+ "box_plane": """
+
+
+
+
+
+
+
+
+
+ """,
+ "box_box_vf": """
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "box_box_vf_flat": """
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "box_box_ee": """
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "box_box_ee_deep": """
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "plane_sphere": """
+
+
+
+
+
+
+
+
+
+ """,
+ "plane_ellipsoid": """
+
+
+
+
+
+
+
+
+
+ """,
+ "plane_capsule": """
+
+
+
+
+
+
+
+
+
+ """,
+ "convex_convex": """
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "capsule_capsule": """
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "sphere_sphere": """
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "sphere_capsule": """
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "sphere_cylinder_corner": """
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "sphere_cylinder_cap": """
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "sphere_cylinder_side": """
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "plane_cylinder_1": """
+
+
+
+
+
+
+
+
+
+ """,
+ "plane_cylinder_2": """
+
+
+
+
+
+
+
+
+
+ """,
+ "plane_cylinder_3": """
+
+
+
+
+
+
+
+
+
+ """,
+ "mesh_plane_simple": """
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "mesh_plane_complex": """
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ "sphere_box_shallow": """
+
+
+
+
+
+
+
+
+
+ """,
+ "sphere_box_deep": """
+
+
+
+
+
+
+
+
+
+ """,
+ "capsule_box_edge": """
+
+
+
+
+
+
+
+
+
+ """,
+ "capsule_box_corner": """
+
+
+
+
+
+
+
+
+
+ """,
+ "capsule_box_face_tip": """
+
+
+
+
+
+
+
+
+
+ """,
+ "capsule_box_face_flat": """
+
+
+
+
+
+
+
+
+
+ """,
+ }
+
+ @classmethod
+ def setUpClass(cls):
+ register_sdf_plugins(mjwarp._src.collision_sdf)
+
+ @parameterized.parameters(_SDF_SDF.keys())
+ def test_sdf_collision(self, fixture):
+ """Tests collisions with different geometries."""
+
+ mjm, mjd, m, d = test_util.fixture(xml=self._SDF_SDF[fixture], qpos0=True)
+
+ mujoco.mj_collision(mjm, mjd)
+ mjwarp.collision(m, d)
+ for i in range(min(mjd.ncon, d.ncon.numpy()[0])):
+ actual_dist = mjd.contact.dist[i]
+ actual_pos = mjd.contact.pos[i]
+ actual_frame = mjd.contact.frame[i][0:3]
+ result = False
+ test_dist = d.contact.dist.numpy()[i]
+ test_pos = d.contact.pos.numpy()[i, :]
+ test_frame = d.contact.frame.numpy()[i].flatten()[0:3]
+ check_dist = np.allclose(actual_dist, test_dist, rtol=5e-2, atol=1.0e-1)
+ check_frame = np.allclose(actual_frame, test_frame, rtol=5e-2, atol=1.0e-1)
+ check_pos = np.allclose(actual_pos, test_pos, rtol=5e-2, atol=1.0e-1)
+ result = check_dist
+ np.testing.assert_equal(result, True, f"Contact {i} not found in Gjk results")
+
+ @parameterized.parameters(_FIXTURES.keys())
+ def test_collision(self, fixture):
+ """Tests collisions with different geometries."""
+ mjm, mjd, m, d = test_util.fixture(xml=self._FIXTURES[fixture], qpos0=True)
+
+ # Exempt GJK collisions from exact contact count check
+ # because GJK generates more contacts
+ allow_different_contact_count = False
+
+ mujoco.mj_collision(mjm, mjd)
+ mjwarp.collision(m, d)
+
+ self.assertGreater(d.ncon.numpy()[0], 0)
+ self.assertGreater(mjd.ncon, 0)
+
+ for i in range(mjd.ncon):
+ actual_dist = mjd.contact.dist[i]
+ actual_pos = mjd.contact.pos[i]
+ actual_frame = mjd.contact.frame[i]
+ result = False
+ for j in range(d.ncon.numpy()[0]):
+ test_dist = d.contact.dist.numpy()[j]
+ test_pos = d.contact.pos.numpy()[j, :]
+ test_frame = d.contact.frame.numpy()[j].flatten()
+ check_dist = np.allclose(actual_dist, test_dist, rtol=5e-2, atol=1.0e-2)
+ check_pos = np.allclose(actual_pos, test_pos, rtol=5e-2, atol=1.0e-2)
+ check_frame = np.allclose(actual_frame, test_frame, rtol=5e-2, atol=1.0e-2)
+ if check_dist and check_pos and check_frame:
+ result = True
+ break
+ np.testing.assert_equal(result, True, f"Contact {i} not found in Gjk results")
+
+ if not allow_different_contact_count:
+ self.assertEqual(d.ncon.numpy()[0], mjd.ncon)
+
+ _HFIELD_FIXTURES = {
+ "hfield_box": """
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ }
+
+ @parameterized.parameters(_HFIELD_FIXTURES.keys())
+ def test_hfield_collision(self, fixture):
+ """Tests hfield collision with different geometries."""
+ mjm, mjd, m, d = test_util.fixture(xml=self._HFIELD_FIXTURES[fixture])
+
+ mujoco.mj_collision(mjm, mjd)
+ mjwarp.collision(m, d)
+
+ self.assertEqual(mjd.ncon > 0, d.ncon.numpy()[0] > 0, "If MJ collides, MJW should too")
+
+ def test_contact_exclude(self):
+ """Tests contact exclude."""
+ _, _, m, _ = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """
+ )
+ self.assertEqual(m.nxn_geom_pair.numpy().shape[0], 3)
+ np.testing.assert_equal(m.nxn_pairid.numpy(), np.array([-2, -1, -1]))
+
+ def test_contact_pair(self):
+ """Tests contact pair."""
+ # no pairs
+ _, _, m, _ = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+ """
+ )
+ self.assertTrue((m.nxn_pairid.numpy() == -1).all())
+
+ # 1 pair
+ _, _, m, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ qpos0=True,
+ )
+ self.assertTrue((m.nxn_pairid.numpy() == 0).all())
+
+ for arr in (
+ d.ncon,
+ d.contact.includemargin,
+ d.contact.dim,
+ d.contact.friction,
+ d.contact.solref,
+ d.contact.solreffriction,
+ d.contact.solimp,
+ ):
+ arr.zero_()
+
+ mjwarp.collision(m, d)
+
+ self.assertEqual(d.ncon.numpy()[0], 1)
+ self.assertEqual(d.contact.includemargin.numpy()[0], -1)
+ self.assertEqual(d.contact.dim.numpy()[0], 6)
+ np.testing.assert_allclose(d.contact.friction.numpy()[0], np.array([5, 4, 3, 2, 1]))
+ np.testing.assert_allclose(d.contact.solref.numpy()[0], np.array([-0.25, -0.5]))
+ np.testing.assert_allclose(d.contact.solreffriction.numpy()[0], np.array([2.0, 4.0]))
+ np.testing.assert_allclose(d.contact.solimp.numpy()[0], np.array([0.1, 0.2, 0.3, 0.4, 0.5]))
+
+ # 1 pair: override contype and conaffinity
+ _, _, m, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ qpos0=True,
+ )
+ self.assertTrue((m.nxn_pairid.numpy() == 0).all())
+
+ for arr in (
+ d.ncon,
+ d.contact.includemargin,
+ d.contact.dim,
+ d.contact.friction,
+ d.contact.solref,
+ d.contact.solreffriction,
+ d.contact.solimp,
+ ):
+ arr.zero_()
+
+ mjwarp.collision(m, d)
+
+ self.assertEqual(d.ncon.numpy()[0], 1)
+ self.assertEqual(d.contact.includemargin.numpy()[0], -1)
+ self.assertEqual(d.contact.dim.numpy()[0], 6)
+ np.testing.assert_allclose(d.contact.friction.numpy()[0], np.array([5, 4, 3, 2, 1]))
+ np.testing.assert_allclose(d.contact.solref.numpy()[0], np.array([-0.25, -0.5]))
+ np.testing.assert_allclose(d.contact.solreffriction.numpy()[0], np.array([2.0, 4.0]))
+ np.testing.assert_allclose(d.contact.solimp.numpy()[0], np.array([0.1, 0.2, 0.3, 0.4, 0.5]))
+
+ # 1 pair: override exclude
+ _, _, m, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ qpos0=True,
+ )
+ self.assertTrue((m.nxn_pairid.numpy() == 0).all())
+
+ for arr in (
+ d.ncon,
+ d.contact.includemargin,
+ d.contact.dim,
+ d.contact.friction,
+ d.contact.solref,
+ d.contact.solreffriction,
+ d.contact.solimp,
+ ):
+ arr.zero_()
+
+ mjwarp.collision(m, d)
+
+ self.assertEqual(d.ncon.numpy()[0], 1)
+ self.assertEqual(d.contact.includemargin.numpy()[0], -1)
+ self.assertEqual(d.contact.dim.numpy()[0], 6)
+ np.testing.assert_allclose(d.contact.friction.numpy()[0], np.array([5, 4, 3, 2, 1]))
+ np.testing.assert_allclose(d.contact.solref.numpy()[0], np.array([-0.25, -0.5]))
+ np.testing.assert_allclose(d.contact.solreffriction.numpy()[0], np.array([2.0, 4.0]))
+ np.testing.assert_allclose(d.contact.solimp.numpy()[0], np.array([0.1, 0.2, 0.3, 0.4, 0.5]))
+
+ # 1 pair 1 exclude
+ _, _, m, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ qpos0=True,
+ )
+ np.testing.assert_equal(m.nxn_pairid.numpy(), np.array([-2, -1, 0]))
+
+ for arr in (
+ d.ncon,
+ d.contact.includemargin,
+ d.contact.dim,
+ d.contact.friction,
+ d.contact.solref,
+ d.contact.solreffriction,
+ d.contact.solimp,
+ ):
+ arr.zero_()
+
+ mjwarp.collision(m, d)
+
+ self.assertEqual(d.ncon.numpy()[0], 2)
+ self.assertEqual(d.contact.includemargin.numpy()[1], -1)
+ self.assertEqual(d.contact.dim.numpy()[1], 6)
+ np.testing.assert_allclose(d.contact.friction.numpy()[1], np.array([5, 4, 3, 2, 1]))
+ np.testing.assert_allclose(d.contact.solref.numpy()[1], np.array([-0.25, -0.5]))
+ np.testing.assert_allclose(d.contact.solreffriction.numpy()[1], np.array([2.0, 4.0]))
+ np.testing.assert_allclose(d.contact.solimp.numpy()[1], np.array([0.1, 0.2, 0.3, 0.4, 0.5]))
+
+ # TODO(team): test sap_broadphase
+
+ @parameterized.parameters(
+ (True, True),
+ (True, False),
+ (False, True),
+ (False, False),
+ )
+ def test_collision_disableflags(self, constraint, contact):
+ """Tests collision disableflags."""
+ mjm, mjd, m, d = test_util.fixture(
+ "humanoid/humanoid.xml",
+ keyframe=0,
+ constraint=constraint,
+ contact=contact,
+ kick=False,
+ )
+
+ mujoco.mj_collision(mjm, mjd)
+ mjwarp.collision(m, d)
+
+ self.assertEqual(d.ncon.numpy()[0], mjd.ncon)
+
+ def test_hfield_maxconpair(self):
+ _XML = f"""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """
+
+ _, _, m, d = test_util.fixture(xml=_XML, keyframe=0)
+
+ mjwarp.collision(m, d)
+
+ np.testing.assert_equal(d.ncon.numpy()[0], types.MJ_MAXCONPAIR)
+
+ def test_min_friction(self):
+ _, _, _, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ keyframe=0,
+ )
+
+ self.assertEqual(d.ncon.numpy()[0], 1)
+ np.testing.assert_allclose(d.contact.friction.numpy()[0], types.MJ_MINMU)
+
+ # TODO(team): test contact parameter mixing
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py
new file mode 100644
index 00000000..8128262a
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py
@@ -0,0 +1,1361 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_prism_vertex
+from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import Geom
+from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
+from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
+
+# TODO(team): improve compile time to enable backward pass
+wp.config.enable_backward = False
+
+FLOAT_MIN = -1e30
+FLOAT_MAX = 1e30
+MJ_MINVAL2 = MJ_MINVAL * MJ_MINVAL
+
+mat43 = wp.types.matrix(shape=(4, 3), dtype=float)
+
+
+@wp.struct
+class GJKResult:
+ dist: float
+ x1: wp.vec3
+ x2: wp.vec3
+ dim: int
+ simplex: mat43
+ simplex1: mat43
+ simplex2: mat43
+ simplex_index1: wp.vec4i
+ simplex_index2: wp.vec4i
+
+
+@wp.struct
+class Polytope:
+ status: int
+
+ # vertices in polytope
+ vert: wp.array(dtype=wp.vec3)
+ vert1: wp.array(dtype=wp.vec3)
+ vert2: wp.array(dtype=wp.vec3)
+ vert_index1: wp.array(dtype=int)
+ vert_index2: wp.array(dtype=int)
+ nvert: int
+
+ # faces in polytope
+ face: wp.array(dtype=wp.vec3i)
+ face_pr: wp.array(dtype=wp.vec3)
+ face_norm2: wp.array(dtype=float)
+ face_index: wp.array(dtype=int)
+ nface: int
+
+ # TODO(kbayes): look into if a linear map actually improves performance
+ face_map: wp.array(dtype=int)
+ nmap: int
+
+ # edges that make up the horizon when adding new vertices to polytope
+ horizon: wp.array(dtype=int)
+ nhorizon: int
+
+
+@wp.func
+def _support(geom: Geom, geomtype: int, dir: wp.vec3):
+ cached_index = -1
+ vertex_index = -1
+ local_dir = wp.transpose(geom.rot) @ dir
+ if geomtype == int(GeomType.SPHERE.value):
+ support_pt = geom.pos + geom.size[0] * dir
+ elif geomtype == int(GeomType.BOX.value):
+ tmp = wp.sign(local_dir)
+ res = wp.cw_mul(tmp, geom.size)
+ support_pt = geom.rot @ res + geom.pos
+ vertex_index = 0
+ if tmp[0] > 0:
+ vertex_index += 1
+ if tmp[1] > 0:
+ vertex_index += 2
+ if tmp[2] > 0:
+ vertex_index += 4
+ elif geomtype == int(GeomType.CAPSULE.value):
+ res = local_dir * geom.size[0]
+ # add cylinder contribution
+ res[2] += wp.sign(local_dir[2]) * geom.size[1]
+ support_pt = geom.rot @ res + geom.pos
+ elif geomtype == int(GeomType.ELLIPSOID.value):
+ res = wp.cw_mul(local_dir, geom.size)
+ res = wp.normalize(res)
+ # transform to ellipsoid
+ res = wp.cw_mul(res, geom.size)
+ support_pt = geom.rot @ res + geom.pos
+ elif geomtype == int(GeomType.CYLINDER.value):
+ res = wp.vec3(0.0, 0.0, 0.0)
+ # set result in XY plane: support on circle
+ d = wp.sqrt(local_dir[0] * local_dir[0] + local_dir[1] * local_dir[1])
+ if d > MJ_MINVAL:
+ scl = geom.size[0] / d
+ res[0] = local_dir[0] * scl
+ res[1] = local_dir[1] * scl
+ # set result in Z direction
+ res[2] = wp.sign(local_dir[2]) * geom.size[1]
+ support_pt = geom.rot @ res + geom.pos
+ elif geomtype == int(GeomType.MESH.value):
+ max_dist = float(FLOAT_MIN)
+ if geom.graphadr == -1 or geom.vertnum < 10:
+ if geom.index > -1:
+ cached_index = geom.index
+ max_dist = wp.dot(geom.vert[geom.index], local_dir)
+ support_pt = geom.vert[geom.index]
+ # exhaustive search over all vertices
+ for i in range(geom.vertnum):
+ vert = geom.vert[geom.vertadr + i]
+ dist = wp.dot(vert, local_dir)
+ if dist > max_dist:
+ max_dist = dist
+ support_pt = vert
+ cached_index = geom.vertadr + i
+ vertex_index = cached_index - geom.vertadr
+ else:
+ numvert = geom.graph[geom.graphadr]
+ vert_edgeadr = geom.graphadr + 2
+ vert_globalid = geom.graphadr + 2 + numvert
+ edge_localid = geom.graphadr + 2 + 2 * numvert
+ # hillclimb until no change
+ prev = int(-1)
+ imax = int(0)
+ if geom.index > -1:
+ imax = geom.index
+ cached_index = geom.index
+
+ while True:
+ prev = int(imax)
+ i = int(geom.graph[vert_edgeadr + imax])
+ while geom.graph[edge_localid + i] >= 0:
+ subidx = geom.graph[edge_localid + i]
+ idx = geom.graph[vert_globalid + subidx]
+ dist = wp.dot(local_dir, geom.vert[geom.vertadr + idx])
+ if dist > max_dist:
+ max_dist = dist
+ imax = int(subidx)
+ i += int(1)
+ if imax == prev:
+ break
+ cached_index = imax
+ imax = geom.graph[vert_globalid + imax]
+ vertex_index = imax
+ support_pt = geom.vert[geom.vertadr + imax]
+
+ support_pt = geom.rot @ support_pt + geom.pos
+ elif geomtype == int(GeomType.HFIELD.value):
+ max_dist = float(FLOAT_MIN)
+ for i in range(6):
+ vert = hfield_prism_vertex(geom.hfprism, i)
+ dist = wp.dot(vert, local_dir)
+ if dist > max_dist:
+ max_dist = dist
+ support_pt = vert
+ support_pt = geom.rot @ support_pt + geom.pos
+
+ return support_pt, cached_index, vertex_index
+
+
+@wp.func
+def _attach_face(pt: Polytope, idx: int, v1: int, v2: int, v3: int):
+ # out of memory, returning 0 will force EPA to return early without contact
+ if pt.nface == pt.face.shape[0]:
+ return 0.0
+
+ # compute witness point v
+ r, ret = _project_origin_plane(pt.vert[v3], pt.vert[v2], pt.vert[v1])
+ if ret:
+ return 0.0
+
+ face = wp.vec3i(v1, v2, v3)
+ pt.face[idx] = face
+ pt.face_pr[idx] = r
+
+ pt.face_norm2[idx] = wp.dot(r, r)
+ pt.face_index[idx] = -1
+ return pt.face_norm2[idx]
+
+
+@wp.func
+def _epa_support(pt: Polytope, idx: int, geom1: Geom, geom2: Geom, geom1_type: int, geom2_type: int, dir: wp.vec3):
+ s1, index1, vertex_index1 = _support(geom1, geom1_type, dir)
+ s2, index2, vertex_index2 = _support(geom2, geom2_type, -dir)
+
+ pt.vert[idx] = s1 - s2
+ pt.vert1[idx] = s1
+ pt.vert2[idx] = s2
+ pt.vert_index1[idx] = vertex_index1
+ pt.vert_index2[idx] = vertex_index2
+ return index1, index2
+
+
+@wp.func
+def _linear_combine(n: int, coefs: wp.vec4, mat: mat43):
+ v = wp.vec3(0.0)
+ if n == 1:
+ v = coefs[0] * mat[0]
+ elif n == 2:
+ v = coefs[0] * mat[0] + coefs[1] * mat[1]
+ elif n == 3:
+ v = coefs[0] * mat[0] + coefs[1] * mat[1] + coefs[2] * mat[2]
+ else:
+ v = coefs[0] * mat[0] + coefs[1] * mat[1] + coefs[2] * mat[2] + coefs[3] * mat[3]
+ return v
+
+
+@wp.func
+def _almost_equal(v1: wp.vec3, v2: wp.vec3):
+ return wp.abs(v1[0] - v2[0]) < MJ_MINVAL and wp.abs(v1[1] - v2[1]) < MJ_MINVAL and wp.abs(v1[2] - v2[2]) < MJ_MINVAL
+
+
+@wp.func
+def _subdistance(n: int, simplex: mat43):
+ if n == 4:
+ return _S3D(simplex[0], simplex[1], simplex[2], simplex[3])
+ if n == 3:
+ coordinates3 = _S2D(simplex[0], simplex[1], simplex[2])
+ return wp.vec4(coordinates3[0], coordinates3[1], coordinates3[2], 0.0)
+ if n == 2:
+ coordinates2 = _S1D(simplex[0], simplex[1])
+ return wp.vec4(coordinates2[0], coordinates2[1], 0.0, 0.0)
+ return wp.vec4(1.0, 0.0, 0.0, 0.0)
+
+
+@wp.func
+def _det3(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3):
+ return wp.dot(v1, wp.cross(v2, v3))
+
+
+@wp.func
+def _same_sign(a: float, b: float):
+ if a > 0 and b > 0:
+ return 1
+ if a < 0 and b < 0:
+ return -1
+ return 0
+
+
+@wp.func
+def _project_origin_line(v1: wp.vec3, v2: wp.vec3):
+ diff = v2 - v1
+ scl = -(wp.dot(v2, diff) / wp.dot(diff, diff))
+ return v2 + scl * diff
+
+
+@wp.func
+def _project_origin_plane(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3):
+ z = wp.vec3(0.0)
+ diff21 = v2 - v1
+ diff31 = v3 - v1
+ diff32 = v3 - v2
+
+ # n = (v1 - v2) x (v3 - v2)
+ n = wp.cross(diff32, diff21)
+ nv = wp.dot(n, v2)
+ nn = wp.dot(n, n)
+ if nn == 0:
+ return z, 1
+ if nv != 0 and nn > MJ_MINVAL:
+ v = (nv / nn) * n
+ return v, 0
+
+ # n = (v2 - v1) x (v3 - v1)
+ n = wp.cross(diff21, diff31)
+ nv = wp.dot(n, v1)
+ nn = wp.dot(n, n)
+ if nn == 0:
+ return z, 1
+ if nv != 0 and nn > MJ_MINVAL:
+ v = (nv / nn) * n
+ return v, 0
+
+ # n = (v1 - v3) x (v2 - v3)
+ n = wp.cross(diff31, diff32)
+ nv = wp.dot(n, v3)
+ nn = wp.dot(n, n)
+ v = (nv / nn) * n
+ return v, 0
+
+
+@wp.func
+def _S3D(s1: wp.vec3, s2: wp.vec3, s3: wp.vec3, s4: wp.vec3):
+ # [[ s1_x, s2_x, s3_x, s4_x ],
+ # [ s1_y, s2_y, s3_y, s4_y ],
+ # [ s1_z, s2_z, s3_z, s4_z ],
+ # [ 1, 1, 1, 1 ]]
+ # we want to solve M*lambda = P, where P = [p_x, p_y, p_z, 1] with [p_x, p_y, p_z] is the
+ # origin projected onto the simplex
+
+ # compute cofactors to find det(M)
+ C41 = -_det3(s2, s3, s4)
+ C42 = _det3(s1, s3, s4)
+ C43 = -_det3(s1, s2, s4)
+ C44 = _det3(s1, s2, s3)
+
+ # NOTE: m_det = 6*SignVol(simplex) with C4i corresponding to the volume of the 3-simplex
+ # with vertices {s1, s2, s3, 0} - si
+ m_det = C41 + C42 + C43 + C44
+
+ comp1 = _same_sign(m_det, C41)
+ comp2 = _same_sign(m_det, C42)
+ comp3 = _same_sign(m_det, C43)
+ comp4 = _same_sign(m_det, C44)
+
+ # if all signs are the same then the origin is inside the simplex
+ if comp1 and comp2 and comp3 and comp4:
+ return wp.vec4(C41 / m_det, C42 / m_det, C43 / m_det, C44 / m_det)
+
+ # find the smallest distance, and use the corresponding barycentric coordinates
+ coordinates = wp.vec4(0.0, 0.0, 0.0, 0.0)
+ dmin = FLOAT_MAX
+
+ if not comp1:
+ subcoord = _S2D(s2, s3, s4)
+ x = subcoord[0] * s2 + subcoord[1] * s3 + subcoord[2] * s4
+ d = wp.dot(x, x)
+ coordinates[0] = 0.0
+ coordinates[1] = subcoord[0]
+ coordinates[2] = subcoord[1]
+ coordinates[3] = subcoord[2]
+ dmin = d
+
+ if not comp2:
+ subcoord = _S2D(s1, s3, s4)
+ x = subcoord[0] * s1 + subcoord[1] * s3 + subcoord[2] * s4
+ d = wp.dot(x, x)
+ if d < dmin:
+ coordinates[0] = subcoord[0]
+ coordinates[1] = 0.0
+ coordinates[2] = subcoord[1]
+ coordinates[3] = subcoord[2]
+ dmin = d
+
+ if not comp3:
+ subcoord = _S2D(s1, s2, s4)
+ x = subcoord[0] * s1 + subcoord[1] * s2 + subcoord[2] * s4
+ d = wp.dot(x, x)
+ if d < dmin:
+ coordinates[0] = subcoord[0]
+ coordinates[1] = subcoord[1]
+ coordinates[2] = 0.0
+ coordinates[3] = subcoord[2]
+ dmin = d
+
+ if not comp4:
+ subcoord = _S2D(s1, s2, s3)
+ x = subcoord[0] * s1 + subcoord[1] * s2 + subcoord[2] * s3
+ d = wp.dot(x, x)
+ if d < dmin:
+ coordinates[0] = subcoord[0]
+ coordinates[1] = subcoord[1]
+ coordinates[2] = subcoord[2]
+ coordinates[3] = 0.0
+ return coordinates
+
+
+@wp.func
+def _S2D(s1: wp.vec3, s2: wp.vec3, s3: wp.vec3):
+ # project origin onto affine hull of the simplex
+ p_o, ret = _project_origin_plane(s1, s2, s3)
+ if ret:
+ v = _S1D(s1, s2)
+ return wp.vec3(v[0], v[1], 0.0)
+
+ # Below are the minors M_i4 of the matrix M given by
+ # [[ s1_x, s2_x, s3_x, s4_x ],
+ # [ s1_y, s2_y, s3_y, s4_y ],
+ # [ s1_z, s2_z, s3_z, s4_z ],
+ # [ 1, 1, 1, 1 ]]
+ M_14 = s2[1] * s3[2] - s2[2] * s3[1] - s1[1] * s3[2] + s1[2] * s3[1] + s1[1] * s2[2] - s1[2] * s2[1]
+ M_24 = s2[0] * s3[2] - s2[2] * s3[0] - s1[0] * s3[2] + s1[2] * s3[0] + s1[0] * s2[2] - s1[2] * s2[0]
+ M_34 = s2[0] * s3[1] - s2[1] * s3[0] - s1[0] * s3[1] + s1[1] * s3[0] + s1[0] * s2[1] - s1[1] * s2[0]
+
+ # exclude the axis with the largest projection of the simplex using the computed minors
+ M_max = 0.0
+ s1_2D = wp.vec2(0.0)
+ s2_2D = wp.vec2(0.0)
+ s3_2D = wp.vec2(0.0)
+ p_o_2D = wp.vec2(0.0)
+
+ mu1 = wp.abs(M_14)
+ mu2 = wp.abs(M_24)
+ mu3 = wp.abs(M_34)
+
+ if mu1 >= mu2 and mu1 >= mu3:
+ M_max = M_14
+ s1_2D[0] = s1[1]
+ s1_2D[1] = s1[2]
+
+ s2_2D[0] = s2[1]
+ s2_2D[1] = s2[2]
+
+ s3_2D[0] = s3[1]
+ s3_2D[1] = s3[2]
+
+ p_o_2D[0] = p_o[1]
+ p_o_2D[1] = p_o[2]
+ elif mu2 >= mu3:
+ M_max = M_24
+ s1_2D[0] = s1[0]
+ s1_2D[1] = s1[2]
+
+ s2_2D[0] = s2[0]
+ s2_2D[1] = s2[2]
+
+ s3_2D[0] = s3[0]
+ s3_2D[1] = s3[2]
+
+ p_o_2D[0] = p_o[0]
+ p_o_2D[1] = p_o[2]
+ else:
+ M_max = M_34
+ s1_2D[0] = s1[0]
+ s1_2D[1] = s1[1]
+
+ s2_2D[0] = s2[0]
+ s2_2D[1] = s2[1]
+
+ s3_2D[0] = s3[0]
+ s3_2D[1] = s3[1]
+
+ p_o_2D[0] = p_o[0]
+ p_o_2D[1] = p_o[1]
+
+ # compute the cofactors C3i of the following matrix:
+ # [[ s1_2D[0] - p_o_2D[0], s2_2D[0] - p_o_2D[0], s3_2D[0] - p_o_2D[0] ],
+ # [ s1_2D[1] - p_o_2D[1], s2_2D[1] - p_o_2D[1], s3_2D[1] - p_o_2D[1] ],
+ # [ 1, 1, 1 ]]
+
+ # C31 corresponds to the signed area of 2-simplex: (p_o_2D, s2_2D, s3_2D)
+ C31 = (
+ p_o_2D[0] * s2_2D[1]
+ + p_o_2D[1] * s3_2D[0]
+ + s2_2D[0] * s3_2D[1]
+ - p_o_2D[0] * s3_2D[1]
+ - p_o_2D[1] * s2_2D[0]
+ - s3_2D[0] * s2_2D[1]
+ )
+
+ # C32 corresponds to the signed area of 2-simplex: (_po_2D, s1_2D, s3_2D)
+ C32 = (
+ p_o_2D[0] * s3_2D[1]
+ + p_o_2D[1] * s1_2D[0]
+ + s3_2D[0] * s1_2D[1]
+ - p_o_2D[0] * s1_2D[1]
+ - p_o_2D[1] * s3_2D[0]
+ - s1_2D[0] * s3_2D[1]
+ )
+
+ # C33 corresponds to the signed area of 2-simplex: (p_o_2D, s1_2D, s2_2D)
+ C33 = (
+ p_o_2D[0] * s1_2D[1]
+ + p_o_2D[1] * s2_2D[0]
+ + s1_2D[0] * s2_2D[1]
+ - p_o_2D[0] * s2_2D[1]
+ - p_o_2D[1] * s1_2D[0]
+ - s2_2D[0] * s1_2D[1]
+ )
+
+ comp1 = _same_sign(M_max, C31)
+ comp2 = _same_sign(M_max, C32)
+ comp3 = _same_sign(M_max, C33)
+
+ # all the same sign, p_o is inside the 2-simplex
+ if comp1 and comp2 and comp3:
+ return wp.vec3(C31 / M_max, C32 / M_max, C33 / M_max)
+
+ # find the smallest distance, and use the corresponding barycentric coordinates
+ dmin = FLOAT_MAX
+ coordinates = wp.vec3(0.0, 0.0, 0.0)
+
+ if not comp1:
+ subcoord = _S1D(s2, s3)
+ x = subcoord[0] * s2 + subcoord[1] * s3
+ d = wp.dot(x, x)
+ coordinates[0] = 0.0
+ coordinates[1] = subcoord[0]
+ coordinates[2] = subcoord[1]
+ dmin = d
+
+ if not comp2:
+ subcoord = _S1D(s1, s3)
+ x = subcoord[0] * s1 + subcoord[1] * s3
+ d = wp.dot(x, x)
+ if d < dmin:
+ coordinates[0] = subcoord[0]
+ coordinates[1] = 0.0
+ coordinates[2] = subcoord[1]
+ dmin = d
+
+ if not comp3:
+ subcoord = _S1D(s1, s2)
+ x = subcoord[0] * s1 + subcoord[1] * s2
+ d = wp.dot(x, x)
+ if d < dmin:
+ coordinates[0] = subcoord[0]
+ coordinates[1] = subcoord[1]
+ coordinates[2] = 0.0
+ return coordinates
+
+
+@wp.func
+def _S1D(s1: wp.vec3, s2: wp.vec3):
+ # find projection of origin onto the 1-simplex:
+ p_o = _project_origin_line(s1, s2)
+
+ # find the axis with the largest projection "shadow" of the simplex
+ mu_max = 0.0
+ index = 0
+ for i in range(3):
+ mu = s1[i] - s2[i]
+ if wp.abs(mu) >= wp.abs(mu_max):
+ mu_max = mu
+ index = i
+
+ C1 = p_o[index] - s2[index]
+ C2 = s1[index] - p_o[index]
+
+ # inside the simplex
+ if _same_sign(mu_max, C1) and _same_sign(mu_max, C2):
+ return wp.vec2(C1 / mu_max, C2 / mu_max)
+ return wp.vec2(0.0, 1.0)
+
+
+@wp.func
+def _gjk(
+ # In:
+ tolerance: float,
+ gjk_iterations: int,
+ geom1: Geom,
+ geom2: Geom,
+ x1_0: wp.vec3,
+ x2_0: wp.vec3,
+ geomtype1: int,
+ geomtype2: int,
+ cutoff: float,
+):
+ """Find distance within a tolerance between two geoms."""
+ cutoff2 = cutoff * cutoff
+ simplex = mat43()
+ simplex1 = mat43()
+ simplex2 = mat43()
+ simplex_index1 = wp.vec4i()
+ simplex_index2 = wp.vec4i()
+ n = int(0)
+ coordinates = wp.vec4() # barycentric coordinates
+ epsilon = 0.5 * tolerance * tolerance
+
+ # set initial guess
+ x_k = x1_0 - x2_0
+
+ for k in range(gjk_iterations):
+ xnorm = wp.dot(x_k, x_k)
+ # TODO(kbayes): determine new constant here
+ if xnorm < 1e-12:
+ break
+ dir_neg = x_k / wp.sqrt(xnorm)
+
+ # compute the kth support point
+ s1_k, i1, vertex_index1 = _support(geom1, geomtype1, -dir_neg)
+ s2_k, i2, vertex_index2 = _support(geom2, geomtype2, dir_neg)
+ geom1.index = i1
+ geom2.index = i2
+ simplex1[n] = s1_k
+ simplex2[n] = s2_k
+ simplex_index1[n] = vertex_index1
+ simplex_index2[n] = vertex_index2
+ simplex[n] = s1_k - s2_k
+
+ if cutoff == 0.0:
+ if wp.dot(x_k, simplex[n]) > 0:
+ result = GJKResult()
+ result.dim = 0
+ result.dist = FLOAT_MAX
+ return result
+ elif cutoff < FLOAT_MAX:
+ vs = wp.dot(x_k, simplex[n])
+ vv = wp.dot(x_k, x_k)
+ if wp.dot(x_k, simplex[n]) > 0 and (vs * vs / vv) >= cutoff2:
+ result = GJKResult()
+ result.dim = 0
+ result.dist = FLOAT_MAX
+ return result
+
+ # stopping criteria using the Frank-Wolfe duality gap given by
+ # |f(x_k) - f(x_min)|^2 <= < grad f(x_k), (x_k - simplex[n]) >
+ if wp.dot(x_k, x_k - simplex[n]) < epsilon:
+ break
+
+ # run the distance subalgorithm to compute the barycentric coordinates
+ # of the closest point to the origin in the simplex
+ coordinates = _subdistance(n + 1, simplex)
+
+ # remove vertices from the simplex no longer needed
+ n = int(0)
+ for i in range(4):
+ if coordinates[i] == 0:
+ continue
+
+ simplex[n] = simplex[i]
+ simplex1[n] = simplex1[i]
+ simplex2[n] = simplex2[i]
+ simplex_index1[n] = simplex_index1[i]
+ simplex_index2[n] = simplex_index2[i]
+ coordinates[n] = coordinates[i]
+ n += int(1)
+
+ # SHOULD NOT OCCUR
+ if n < 1:
+ break
+
+ # get the next iteration of x_k
+ x_next = _linear_combine(n, coordinates, simplex)
+
+ # x_k has converged to minimum
+ if _almost_equal(x_next, x_k):
+ break
+
+ # copy next iteration into x_k
+ x_k = x_next
+
+ # we have a tetrahedron containing the origin so return early
+ if n == 4:
+ break
+
+ result = GJKResult()
+
+ # compute the approximate witness points
+ result.x1 = _linear_combine(n, coordinates, simplex1)
+ result.x2 = _linear_combine(n, coordinates, simplex2)
+ result.dist = wp.norm_l2(x_k)
+
+ result.dim = n
+ result.simplex1 = simplex1
+ result.simplex2 = simplex2
+ result.simplex_index1 = simplex_index1
+ result.simplex_index2 = simplex_index2
+ result.simplex = simplex
+ return result
+
+
+@wp.func
+def _same_side(p0: wp.vec3, p1: wp.vec3, p2: wp.vec3, p3: wp.vec3):
+ n = wp.cross(p1 - p0, p2 - p0)
+ dot1 = wp.dot(n, p3 - p0)
+ dot2 = wp.dot(n, -p0)
+ if dot1 > 0 and dot2 > 0:
+ return 1
+ if dot1 < 0 and dot2 < 0:
+ return 1
+ return 0
+
+
+@wp.func
+def _test_tetra(p0: wp.vec3, p1: wp.vec3, p2: wp.vec3, p3: wp.vec3):
+ return _same_side(p0, p1, p2, p3) and _same_side(p1, p2, p3, p0) and _same_side(p2, p3, p0, p1) and _same_side(p3, p0, p1, p2)
+
+
+@wp.func
+def _tri_affine_coord(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, p: wp.vec3):
+ # compute minors as in S2D
+ M_14 = v2[1] * v3[2] - v2[2] * v3[1] - v1[1] * v3[2] + v1[2] * v3[1] + v1[1] * v2[2] - v1[2] * v2[1]
+ M_24 = v2[0] * v3[2] - v2[2] * v3[0] - v1[0] * v3[2] + v1[2] * v3[0] + v1[0] * v2[2] - v1[2] * v2[0]
+ M_34 = v2[0] * v3[1] - v2[1] * v3[0] - v1[0] * v3[1] + v1[1] * v3[0] + v1[0] * v2[1] - v1[1] * v2[0]
+
+ # exclude one of the axes with the largest projection
+ # of the simplex using the computed minors
+ M_max = 0.0
+ x = 0
+ y = 0
+
+ mu1 = wp.abs(M_14)
+ mu2 = wp.abs(M_24)
+ mu3 = wp.abs(M_34)
+
+ if mu1 >= mu2 and mu1 >= mu3:
+ M_max = M_14
+ x = 1
+ y = 2
+ elif mu2 >= mu3:
+ M_max = M_24
+ x = 0
+ y = 2
+ else:
+ M_max = M_34
+ x = 0
+ y = 1
+
+ # C31 corresponds to the signed area of 2-simplex: (v, s2, s3)
+ C31 = p[x] * v2[y] + p[y] * v3[x] + v2[x] * v3[y] - p[x] * v3[y] - p[y] * v2[x] - v3[x] * v2[y]
+
+ # C32 corresponds to the signed area of 2-simplex: (v, s1, s3)
+ C32 = p[x] * v3[y] + p[y] * v1[x] + v3[x] * v1[y] - p[x] * v1[y] - p[y] * v3[x] - v1[x] * v3[y]
+
+ # C33 corresponds to the signed area of 2-simplex: (v, s1, s2)
+ C33 = p[x] * v1[y] + p[y] * v2[x] + v1[x] * v2[y] - p[x] * v2[y] - p[y] * v1[x] - v2[x] * v1[y]
+
+ # compute affine coordinates
+ return wp.vec3(C31 / M_max, C32 / M_max, C33 / M_max)
+
+
+@wp.func
+def _tri_point_intersect(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, p: wp.vec3):
+ coordinates = _tri_affine_coord(v1, v2, v3, p)
+ l1 = coordinates[0]
+ l2 = coordinates[1]
+ l3 = coordinates[2]
+
+ if l1 < 0 or l2 < 0 or l3 < 0:
+ return False
+
+ pr = wp.vec3()
+ pr[0] = v1[0] * l1 + v2[0] * l2 + v3[0] * l3
+ pr[1] = v1[1] * l1 + v2[1] * l2 + v3[1] * l3
+ pr[2] = v1[2] * l1 + v2[2] * l2 + v3[2] * l3
+ return wp.norm_l2(pr - p) < MJ_MINVAL
+
+
+@wp.func
+def _replace_simplex3(pt: Polytope, v1: int, v2: int, v3: int):
+ result = GJKResult()
+
+ # reset GJK simplex
+ simplex = mat43()
+ simplex[0] = pt.vert[v1]
+ simplex[1] = pt.vert[v2]
+ simplex[2] = pt.vert[v3]
+
+ simplex1 = mat43()
+ simplex1[0] = pt.vert1[v1]
+ simplex1[1] = pt.vert1[v2]
+ simplex1[2] = pt.vert1[v3]
+
+ simplex2 = mat43()
+ simplex2[0] = pt.vert2[v1]
+ simplex2[1] = pt.vert2[v2]
+ simplex2[2] = pt.vert2[v3]
+
+ simplex_index1 = wp.vec4i()
+ simplex_index1[0] = pt.vert_index1[v1]
+ simplex_index1[1] = pt.vert_index1[v2]
+ simplex_index1[2] = pt.vert_index1[v3]
+
+ simplex_index2 = wp.vec4i()
+ simplex_index2[0] = pt.vert_index2[v1]
+ simplex_index2[1] = pt.vert_index2[v2]
+ simplex_index2[2] = pt.vert_index2[v3]
+
+ result.simplex = simplex
+ result.simplex1 = simplex1
+ result.simplex2 = simplex2
+ result.simplex_index1 = simplex_index1
+ result.simplex_index2 = simplex_index2
+
+ return result
+
+
+@wp.func
+def _rotmat(axis: wp.vec3):
+ n = wp.norm_l2(axis)
+ u1 = axis[0] / n
+ u2 = axis[1] / n
+ u3 = axis[2] / n
+
+ sin = 0.86602540378 # sin(120 deg)
+ cos = -0.5 # cos(120 deg)
+ R = wp.mat33()
+ R[0, 0] = cos + u1 * u1 * (1.0 - cos)
+ R[0, 1] = u1 * u2 * (1.0 - cos) - u3 * sin
+ R[0, 2] = u1 * u3 * (1.0 - cos) + u2 * sin
+ R[1, 0] = u2 * u1 * (1.0 - cos) + u3 * sin
+ R[1, 1] = cos + u2 * u2 * (1.0 - cos)
+ R[1, 2] = u2 * u3 * (1.0 - cos) - u1 * sin
+ R[2, 0] = u1 * u3 * (1.0 - cos) - u2 * sin
+ R[2, 1] = u2 * u3 * (1.0 - cos) + u1 * sin
+ R[2, 2] = cos + u3 * u3 * (1.0 - cos)
+ return R
+
+
+@wp.func
+def _ray_triangle(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, v4: wp.vec3, v5: wp.vec3):
+ vol1 = _det3(v3 - v1, v4 - v1, v2 - v1)
+ vol2 = _det3(v4 - v1, v5 - v1, v2 - v1)
+ vol3 = _det3(v5 - v1, v3 - v1, v2 - v1)
+
+ if vol1 >= 0 and vol2 >= 0 and vol3 >= 0:
+ return 1
+ if vol1 <= 0 and vol2 <= 0 and vol3 <= 0:
+ return -1
+ return 0
+
+
+@wp.func
+def _add_edge(pt: Polytope, e1: int, e2: int):
+ n = pt.nhorizon
+
+ if n < 0:
+ return -1
+
+ for i in range(n):
+ old_e1 = pt.horizon[2 * i + 0]
+ old_e2 = pt.horizon[2 * i + 1]
+ if (old_e1 == e1 and old_e2 == e2) or (old_e1 == e2 and old_e2 == e1):
+ pt.horizon[2 * i + 0] = pt.horizon[2 * (n - 1) + 0]
+ pt.horizon[2 * i + 1] = pt.horizon[2 * (n - 1) + 1]
+ return n - 1
+
+ # out of memory, force EPA to return early without contact
+ if n > pt.horizon.shape[0] - 2:
+ return -1
+
+ pt.horizon[2 * n + 0] = e1
+ pt.horizon[2 * n + 1] = e2
+ return n + 1
+
+
+@wp.func
+def _delete_face(pt: Polytope, face_id: int):
+ index = pt.face_index[face_id]
+ # delete from map
+ if index >= 0:
+ last_face = pt.face_map[pt.nmap - 1]
+ pt.face_map[index] = last_face
+ pt.face_index[last_face] = index
+ pt.nmap -= 1
+ # mark face as deleted from polytope
+ pt.face_index[face_id] = -2
+ return pt.nmap
+
+
+@wp.func
+def _epa_witness(pt: Polytope, face_idx: int):
+ # compute affine coordinates for witness points on plane defined by face
+ v1 = pt.vert[pt.face[face_idx][0]]
+ v2 = pt.vert[pt.face[face_idx][1]]
+ v3 = pt.vert[pt.face[face_idx][2]]
+
+ coordinates = _tri_affine_coord(v1, v2, v3, pt.face_pr[face_idx])
+ l1 = coordinates[0]
+ l2 = coordinates[1]
+ l3 = coordinates[2]
+
+ # face on geom 1
+ v1 = pt.vert1[pt.face[face_idx][0]]
+ v2 = pt.vert1[pt.face[face_idx][1]]
+ v3 = pt.vert1[pt.face[face_idx][2]]
+ x1 = wp.vec3()
+ x1[0] = v1[0] * l1 + v2[0] * l2 + v3[0] * l3
+ x1[1] = v1[1] * l1 + v2[1] * l2 + v3[1] * l3
+ x1[2] = v1[2] * l1 + v2[2] * l2 + v3[2] * l3
+
+ # face on geom 2
+ v1 = pt.vert2[pt.face[face_idx][0]]
+ v2 = pt.vert2[pt.face[face_idx][1]]
+ v3 = pt.vert2[pt.face[face_idx][2]]
+ x2 = wp.vec3()
+ x2[0] = v1[0] * l1 + v2[0] * l2 + v3[0] * l3
+ x2[1] = v1[1] * l1 + v2[1] * l2 + v3[1] * l3
+ x2[2] = v1[2] * l1 + v2[2] * l2 + v3[2] * l3
+
+ return x1, x2
+
+
+@wp.func
+def _polytope2(
+ # In:
+ pt: Polytope,
+ dist: float,
+ simplex: mat43,
+ simplex1: mat43,
+ simplex2: mat43,
+ simplex_index1: wp.vec4i,
+ simplex_index2: wp.vec4i,
+ geom1: Geom,
+ geom2: Geom,
+ geomtype1: int,
+ geomtype2: int,
+):
+ """Create polytope for EPA given a 1-simplex from GJK"""
+ diff = simplex[1] - simplex[0]
+
+ # find component with smallest magnitude (so cross product is largest)
+ value = FLOAT_MAX
+ index = 0
+ for i in range(3):
+ if wp.abs(diff[i]) < value:
+ value = wp.abs(diff[i])
+ index = i
+
+ # cross product with best coordinate axis
+ e = wp.vec(0.0, 0.0, 0.0)
+ e[index] = 1.0
+ d1 = wp.cross(e, diff)
+
+ # rotate around the line segment to get three more points spaced 120 degrees apart
+ R = _rotmat(diff)
+ d2 = R @ d1
+ d3 = R @ d2
+
+ # save vertices and get indices for each one
+ pt.vert[0] = simplex[0]
+ pt.vert[1] = simplex[1]
+
+ pt.vert1[0] = simplex1[0]
+ pt.vert1[1] = simplex1[1]
+
+ pt.vert_index1[0] = simplex_index1[0]
+ pt.vert_index1[1] = simplex_index1[1]
+
+ pt.vert2[0] = simplex2[0]
+ pt.vert2[1] = simplex2[1]
+
+ pt.vert_index2[0] = simplex_index2[0]
+ pt.vert_index2[1] = simplex_index2[1]
+
+ _epa_support(pt, 2, geom1, geom2, geomtype1, geomtype2, d1 / wp.norm_l2(d1))
+ _epa_support(pt, 3, geom1, geom2, geomtype1, geomtype2, d2 / wp.norm_l2(d2))
+ _epa_support(pt, 4, geom1, geom2, geomtype1, geomtype2, d3 / wp.norm_l2(d3))
+
+ # build hexahedron
+ if _attach_face(pt, 0, 0, 2, 3) < MJ_MINVAL:
+ pt.status = -1
+ return pt, _replace_simplex3(pt, 0, 2, 3)
+
+ if _attach_face(pt, 1, 0, 4, 2) < MJ_MINVAL2:
+ pt.status = -1
+ return pt, _replace_simplex3(pt, 0, 4, 2)
+
+ if _attach_face(pt, 2, 0, 3, 4) < MJ_MINVAL2:
+ pt.status = -1
+ return pt, _replace_simplex3(pt, 0, 3, 4)
+
+ if _attach_face(pt, 3, 1, 3, 2) < MJ_MINVAL2:
+ pt.status = -1
+ return pt, _replace_simplex3(pt, 1, 3, 2)
+
+ if _attach_face(pt, 4, 1, 2, 4) < MJ_MINVAL2:
+ pt.status = -1
+ return pt, _replace_simplex3(pt, 1, 2, 4)
+
+ if _attach_face(pt, 5, 1, 4, 3) < MJ_MINVAL2:
+ pt.status = -1
+ return pt, _replace_simplex3(pt, 1, 4, 3)
+
+ # check hexahedron is convex
+ if not _ray_triangle(simplex[0], simplex[1], pt.vert[2], pt.vert[3], pt.vert[4]):
+ pt.status = 1
+ return pt, GJKResult()
+
+ # populate face map
+ for i in range(6):
+ pt.face_map[i] = i
+ pt.face_index[i] = i
+
+ # set polytope counts
+ pt.nvert = 5
+ pt.nface = 6
+ pt.nmap = 6
+ pt.status = 0
+ return pt, GJKResult()
+
+
+@wp.func
+def _polytope3(
+ # In:
+ pt: Polytope,
+ dist: float,
+ simplex: mat43,
+ simplex1: mat43,
+ simplex2: mat43,
+ simplex_index1: wp.vec4i,
+ simplex_index2: wp.vec4i,
+ geom1: Geom,
+ geom2: Geom,
+ geomtype1: int,
+ geomtype2: int,
+):
+ """Create polytope for EPA given a 2-simplex from GJK"""
+ # get normals in both directions
+ n = wp.cross(simplex[1] - simplex[0], simplex[2] - simplex[0])
+ if wp.norm_l2(n) < MJ_MINVAL:
+ pt.status = 2
+ return pt
+
+ pt.vert[0] = simplex[0]
+ pt.vert[1] = simplex[1]
+ pt.vert[2] = simplex[2]
+
+ pt.vert1[0] = simplex1[0]
+ pt.vert1[1] = simplex1[1]
+ pt.vert1[2] = simplex1[2]
+
+ pt.vert_index1[0] = simplex_index1[0]
+ pt.vert_index1[1] = simplex_index1[1]
+ pt.vert_index1[2] = simplex_index1[2]
+
+ pt.vert2[0] = simplex2[0]
+ pt.vert2[1] = simplex2[1]
+ pt.vert2[2] = simplex2[2]
+
+ pt.vert_index2[0] = simplex_index2[0]
+ pt.vert_index2[1] = simplex_index2[1]
+ pt.vert_index2[2] = simplex_index2[2]
+
+ _epa_support(pt, 3, geom1, geom2, geomtype1, geomtype2, -n)
+ _epa_support(pt, 4, geom1, geom2, geomtype1, geomtype2, n)
+
+ v1 = simplex[0]
+ v2 = simplex[1]
+ v3 = simplex[2]
+ v4 = pt.vert[3]
+ v5 = pt.vert[4]
+
+ # check that v4 is not contained in the 2-simplex
+ if _tri_point_intersect(v1, v2, v3, v4):
+ pt.status = 3
+ return pt
+
+ # check that v5 is not contained in the 2-simplex
+ if _tri_point_intersect(v1, v2, v3, v5):
+ pt.status = 4
+ return pt
+
+ # if origin does not lie on simplex then we need to check that the hexahedron contains the
+ # origin
+ if dist > 1e-5 and not _test_tetra(v1, v2, v3, v4) and not _test_tetra(v1, v2, v3, v5):
+ pt.status = 5
+ return pt
+
+ # create hexahedron for EPA
+ if _attach_face(pt, 0, 4, 0, 1) < MJ_MINVAL2:
+ pt.status = 6
+ return pt
+ if _attach_face(pt, 1, 4, 2, 0) < MJ_MINVAL2:
+ pt.status = 7
+ return pt
+ if _attach_face(pt, 2, 4, 1, 2) < MJ_MINVAL2:
+ pt.status = 8
+ return pt
+ if _attach_face(pt, 3, 3, 1, 0) < MJ_MINVAL2:
+ pt.status = 9
+ return pt
+ if _attach_face(pt, 4, 3, 0, 2) < MJ_MINVAL2:
+ pt.status = 10
+ return pt
+ if _attach_face(pt, 5, 3, 2, 1) < MJ_MINVAL2:
+ pt.status = 11
+ return pt
+
+ # populate face map
+ for i in range(6):
+ pt.face_map[i] = i
+ pt.face_index[i] = i
+
+ # set polytope counts
+ pt.nvert = 5
+ pt.nface = 6
+ pt.nmap = 6
+ pt.status = 0
+ return pt
+
+
+@wp.func
+def _polytope4(
+ # In:
+ pt: Polytope,
+ dist: float,
+ simplex: mat43,
+ simplex1: mat43,
+ simplex2: mat43,
+ simplex_index1: wp.vec4i,
+ simplex_index2: wp.vec4i,
+ geom1: Geom,
+ geom2: Geom,
+ geomtype1: int,
+ geomtype2: int,
+):
+ """Create polytope for EPA given a 3-simplex from GJK"""
+ pt.vert[0] = simplex[0]
+ pt.vert[1] = simplex[1]
+ pt.vert[2] = simplex[2]
+ pt.vert[3] = simplex[3]
+
+ pt.vert1[0] = simplex1[0]
+ pt.vert1[1] = simplex1[1]
+ pt.vert1[2] = simplex1[2]
+ pt.vert1[3] = simplex1[3]
+
+ pt.vert_index1[0] = simplex_index1[0]
+ pt.vert_index1[1] = simplex_index1[1]
+ pt.vert_index1[2] = simplex_index1[2]
+ pt.vert_index1[3] = simplex_index1[3]
+
+ pt.vert2[0] = simplex2[0]
+ pt.vert2[1] = simplex2[1]
+ pt.vert2[2] = simplex2[2]
+ pt.vert2[3] = simplex2[3]
+
+ pt.vert_index2[0] = simplex_index2[0]
+ pt.vert_index2[1] = simplex_index2[1]
+ pt.vert_index2[2] = simplex_index2[2]
+ pt.vert_index2[3] = simplex_index2[3]
+
+ # if the origin is on a face, replace the 3-simplex with a 2-simplex
+ if _attach_face(pt, 0, 0, 1, 2) < MJ_MINVAL2:
+ pt.status = -1
+ return pt, _replace_simplex3(pt, 0, 1, 2)
+
+ if _attach_face(pt, 1, 0, 3, 1) < MJ_MINVAL2:
+ pt.status = -1
+ return pt, _replace_simplex3(pt, 0, 3, 1)
+
+ if _attach_face(pt, 2, 0, 2, 3) < MJ_MINVAL2:
+ pt.status = -1
+ return pt, _replace_simplex3(pt, 0, 2, 3)
+
+ if _attach_face(pt, 3, 3, 2, 1) < MJ_MINVAL2:
+ pt.status = -1
+ return pt, _replace_simplex3(pt, 3, 2, 1)
+
+ if not _test_tetra(pt.vert[0], pt.vert[1], pt.vert[2], pt.vert[3]):
+ pt.status = 12
+ return pt, GJKResult()
+
+ # populate face map
+ for i in range(4):
+ pt.face_map[i] = i
+ pt.face_index[i] = i
+
+ # set polytope counts
+ pt.nvert = 4
+ pt.nface = 4
+ pt.nmap = 4
+ pt.status = 0
+ return pt, GJKResult()
+
+
+@wp.func
+def _epa(tolerance2: float, epa_iterations: int, pt: Polytope, geom1: Geom, geom2: Geom, geomtype1: int, geomtype2: int):
+ """Recover penetration data from two geoms in contact given an initial polytope."""
+ upper = FLOAT_MAX
+ upper2 = FLOAT_MAX
+ idx = int(-1)
+ pidx = int(-1)
+
+ for k in range(epa_iterations):
+ pidx = int(idx)
+ idx = int(-1)
+
+ # find the face closest to the origin (lower bound for penetration depth)
+ lower2 = float(FLOAT_MAX)
+ for i in range(pt.nmap):
+ face_idx = pt.face_map[i]
+ if pt.face_norm2[face_idx] < lower2:
+ idx = int(face_idx)
+ lower2 = float(pt.face_norm2[face_idx])
+
+ # face not valid, return previous face
+ if lower2 > upper2 or idx < 0:
+ idx = pidx
+ break
+
+ # check if lower bound is 0
+ if lower2 <= 0:
+ break
+
+ # compute support point w from the closest face's normal
+ lower = wp.sqrt(lower2)
+ wi = pt.nvert
+ i1, i2 = _epa_support(pt, wi, geom1, geom2, geomtype1, geomtype2, pt.face_pr[idx] / lower)
+ geom1.index = i1
+ geom2.index = i2
+ pt.nvert += 1
+
+ # upper bound for kth iteration
+ upper_k = wp.dot(pt.face_pr[idx], pt.vert[wi]) / lower
+ if upper_k < upper:
+ upper = upper_k
+ upper2 = upper * upper
+
+ if upper - lower < tolerance2:
+ break
+
+ pt.nmap = _delete_face(pt, idx)
+ pt.nhorizon = _add_edge(pt, pt.face[idx][0], pt.face[idx][1])
+ pt.nhorizon = _add_edge(pt, pt.face[idx][1], pt.face[idx][2])
+ pt.nhorizon = _add_edge(pt, pt.face[idx][2], pt.face[idx][0])
+ if pt.nhorizon == -1:
+ idx = -1
+ break
+
+ # compute horizon for w
+ for i in range(pt.nface):
+ if pt.face_index[i] == -2:
+ continue
+
+ if wp.dot(pt.face_pr[i], pt.vert[wi]) - pt.face_norm2[i] > MJ_MINVAL:
+ pt.nmap = _delete_face(pt, i)
+ pt.nhorizon = _add_edge(pt, pt.face[i][0], pt.face[i][1])
+ pt.nhorizon = _add_edge(pt, pt.face[i][1], pt.face[i][2])
+ pt.nhorizon = _add_edge(pt, pt.face[i][2], pt.face[i][0])
+ if pt.nhorizon == -1:
+ idx = -1
+ break
+
+ # insert w as new vertex and attach faces along the horizon
+ for i in range(pt.nhorizon):
+ dist2 = _attach_face(pt, pt.nface, wi, pt.horizon[2 * i + 0], pt.horizon[2 * i + 1])
+ if dist2 == 0:
+ idx = -1
+ break
+
+ pt.nface += 1
+
+ # store face in map
+ if dist2 >= lower2 and dist2 <= upper2:
+ pt.face_map[pt.nmap] = pt.nface - 1
+ pt.face_index[pt.nface - 1] = pt.nmap
+ pt.nmap += 1
+
+ # no face candidates left
+ if pt.nmap == 0 or idx == -1:
+ break
+
+ # clear horizon
+ pt.nhorizon = 0
+
+ # return from valid face
+ if idx > -1:
+ x1, x2 = _epa_witness(pt, idx)
+ return -wp.sqrt(pt.face_norm2[idx]), x1, x2
+ return 0.0, wp.vec3(), wp.vec3()
+
+
+@wp.func
+def ccd(
+ # In:
+ tolerance: float,
+ cutoff: float,
+ gjk_iterations: int,
+ epa_iterations: int,
+ geom1: Geom,
+ geom2: Geom,
+ geomtype1: int,
+ geomtype2: int,
+ x_1: wp.vec3,
+ x_2: wp.vec3,
+ vert: wp.array(dtype=wp.vec3),
+ vert1: wp.array(dtype=wp.vec3),
+ vert2: wp.array(dtype=wp.vec3),
+ vert_index1: wp.array(dtype=int),
+ vert_index2: wp.array(dtype=int),
+ face: wp.array(dtype=wp.vec3i),
+ face_pr: wp.array(dtype=wp.vec3),
+ face_norm2: wp.array(dtype=float),
+ face_index: wp.array(dtype=int),
+ face_map: wp.array(dtype=int),
+ horizon: wp.array(dtype=int),
+):
+ """General convex collision detection via GJK/EPA."""
+ result = _gjk(tolerance, gjk_iterations, geom1, geom2, x_1, x_2, geomtype1, geomtype2, cutoff)
+
+ # no penetration depth to recover
+ if result.dist > tolerance or result.dim < 2:
+ return result.dist, result.x1, result.x2
+
+ pt = Polytope()
+ pt.nface = 0
+ pt.nmap = 0
+ pt.nvert = 0
+ pt.nhorizon = 0
+ pt.vert = vert
+ pt.vert1 = vert1
+ pt.vert2 = vert2
+ pt.vert_index1 = vert_index1
+ pt.vert_index2 = vert_index2
+ pt.face = face
+ pt.face_pr = face_pr
+ pt.face_norm2 = face_norm2
+ pt.face_index = face_index
+ pt.face_map = face_map
+ pt.horizon = horizon
+
+ if result.dim == 2:
+ pt, new_result = _polytope2(
+ pt,
+ result.dist,
+ result.simplex,
+ result.simplex1,
+ result.simplex2,
+ result.simplex_index1,
+ result.simplex_index2,
+ geom1,
+ geom2,
+ geomtype1,
+ geomtype2,
+ )
+ if pt.status == -1:
+ result.simplex = new_result.simplex
+ result.simplex1 = new_result.simplex1
+ result.simplex2 = new_result.simplex2
+ result.simplex_index1 = new_result.simplex_index1
+ result.simplex_index2 = new_result.simplex_index2
+ result.dim = 3
+ elif result.dim == 4:
+ pt, new_result = _polytope4(
+ pt,
+ result.dist,
+ result.simplex,
+ result.simplex1,
+ result.simplex2,
+ result.simplex_index1,
+ result.simplex_index2,
+ geom1,
+ geom2,
+ geomtype1,
+ geomtype2,
+ )
+ if pt.status == -1:
+ result.simplex = new_result.simplex
+ result.simplex1 = new_result.simplex1
+ result.simplex2 = new_result.simplex2
+ result.simplex_index1 = new_result.simplex_index1
+ result.simplex_index2 = new_result.simplex_index2
+ result.dim = 3
+
+ # polytope2 and polytope4 may need to fallback here
+ if result.dim == 3:
+ pt = _polytope3(
+ pt,
+ result.dist,
+ result.simplex,
+ result.simplex1,
+ result.simplex2,
+ result.simplex_index1,
+ result.simplex_index2,
+ geom1,
+ geom2,
+ geomtype1,
+ geomtype2,
+ )
+
+ # origin on boundary (objects are not considered penetrating)
+ if pt.status:
+ return result.dist, result.x1, result.x2
+
+ return _epa(tolerance * tolerance, epa_iterations, pt, geom1, geom2, geomtype1, geomtype2)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py
new file mode 100644
index 00000000..e8313be6
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py
@@ -0,0 +1,714 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_prism_vertex
+from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import Geom
+from mujoco.mjx.third_party.mujoco_warp._src.math import gjk_normalize
+from mujoco.mjx.third_party.mujoco_warp._src.math import orthonormal
+from mujoco.mjx.third_party.mujoco_warp._src.support import all_same
+from mujoco.mjx.third_party.mujoco_warp._src.support import any_different
+from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
+from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
+
+# TODO(team): improve compile time to enable backward pass
+wp.config.enable_backward = False
+
+FLOAT_MIN = -1e30
+FLOAT_MAX = 1e30
+EPS_BEST_COUNT = 12
+MULTI_CONTACT_COUNT = 4
+MULTI_POLYGON_COUNT = 8
+
+matc3 = wp.types.matrix(shape=(EPS_BEST_COUNT, 3), dtype=float)
+vecc3 = wp.types.vector(EPS_BEST_COUNT * 3, dtype=float)
+
+# Matrix definition for the `tris` scratch space which is used to store the
+# triangles of the polytope. Note that the first dimension is 2, as we need
+# to store the previous and current polytope. But since Warp doesn't support
+# 3D matrices yet, we use 2 * 3 * EPS_BEST_COUNT as the first dimension.
+TRIS_DIM = 3 * EPS_BEST_COUNT
+mat2c3 = wp.types.matrix(shape=(2 * TRIS_DIM, 3), dtype=float)
+mat3p = wp.types.matrix(shape=(MULTI_POLYGON_COUNT, 3), dtype=float)
+mat3c = wp.types.matrix(shape=(MULTI_CONTACT_COUNT, 3), dtype=float)
+mat43 = wp.types.matrix(shape=(4, 3), dtype=float)
+
+vec6 = wp.types.vector(6, dtype=int)
+VECI1 = vec6(0, 0, 0, 1, 1, 2)
+VECI2 = vec6(1, 2, 3, 2, 3, 3)
+
+
+@wp.func
+def _gjk_support_geom(geom: Geom, geomtype: int, dir: wp.vec3):
+ local_dir = wp.transpose(geom.rot) @ dir
+ if geomtype == int(GeomType.SPHERE.value):
+ support_pt = geom.pos + geom.size[0] * dir
+ elif geomtype == int(GeomType.BOX.value):
+ res = wp.cw_mul(wp.sign(local_dir), geom.size)
+ support_pt = geom.rot @ res + geom.pos
+ elif geomtype == int(GeomType.CAPSULE.value):
+ res = local_dir * geom.size[0]
+ # add cylinder contribution
+ res[2] += wp.sign(local_dir[2]) * geom.size[1]
+ support_pt = geom.rot @ res + geom.pos
+ elif geomtype == int(GeomType.ELLIPSOID.value):
+ res = wp.cw_mul(local_dir, geom.size)
+ res = wp.normalize(res)
+ # transform to ellipsoid
+ res = wp.cw_mul(res, geom.size)
+ support_pt = geom.rot @ res + geom.pos
+ elif geomtype == int(GeomType.CYLINDER.value):
+ res = wp.vec3(0.0, 0.0, 0.0)
+ # set result in XY plane: support on circle
+ d = wp.sqrt(wp.dot(local_dir, local_dir))
+ if d > MJ_MINVAL:
+ scl = geom.size[0] / d
+ res[0] = local_dir[0] * scl
+ res[1] = local_dir[1] * scl
+ # set result in Z direction
+ res[2] = wp.sign(local_dir[2]) * geom.size[1]
+ support_pt = geom.rot @ res + geom.pos
+ elif geomtype == int(GeomType.MESH.value):
+ max_dist = float(FLOAT_MIN)
+ if geom.graphadr == -1 or geom.vertnum < 10:
+ # exhaustive search over all vertices
+ for i in range(geom.vertnum):
+ vert = geom.vert[geom.vertadr + i]
+ dist = wp.dot(vert, local_dir)
+ if dist > max_dist:
+ max_dist = dist
+ support_pt = vert
+ else:
+ numvert = geom.graph[geom.graphadr]
+ vert_edgeadr = geom.graphadr + 2
+ vert_globalid = geom.graphadr + 2 + numvert
+ edge_localid = geom.graphadr + 2 + 2 * numvert
+ # hillclimb until no change
+ prev = int(-1)
+ imax = int(0)
+
+ while True:
+ prev = int(imax)
+ i = int(geom.graph[vert_edgeadr + imax])
+ while geom.graph[edge_localid + i] >= 0:
+ subidx = geom.graph[edge_localid + i]
+ idx = geom.graph[vert_globalid + subidx]
+ dist = wp.dot(local_dir, geom.vert[geom.vertadr + idx])
+ if dist > max_dist:
+ max_dist = dist
+ imax = int(subidx)
+ i += int(1)
+ if imax == prev:
+ break
+ imax = geom.graph[vert_globalid + imax]
+ support_pt = geom.vert[geom.vertadr + imax]
+
+ support_pt = geom.rot @ support_pt + geom.pos
+ elif geomtype == int(GeomType.HFIELD.value):
+ max_dist = float(FLOAT_MIN)
+ for i in range(6):
+ vert = hfield_prism_vertex(geom.hfprism, i)
+ dist = wp.dot(vert, local_dir)
+ if dist > max_dist:
+ max_dist = dist
+ support_pt = vert
+ support_pt = geom.rot @ support_pt + geom.pos
+
+ return wp.dot(support_pt, dir), support_pt
+
+
+@wp.func
+def _gjk_support(
+ # In:
+ geom1: Geom,
+ geom2: Geom,
+ geomtype1: int,
+ geomtype2: int,
+ dir: wp.vec3,
+):
+ # Returns the distance between support points on two geoms, and the support point.
+ # Negative distance means objects are not intersecting along direction `dir`.
+ # Positive distance means objects are intersecting along the given direction `dir`.
+
+ dist1, s1 = _gjk_support_geom(geom1, geomtype1, dir)
+ dist2, s2 = _gjk_support_geom(geom2, geomtype2, -dir)
+
+ support_pt = s1 - s2
+ return dist1 + dist2, support_pt
+
+
+@wp.func
+def _expand_polytope(count: int, prev_count: int, dists: vecc3, tris: mat2c3, p: matc3):
+ # expand polytope greedily
+ for j in range(count):
+ best = int(0)
+ dd = dists[0]
+ for i in range(1, 3 * prev_count):
+ if dists[i] < dd:
+ dd = dists[i]
+ best = i
+
+ dists[best] = float(wp.static(2 * FLOAT_MAX))
+
+ parent_index = best // 3
+ child_index = best % 3
+
+ # fill in the new triangle at the next index
+ tris[TRIS_DIM + j * 3 + 0] = tris[parent_index * 3 + child_index]
+ tris[TRIS_DIM + j * 3 + 1] = tris[parent_index * 3 + ((child_index + 1) % 3)]
+ tris[TRIS_DIM + j * 3 + 2] = p[parent_index]
+
+ for r in range(wp.static(EPS_BEST_COUNT * 3)):
+ # swap triangles
+ swap = tris[TRIS_DIM + r]
+ tris[TRIS_DIM + r] = tris[r]
+ tris[r] = swap
+
+ return dists, tris
+
+
+@wp.func
+def gjk_legacy(
+ # In:
+ gjk_iterations: int,
+ geom1: Geom,
+ geom2: Geom,
+ geomtype1: int,
+ geomtype2: int,
+):
+ dir = wp.vec3(0.0, 0.0, 1.0)
+ dir_n = -dir
+ depth = float(FLOAT_MAX)
+
+ dist_max, simplex0 = _gjk_support(geom1, geom2, geomtype1, geomtype2, dir)
+ dist_min, simplex1 = _gjk_support(geom1, geom2, geomtype1, geomtype2, dir_n)
+
+ if dist_max < dist_min:
+ depth = dist_max
+ normal = dir
+ else:
+ depth = dist_min
+ normal = dir_n
+
+ sd = simplex0 - simplex1
+ dir = orthonormal(sd)
+
+ dist_max, simplex3 = _gjk_support(geom1, geom2, geomtype1, geomtype2, dir)
+
+ # Initialize a 2-simplex with simplex[2]==simplex[1]. This ensures the
+ # correct winding order for face normals defined below. Face 0 and face 3
+ # are degenerate, and face 1 and 2 have opposing normals.
+ simplex = mat43()
+ simplex[0] = simplex0
+ simplex[1] = simplex1
+ simplex[2] = simplex[1]
+ simplex[3] = simplex3
+
+ if dist_max < depth:
+ depth = dist_max
+ normal = dir
+ if dist_min < depth:
+ depth = dist_min
+ normal = dir_n
+
+ plane = mat43()
+ for _ in range(gjk_iterations):
+ # winding orders: plane[0] ccw, plane[1] cw, plane[2] ccw, plane[3] cw
+ plane[0] = wp.cross(simplex[3] - simplex[2], simplex[1] - simplex[2])
+ plane[1] = wp.cross(simplex[3] - simplex[0], simplex[2] - simplex[0])
+ plane[2] = wp.cross(simplex[3] - simplex[1], simplex[0] - simplex[1])
+ plane[3] = wp.cross(simplex[2] - simplex[0], simplex[1] - simplex[0])
+
+ # Compute distance of each face halfspace to the origin. If dplane<0, then the
+ # origin is outside the halfspace. If dplane>0 then the origin is inside
+ # the halfspace defined by the face plane.
+
+ dplane = wp.vec4(float(FLOAT_MAX))
+
+ plane0, p0 = gjk_normalize(plane[0])
+ plane1, p1 = gjk_normalize(plane[1])
+ plane2, p2 = gjk_normalize(plane[2])
+ plane3, p3 = gjk_normalize(plane[3])
+
+ plane[0] = plane0
+ plane[1] = plane1
+ plane[2] = plane2
+ plane[3] = plane3
+
+ if p0:
+ dplane[0] = wp.dot(plane[0], simplex[2])
+
+ if p1:
+ dplane[1] = wp.dot(plane[1], simplex[0])
+
+ if p2:
+ dplane[2] = wp.dot(plane[2], simplex[1])
+
+ if p3:
+ dplane[3] = wp.dot(plane[3], simplex[0])
+
+ # pick plane normal with minimum distance to the origin
+ i1 = wp.where(dplane[0] < dplane[1], 0, 1)
+ i2 = wp.where(dplane[2] < dplane[3], 2, 3)
+ index = wp.where(dplane[i1] < dplane[i2], i1, i2)
+
+ if dplane[index] > 0.0:
+ # origin is inside the simplex, objects are intersecting
+ break
+
+ # add new support point to the simplex
+ dist, simplex_i = _gjk_support(geom1, geom2, geomtype1, geomtype2, plane[index])
+ simplex[index] = simplex_i
+
+ if dist < depth:
+ depth = dist
+ normal = plane[index]
+
+ # preserve winding order of the simplex faces
+ index1 = (index + 1) & 3
+ index2 = (index + 2) & 3
+ swap = simplex[index1]
+ simplex[index1] = simplex[index2]
+ simplex[index2] = swap
+
+ if dist < 0.0:
+ break # objects are likely non-intersecting
+
+ return simplex, normal
+
+
+@wp.func
+def epa_legacy(
+ # In:
+ epa_iterations: int,
+ geom1: Geom,
+ geom2: Geom,
+ geomtype1: int,
+ geomtype2: int,
+ depth_extension: float,
+ epa_exact_neg_distance: bool,
+ simplex: mat43,
+ normal: wp.vec3,
+):
+ # get the support, if depth < 0: objects do not intersect
+ depth, _ = _gjk_support(geom1, geom2, geomtype1, geomtype2, normal)
+
+ if depth < -depth_extension:
+ # Objects are not intersecting, and we do not obtain the closest points as
+ # specified by depth_extension.
+ return wp.nan, wp.vec3(wp.nan, wp.nan, wp.nan)
+
+ if wp.static(epa_exact_neg_distance):
+ # Check closest points to all edges of the simplex, rather than just the
+ # face normals. This gives the exact depth/normal for the non-intersecting
+ # case.
+ for i in range(6):
+ i1 = VECI1[i]
+ i2 = VECI2[i]
+
+ si1 = simplex[i1]
+ si2 = simplex[i2]
+
+ if si1[0] != si2[0] or si1[1] != si2[1] or si1[2] != si2[2]:
+ v = si1 - si2
+ alpha = wp.dot(si1, v) / wp.dot(v, v)
+
+ # p0 is the closest segment point to the origin
+ p0 = wp.clamp(alpha, 0.0, 1.0) * v - si1
+ p0, pf = gjk_normalize(p0)
+
+ if pf:
+ depth2, _ = _gjk_support(geom1, geom2, geomtype1, geomtype2, p0)
+
+ if depth2 < depth:
+ depth = depth2
+ normal = p0
+
+ # supporting points for each triangle
+ p = matc3()
+
+ # distance to the origin for candidate triangles
+ dists = vecc3()
+
+ tris = mat2c3()
+ tris[0] = simplex[2]
+ tris[1] = simplex[1]
+ tris[2] = simplex[3]
+
+ tris[3] = simplex[0]
+ tris[4] = simplex[2]
+ tris[5] = simplex[3]
+
+ tris[6] = simplex[1]
+ tris[7] = simplex[0]
+ tris[8] = simplex[3]
+
+ tris[9] = simplex[0]
+ tris[10] = simplex[1]
+ tris[11] = simplex[2]
+
+ # Calculate the total number of iterations to avoid nested loop
+ # This is a hack to reduce compile time
+ count = int(4)
+ it = int(0)
+ for _ in range(wp.static(epa_iterations)):
+ it += count
+ count = wp.min(count * 3, EPS_BEST_COUNT)
+
+ count = int(4)
+ i = int(0)
+ for _ in range(it):
+ # Loop through all triangles, and obtain distances to the origin for each
+ # new triangle candidate.
+ ti = 3 * i
+ n = wp.cross(tris[ti + 2] - tris[ti + 0], tris[ti + 1] - tris[ti + 0])
+
+ n, nf = gjk_normalize(n)
+ if not nf:
+ for j in range(3):
+ dists[i * 3 + j] = wp.static(float(2 * FLOAT_MAX))
+ continue
+
+ dist, pi = _gjk_support(geom1, geom2, geomtype1, geomtype2, n)
+ p[i] = pi
+
+ if dist < depth:
+ depth = dist
+ normal = n
+
+ # iterate over edges and get distance using support point
+ for j in range(3):
+ if wp.static(epa_exact_neg_distance):
+ # obtain closest point between new triangle edge and origin
+ tqj = tris[ti + j]
+
+ if (p[i, 0] != tqj[0]) or (p[i, 1] != tqj[1]) or (p[i, 2] != tqj[2]):
+ v = p[i] - tris[ti + j]
+ alpha = wp.dot(p[i], v) / wp.dot(v, v)
+ p0 = wp.clamp(alpha, 0.0, 1.0) * v - p[i]
+ p0, pf = gjk_normalize(p0)
+
+ if pf:
+ dist2, v = _gjk_support(geom1, geom2, geomtype1, geomtype2, p0)
+
+ if dist2 < depth:
+ depth = dist2
+ normal = p0
+
+ plane = wp.cross(p[i] - tris[ti + j], tris[ti + ((j + 1) % 3)] - tris[ti + j])
+ plane, pf = gjk_normalize(plane)
+
+ if pf:
+ dd = wp.dot(plane, tris[ti + j])
+ else:
+ dd = float(FLOAT_MAX)
+
+ if (dd < 0 and depth >= 0) or (
+ tris[ti + ((j + 2) % 3)][0] == p[i][0]
+ and tris[ti + ((j + 2) % 3)][1] == p[i][1]
+ and tris[ti + ((j + 2) % 3)][2] == p[i][2]
+ ):
+ dists[i * 3 + j] = float(FLOAT_MAX)
+ else:
+ dists[i * 3 + j] = dd
+
+ if i == count - 1:
+ prev_count = count
+ count = wp.min(count * 3, EPS_BEST_COUNT)
+ dists, tris = _expand_polytope(count, prev_count, dists, tris, p)
+ i = int(0)
+ else:
+ i += 1
+
+ return depth, normal
+
+
+@wp.func
+def multicontact_legacy(
+ # In:
+ geom1: Geom,
+ geom2: Geom,
+ geomtype1: int,
+ geomtype2: int,
+ depth_extension: float,
+ depth: float,
+ normal: wp.vec3,
+ ncontact: int,
+ npolygon: int,
+ perturbation_angle: float,
+):
+ # Calculates multiple contact points given the normal from EPA.
+ # 1. Calculates the polygon on each shape by tiling the normal
+ # "perturbation_angle" (radians) in the orthogonal component of the normal.
+ # The "perturbation_angle" can be changed to depend on the depth of the
+ # contact, in a future version.
+ # 2. The normal is tilted "npolygon" times in the directions evenly
+ # spaced in the orthogonal component of the normal.
+ # (works well for >= 6, default is 8).
+ # 3. The intersection between these two polygons is calculated in 2D space
+ # (complement to the normal). If they intersect, extreme points in both
+ # directions are found. This can be modified to the extremes in the
+ # direction of eigenvectors of the variance of points of each polygon. If
+ # they do not intersect, the closest points of both polygons are found.
+
+ assert ncontact <= MULTI_CONTACT_COUNT
+ assert npolygon <= MULTI_POLYGON_COUNT
+
+ if depth < -depth_extension:
+ return 0, mat3c()
+
+ dir = orthonormal(normal)
+ dir2 = wp.cross(normal, dir)
+
+ angle = perturbation_angle
+ c = wp.cos(angle)
+ s = wp.sin(angle)
+ tc = 1.0 - c
+
+ v1 = mat3p()
+ v2 = mat3p()
+
+ contact_points = mat3c()
+
+ # Obtain points on the polygon determined by the support and tilt angle,
+ # in the basis of the contact frame.
+ v1count = int(0)
+ v2count = int(0)
+ angle_ratio = wp.static(2.0 * wp.pi) / float(npolygon)
+
+ for i in range(npolygon):
+ angle = angle_ratio * float(i)
+ axis = wp.cos(angle) * dir + wp.sin(angle) * dir2
+
+ # Axis-angle rotation matrix. See
+ # https://en.wikipedia.org/wiki/Rotation_matrix#Rotation_matrix_from_axis_and_angle
+ mat0 = c + axis[0] * axis[0] * tc
+ mat5 = c + axis[1] * axis[1] * tc
+ mat10 = c + axis[2] * axis[2] * tc
+ t1 = axis[0] * axis[1] * tc
+ t2 = axis[2] * s
+ mat4 = t1 + t2
+ mat1 = t1 - t2
+ t1 = axis[0] * axis[2] * tc
+ t2 = axis[1] * s
+ mat8 = t1 - t2
+ mat2 = t1 + t2
+ t1 = axis[1] * axis[2] * tc
+ t2 = axis[0] * s
+ mat9 = t1 + t2
+ mat6 = t1 - t2
+
+ n = wp.vec3(
+ mat0 * normal[0] + mat1 * normal[1] + mat2 * normal[2],
+ mat4 * normal[0] + mat5 * normal[1] + mat6 * normal[2],
+ mat8 * normal[0] + mat9 * normal[1] + mat10 * normal[2],
+ )
+
+ _, p = _gjk_support_geom(geom1, geomtype1, n)
+ v1[v1count] = wp.vec3(wp.dot(p, dir), wp.dot(p, dir2), wp.dot(p, normal))
+
+ if i == 0:
+ v1count += 1
+ elif any_different(v1[v1count], v1[v1count - 1]):
+ v1count += 1
+
+ n = -n
+ _, p = _gjk_support_geom(geom2, geomtype2, n)
+ v2[v2count] = wp.vec3(wp.dot(p, dir), wp.dot(p, dir2), wp.dot(p, normal))
+
+ if i == 0:
+ v2count += 1
+ elif any_different(v2[v2count], v2[v2count - 1]):
+ v2count += 1
+
+ # remove duplicate vertices on the array boundary
+ if v1count > 1 and all_same(v1[v1count - 1], v1[0]):
+ v1count -= 1
+
+ if v2count > 1 and all_same(v2[v2count - 1], v2[0]):
+ v2count -= 1
+
+ # find an intersecting polygon between v1 and v2 in the 2D plane
+ out = mat43()
+ candCount = int(0)
+
+ if v2count > 1:
+ for i in range(v1count):
+ m1a = v1[i]
+ is_in = bool(True)
+
+ # check if point m1a is inside the v2 polygon on the 2D plane
+ for j in range(v2count):
+ j2 = (j + 1) % v2count
+
+ # Checks that orientation of the triangle (v2[j], v2[j2], m1a) is
+ # counter-clockwise. If so, point m1a is inside the v2 polygon.
+ is_in = is_in and ((v2[j2][0] - v2[j][0]) * (m1a[1] - v2[j][1]) - (v2[j2][1] - v2[j][1]) * (m1a[0] - v2[j][0]) >= 0.0)
+
+ if not is_in:
+ break
+
+ if is_in:
+ if not candCount or m1a[0] < out[0, 0]:
+ out[0] = m1a
+ if not candCount or m1a[0] > out[1, 0]:
+ out[1] = m1a
+ if not candCount or m1a[1] < out[2, 1]:
+ out[2] = m1a
+ if not candCount or m1a[1] > out[3, 1]:
+ out[3] = m1a
+ candCount += 1
+
+ if v1count > 1:
+ for i in range(v2count):
+ m1a = v2[i]
+ is_in = bool(True)
+
+ for j in range(v1count):
+ j2 = (j + 1) % v1count
+ is_in = is_in and (v1[j2][0] - v1[j][0]) * (m1a[1] - v1[j][1]) - (v1[j2][1] - v1[j][1]) * (m1a[0] - v1[j][0]) >= 0.0
+ if not is_in:
+ break
+
+ if is_in:
+ if not candCount or m1a[0] < out[0, 0]:
+ out[0] = m1a
+ if not candCount or m1a[0] > out[1, 0]:
+ out[1] = m1a
+ if not candCount or m1a[1] < out[2, 1]:
+ out[2] = m1a
+ if not candCount or m1a[1] > out[3, 1]:
+ out[3] = m1a
+ candCount += 1
+
+ if v1count > 1 and v2count > 1:
+ # Check all edge pairs, and store line segment intersections if they are
+ # on the edge of the boundary.
+ for i in range(v1count):
+ for j in range(v2count):
+ m1a = v1[i]
+ m1b = v1[(i + 1) % v1count]
+ m2a = v2[j]
+ m2b = v2[(j + 1) % v2count]
+
+ det = (m2a[1] - m2b[1]) * (m1b[0] - m1a[0]) - (m1a[1] - m1b[1]) * (m2b[0] - m2a[0])
+
+ if wp.abs(det) > 1e-12:
+ a11 = (m2a[1] - m2b[1]) / det
+ a12 = (m2b[0] - m2a[0]) / det
+ a21 = (m1a[1] - m1b[1]) / det
+ a22 = (m1b[0] - m1a[0]) / det
+ b1 = m2a[0] - m1a[0]
+ b2 = m2a[1] - m1a[1]
+
+ alpha = a11 * b1 + a12 * b2
+ beta = a21 * b1 + a22 * b2
+ if alpha >= 0.0 and alpha <= 1.0 and beta >= 0.0 and beta <= 1.0:
+ m0 = wp.vec3(
+ m1a[0] + alpha * (m1b[0] - m1a[0]),
+ m1a[1] + alpha * (m1b[1] - m1a[1]),
+ (m1a[2] + alpha * (m1b[2] - m1a[2]) + m2a[2] + beta * (m2b[2] - m2a[2])) * 0.5,
+ )
+ if not candCount or m0[0] < out[0, 0]:
+ out[0] = m0
+ if not candCount or m0[0] > out[1, 0]:
+ out[1] = m0
+ if not candCount or m0[1] < out[2, 1]:
+ out[2] = m0
+ if not candCount or m0[1] > out[3, 1]:
+ out[3] = m0
+ candCount += 1
+
+ var_rx = wp.vec3(0.0)
+ contact_count = int(0)
+ if candCount > 0:
+ # Polygon intersection was found.
+ # TODO(btaba): replace the above routine with the manifold point routine
+ # from MJX. Deduplicate the points properly.
+ last_pt = wp.vec3(FLOAT_MAX, FLOAT_MAX, FLOAT_MAX)
+
+ for k in range(ncontact):
+ pt = out[k, 0] * dir + out[k, 1] * dir2 + out[k, 2] * normal
+
+ # skip contact points that are too close
+ if wp.length(pt - last_pt) <= 1e-6:
+ continue
+
+ contact_points[contact_count] = pt
+ last_pt = pt
+ contact_count += 1
+
+ else:
+ # Polygon intersection was not found. Loop through all vertex pairs and
+ # calculate an approximate contact point.
+ minDist = float(0.0)
+ for i in range(v1count):
+ for j in range(v2count):
+ # Find the closest vertex pair. Calculate a contact point var_rx as the
+ # midpoint between the closest vertex pair.
+ m1 = v1[i]
+ m2 = v2[j]
+ dd = (m1[0] - m2[0]) * (m1[0] - m2[0]) + (m1[1] - m2[1]) * (m1[1] - m2[1])
+
+ if i != 0 and j != 0 or dd < minDist:
+ minDist = dd
+ var_rx = ((m1[0] + m2[0]) * dir + (m1[1] + m2[1]) * dir2 + (m1[2] + m2[2]) * normal) * 0.5
+
+ # Check for a closer point between a point on v2 and an edge on v1.
+ m1b = v1[(i + 1) % v1count]
+ m2b = v2[(j + 1) % v2count]
+
+ if v1count > 1:
+ dd = (m1b[0] - m1[0]) * (m1b[0] - m1[0]) + (m1b[1] - m1[1]) * (m1b[1] - m1[1])
+ t = ((m2[1] - m1[1]) * (m1b[0] - m1[0]) - (m2[0] - m1[0]) * (m1b[1] - m1[1])) / dd
+ dx = m2[0] + (m1b[1] - m1[1]) * t
+ dy = m2[1] - (m1b[0] - m1[0]) * t
+ dist = (dx - m2[0]) * (dx - m2[0]) + (dy - m2[1]) * (dy - m2[1])
+
+ if (
+ (dist < minDist)
+ and ((dx - m1[0]) * (m1b[0] - m1[0]) + (dy - m1[1]) * (m1b[1] - m1[1]) >= 0)
+ and ((dx - m1b[0]) * (m1[0] - m1b[0]) + (dy - m1b[1]) * (m1[1] - m1b[1]) >= 0)
+ ):
+ alpha = wp.sqrt(((dx - m1[0]) * (dx - m1[0]) + (dy - m1[1]) * (dy - m1[1])) / dd)
+ minDist = dist
+ w = ((1.0 - alpha) * m1 + alpha * m1b + m2) * 0.5
+ var_rx = w[0] * dir + w[1] * dir2 + w[2] * normal
+
+ # check for a closer point between a point on v1 and an edge on v2
+ if v2count > 1:
+ dd = (m2b[0] - m2[0]) * (m2b[0] - m2[0]) + (m2b[1] - m2[1]) * (m2b[1] - m2[1])
+ t = ((m1[1] - m2[1]) * (m2b[0] - m2[0]) - (m1[0] - m2[0]) * (m2b[1] - m2[1])) / dd
+ dx = m1[0] + (m2b[1] - m2[1]) * t
+ dy = m1[1] - (m2b[0] - m2[0]) * t
+ dist = (dx - m1[0]) * (dx - m1[0]) + (dy - m1[1]) * (dy - m1[1])
+
+ if (
+ dist < minDist
+ and (dx - m2[0]) * (m2b[0] - m2[0]) + (dy - m2[1]) * (m2b[1] - m2[1]) >= 0
+ and (dx - m2b[0]) * (m2[0] - m2b[0]) + (dy - m2b[1]) * (m2[1] - m2b[1]) >= 0
+ ):
+ alpha = wp.sqrt(((dx - m2[0]) * (dx - m2[0]) + (dy - m2[1]) * (dy - m2[1])) / dd)
+ minDist = dist
+ w = (m1 + (1.0 - alpha) * m2 + alpha * m2b) * 0.5
+ var_rx = w[0] * dir + w[1] * dir2 + w[2] * normal
+
+ for k in range(ncontact):
+ contact_points[k] = var_rx
+
+ contact_count = 1
+
+ return contact_count, contact_points
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py
new file mode 100644
index 00000000..8ccb259a
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py
@@ -0,0 +1,320 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import warp as wp
+from absl.testing import absltest
+
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel
+
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import ccd
+from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import Geom
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+
+MAX_ITERATIONS = 10
+
+
+def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int):
+ @nested_kernel
+ def _gjk_kernel(
+ # Model:
+ geom_type: wp.array(dtype=int),
+ geom_dataid: wp.array(dtype=int),
+ geom_size: wp.array2d(dtype=wp.vec3),
+ mesh_vertadr: wp.array(dtype=int),
+ mesh_vertnum: wp.array(dtype=int),
+ mesh_vert: wp.array(dtype=wp.vec3),
+ # Data in:
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ # In:
+ gid1: int,
+ gid2: int,
+ iterations: int,
+ vert: wp.array(dtype=wp.vec3),
+ vert1: wp.array(dtype=wp.vec3),
+ vert2: wp.array(dtype=wp.vec3),
+ vert_index1: wp.array(dtype=int),
+ vert_index2: wp.array(dtype=int),
+ face: wp.array(dtype=wp.vec3i),
+ face_pr: wp.array(dtype=wp.vec3),
+ face_norm2: wp.array(dtype=float),
+ face_index: wp.array(dtype=int),
+ face_map: wp.array(dtype=int),
+ horizon: wp.array(dtype=int),
+ # Out:
+ dist_out: wp.array(dtype=float),
+ pos_out: wp.array(dtype=wp.vec3),
+ ):
+ MESHGEOM = int(GeomType.MESH.value)
+
+ geom1 = Geom()
+ geom1.index = -1
+ geomtype1 = geom_type[gid1]
+ geom1.pos = geom_xpos_in[0, gid1]
+ geom1.rot = geom_xmat_in[0, gid1]
+ geom1.size = geom_size[0, gid1]
+ geom1.graphadr = -1
+
+ if geom_dataid[gid1] >= 0 and geom_type[gid1] == MESHGEOM:
+ dataid = geom_dataid[gid1]
+ geom1.vertadr = mesh_vertadr[dataid]
+ geom1.vertnum = mesh_vertnum[dataid]
+ geom1.vert = mesh_vert
+
+ geom2 = Geom()
+ geom2.index = -1
+ geomtype2 = geom_type[gid2]
+ geom2.pos = geom_xpos_in[0, gid2]
+ geom2.rot = geom_xmat_in[0, gid2]
+ geom2.size = geom_size[0, gid2]
+ geom2.graphadr = -1
+
+ if geom_dataid[gid2] >= 0 and geom_type[gid2] == MESHGEOM:
+ dataid = geom_dataid[gid2]
+ geom2.vertadr = mesh_vertadr[dataid]
+ geom2.vertnum = mesh_vertnum[dataid]
+ geom2.vert = mesh_vert
+
+ x_1 = geom_xpos_in[0, gid1]
+ x_2 = geom_xpos_in[0, gid2]
+
+ (
+ dist,
+ x1,
+ x2,
+ ) = ccd(
+ 1e-6,
+ 1.0e30,
+ iterations,
+ iterations,
+ geom1,
+ geom2,
+ geomtype1,
+ geomtype2,
+ x_1,
+ x_2,
+ vert,
+ vert1,
+ vert2,
+ vert_index1,
+ vert_index2,
+ face,
+ face_pr,
+ face_norm2,
+ face_index,
+ face_map,
+ horizon,
+ )
+
+ dist_out[0] = dist
+ pos_out[0] = x1
+ pos_out[1] = x2
+
+ vert = wp.array(shape=(iterations,), dtype=wp.vec3)
+ vert1 = wp.array(shape=(iterations,), dtype=wp.vec3)
+ vert2 = wp.array(shape=(iterations,), dtype=wp.vec3)
+ vert_index1 = wp.array(shape=(iterations,), dtype=int)
+ vert_index2 = wp.array(shape=(iterations,), dtype=int)
+ face = wp.array(shape=(2 * iterations,), dtype=wp.vec3i)
+ face_pr = wp.array(shape=(2 * iterations,), dtype=wp.vec3)
+ face_norm2 = wp.array(shape=(2 * iterations,), dtype=float)
+ face_index = wp.array(shape=(2 * iterations,), dtype=int)
+ face_map = wp.array(shape=(2 * iterations,), dtype=int)
+ horizon = wp.array(shape=(2 * iterations,), dtype=int)
+ dist_out = wp.array(shape=(1,), dtype=float)
+ pos_out = wp.array(shape=(2,), dtype=wp.vec3)
+ wp.launch(
+ _gjk_kernel,
+ dim=(1,),
+ inputs=[
+ m.geom_type,
+ m.geom_dataid,
+ m.geom_size,
+ m.mesh_vertadr,
+ m.mesh_vertnum,
+ m.mesh_vert,
+ d.geom_xpos,
+ d.geom_xmat,
+ gid1,
+ gid2,
+ iterations,
+ vert,
+ vert1,
+ vert2,
+ vert_index1,
+ vert_index2,
+ face,
+ face_pr,
+ face_norm2,
+ face_index,
+ face_map,
+ horizon,
+ ],
+ outputs=[
+ dist_out,
+ pos_out,
+ ],
+ )
+ return dist_out.numpy()[0], pos_out.numpy()[0], pos_out.numpy()[1]
+
+
+class GJKTest(absltest.TestCase):
+ """Tests for GJK/EPA."""
+
+ def test_spheres_distance(self):
+ """Test distance between two spheres."""
+
+ _, _, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+
+ """
+ )
+
+ dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
+ self.assertEqual(1.0, dist)
+
+ def test_spheres_touching(self):
+ """Test two touching spheres have zero distance"""
+
+ _, _, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+
+ """
+ )
+
+ dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
+ self.assertEqual(0.0, dist)
+
+ def test_box_mesh_distance(self):
+ """Test distance between a mesh and box"""
+
+ _, _, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+
+
+
+
+ """
+ )
+
+ dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
+ self.assertAlmostEqual(0.1, dist)
+
+ def test_sphere_sphere_contact(self):
+ """Test penetration depth between two spheres."""
+
+ _, _, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+
+ """
+ )
+
+ # TODO(kbayes): use margin trick instead of EPA for penetration recovery
+ dist, _, _ = _geom_dist(m, d, 0, 1, 500)
+ self.assertAlmostEqual(-2, dist)
+
+ def test_box_box_contact(self):
+ """Test penetration between two boxes."""
+
+ _, _, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+
+ """
+ )
+ dist, x1, x2 = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
+ self.assertAlmostEqual(-1, dist)
+ normal = wp.normalize(x1 - x2)
+ self.assertAlmostEqual(normal[0], 1)
+ self.assertAlmostEqual(normal[1], 0)
+ self.assertAlmostEqual(normal[2], 0)
+
+ def test_mesh_mesh_contact(self):
+ """Test penetration between two meshes."""
+
+ _, _, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+
+
+
+
+
+
+ """
+ )
+ dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
+ self.assertAlmostEqual(-0.01, dist)
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py
new file mode 100644
index 00000000..8397700c
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py
@@ -0,0 +1,397 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+from typing import Tuple
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+
+
+@wp.func
+def _hfield_overlap_range(
+ # Model:
+ geom_dataid: wp.array(dtype=int),
+ geom_rbound: wp.array2d(dtype=float),
+ geom_margin: wp.array2d(dtype=float),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ # Data in:
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ # In:
+ hfieldid: int,
+ geomid: int,
+ worldid: int,
+) -> Tuple[int, int, int, int]:
+ """Returns min/max grid coordinates of height field cells overlapped by a geom's bounds.
+
+ Args:
+ geom_dataid: Array of geom data IDs
+ geom_rbound: Array of geom bounding radii
+ geom_margin: Array of geom margins
+ hfield_nrow: Array of heightfield rows
+ hfield_ncol: Array of heightfield columns
+ hfield_size: Array of heightfield sizes
+ geom_xpos_in: Array of geom positions
+ geom_xmat_in: Array of geom orientation matrices
+ hfieldid: Index of the height field geom
+ geomid: Index of the other geom
+ worldid: Current world index
+
+ Returns:
+ min_i, min_j, max_i, max_j: Grid coordinate bounds
+ """
+ # get height field dimensions
+ dataid = geom_dataid[hfieldid]
+ nrow = hfield_nrow[dataid]
+ ncol = hfield_ncol[dataid]
+ size = hfield_size[dataid] # (x, y, z_top, z_bottom)
+
+ # get positions and transforms
+ hf_pos = geom_xpos_in[worldid, hfieldid]
+ hf_mat = geom_xmat_in[worldid, hfieldid]
+ geom_pos = geom_xpos_in[worldid, geomid]
+
+ # transform geom_pos to height field local space
+ local_pos = wp.transpose(hf_mat) @ (geom_pos - hf_pos)
+
+ # get bounding radius of other geometry (including margin)
+ bound_radius = geom_rbound[worldid, geomid] + geom_margin[worldid, geomid]
+
+ # calculate grid resolution
+ x_scale = 2.0 * size[0] / float(ncol - 1)
+ y_scale = 2.0 * size[1] / float(nrow - 1)
+
+ # calculate min/max grid coordinates that could contain the object
+ min_i = wp.max(0, int((local_pos[0] - bound_radius + size[0]) / x_scale))
+ max_i = wp.min(ncol - 2, int((local_pos[0] + bound_radius + size[0]) / x_scale) + 1)
+ min_j = wp.max(0, int((local_pos[1] - bound_radius + size[1]) / y_scale))
+ max_j = wp.min(nrow - 2, int((local_pos[1] + bound_radius + size[1]) / y_scale) + 1)
+
+ return min_i, min_j, max_i, max_j
+
+
+@wp.func
+def hfield_triangle_prism(
+ # Model:
+ geom_dataid: wp.array(dtype=int),
+ hfield_adr: wp.array(dtype=int),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ hfield_data: wp.array(dtype=float),
+ # In:
+ hfieldid: int,
+ hftri_index: int,
+) -> wp.mat33:
+ """Returns the vertices of a triangular prism for a heightfield triangle.
+
+ Args:
+ geom_dataid: Array of geometry data IDs
+ hfield_adr: Array of heightfield addresses
+ hfield_nrow: Array of heightfield rows
+ hfield_ncol: Array of heightfield columns
+ hfield_size: Array of heightfield sizes
+ hfield_data: Array of heightfield data
+ hfieldid: Index of the height field geometry
+ hftri_index: Index of the triangle in the heightfield
+
+ Returns:
+ 3x3 matrix containing the vertices of the triangular prism
+ """
+ # https://mujoco.readthedocs.io/en/stable/XMLreference.html#asset-hfield
+
+ # get heightfield dimensions
+ dataid = geom_dataid[hfieldid]
+ if dataid < 0 or hftri_index < 0:
+ return wp.mat33(0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0)
+
+ nrow = hfield_nrow[dataid]
+ ncol = hfield_ncol[dataid]
+ size = hfield_size[dataid] # (x, y, z_top, z_bottom)
+
+ # calculate which triangle in the grid
+ row = (hftri_index // 2) // (ncol - 1)
+ col = (hftri_index // 2) % (ncol - 1)
+
+ # calculate vertices in 2D grid
+ x_scale = 2.0 * size[0] / float(ncol - 1)
+ y_scale = 2.0 * size[1] / float(nrow - 1)
+
+ # grid coordinates (i, j) for triangle corners
+ i0 = col
+ j0 = row
+ i1 = i0 + 1
+ j1 = j0 + 1
+
+ # convert grid coordinates to local space x, y coordinates
+ x0 = float(i0) * x_scale - size[0]
+ y0 = float(j0) * y_scale - size[1]
+ x1 = float(i1) * x_scale - size[0]
+ y1 = float(j1) * y_scale - size[1]
+
+ # get height values at corners from hfield_data
+ base_addr = hfield_adr[dataid]
+ z00 = hfield_data[base_addr + j0 * ncol + i0]
+ z01 = hfield_data[base_addr + j1 * ncol + i0]
+ z10 = hfield_data[base_addr + j0 * ncol + i1]
+ z11 = hfield_data[base_addr + j1 * ncol + i1]
+
+ # scale heights from range [0, 1] to [0, z_top]
+ z_top = size[2]
+ z00 = z00 * z_top
+ z01 = z01 * z_top
+ z10 = z10 * z_top
+ z11 = z11 * z_top
+
+ # set bottom z-value
+ z_bottom = -size[3]
+
+ # compress 6 prism vertices into 3x3 matrix, see hfield_prism_vertex for details
+ return wp.mat33(
+ x0,
+ y0,
+ z00,
+ x1,
+ y1,
+ z11,
+ wp.where(hftri_index % 2, 1.0, 0.0),
+ wp.where(hftri_index % 2, z10, z01),
+ z_bottom,
+ )
+
+
+@wp.func
+def hfield_prism_vertex(prism: wp.mat33, vert_index: int) -> wp.vec3:
+ """Extracts vertices from a compressed triangular prism representation.
+
+ The compression scheme stores a 6-vertex triangular prism using a 3x3 matrix:
+ - prism[0] = First vertex (x,y,z) - corner (i,j)
+ - prism[1] = Second vertex (x,y,z) - corner (i+1,j+1)
+ - prism[2,0] = Triangle type flag: 0 for even triangle (using corner (i,j+1)),
+ non-zero for odd triangle (using corner (i+1,j))
+ - prism[2,1] = Z-coordinate of the third vertex
+ - prism[2,2] = Z-coordinate used for all bottom vertices (common z)
+
+ In this way, we can reconstruct all 6 vertices of the prism by reusing
+ coordinates from the stored vertices.
+
+ Args:
+ prism: 3x3 compressed representation of a triangular prism
+ vert_index: Index of vertex to extract (0-5)
+
+ Returns:
+ The 3D coordinates of the requested vertex
+ """
+ if vert_index == 0 or vert_index == 1:
+ return prism[vert_index] # first two vertices stored directly
+
+ if vert_index == 2: # third vertex
+ if prism[2][0] == 0: # even triangle (i, j+1)
+ return wp.vec3(prism[0][0], prism[1][1], prism[2][1])
+ else: # odd triangle (i+1, j)
+ return wp.vec3(prism[1][0], prism[0][1], prism[2][1])
+
+ if vert_index == 3 or vert_index == 4: # bottom vertices below 0 and 1
+ return wp.vec3(prism[vert_index - 3][0], prism[vert_index - 3][1], prism[2][2])
+
+ if vert_index == 5: # bottom vertex below 2
+ if prism[2][0] == 0: # even triangle
+ return wp.vec3(prism[0][0], prism[1][1], prism[2][2])
+ else: # odd triangle
+ return wp.vec3(prism[1][0], prism[0][1], prism[2][2])
+
+
+@wp.kernel
+def _hfield_midphase(
+ # Model:
+ geom_type: wp.array(dtype=int),
+ geom_dataid: wp.array(dtype=int),
+ geom_rbound: wp.array2d(dtype=float),
+ geom_margin: wp.array2d(dtype=float),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ # Data in:
+ nconmax_in: int,
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ collision_pair_in: wp.array(dtype=wp.vec2i),
+ collision_hftri_index_in: wp.array(dtype=int),
+ collision_pairid_in: wp.array(dtype=int),
+ collision_worldid_in: wp.array(dtype=int),
+ # Data out:
+ collision_pair_out: wp.array(dtype=wp.vec2i),
+ collision_hftri_index_out: wp.array(dtype=int),
+ collision_pairid_out: wp.array(dtype=int),
+ collision_worldid_out: wp.array(dtype=int),
+ ncollision_out: wp.array(dtype=int),
+):
+ """Midphase collision detection for heightfield triangles with other geoms.
+
+ This kernel processes collision pairs where one geom is a heightfield (identified by
+ collision_hftri_index_in[pairid] == -1) and expands them into multiple collision pairs,
+ one for each potentially colliding triangle.
+
+ Args:
+ geom_type: Array of geometry types
+ geom_dataid: Array of geometry data IDs
+ geom_rbound: Array of geometry bounding radii
+ geom_margin: Array of geometry margins
+ hfield_nrow: Array of heightfield rows
+ hfield_ncol: Array of heightfield columns
+ hfield_size: Array of heightfield sizes
+ nconmax_in: Max number of collisions
+ geom_xpos_in: Array of geometry positions
+ geom_xmat_in: Array of geometry orientation matrices
+ collision_pair_in: Array of collision pairs
+ collision_hftri_index_in: Array of heightfield triangle indices, -1 for heightfield
+ pairs
+ collision_pairid_in: Array of collision pair IDs
+ collision_worldid_in: Array of collision world IDs
+
+ collision_pair_out: Output array of collision pairs
+ collision_hftri_index_out: Output array of heightfield triangle indices
+ collision_pairid_out: Output array of collision pair IDs
+ collision_worldid_out: Output array of collision world IDs
+ ncollision_out: Output counter for number of collisions
+ """
+ pairid = wp.tid()
+
+ # only process pairs that are marked for heightfield collision (-1)
+ # the buffer is cleared at the start of each frame in collision_driver.py
+ if collision_hftri_index_in[pairid] != -1:
+ return
+
+ # get the collision pair info
+ pair = collision_pair_in[pairid]
+ worldid = collision_worldid_in[pairid]
+ pair_id = collision_pairid_in[pairid]
+
+ # identify which geom is the heightfield
+ g1 = pair[0]
+ g2 = pair[1]
+
+ hfieldid = g1
+ geomid = g2
+
+ # if the first geom is not a heightfield, swap them
+ # in theory, shouldn't happen as _add_geom_pair already sorted the pair
+ if geom_type[g1] != int(GeomType.HFIELD.value):
+ hfieldid = g2
+ geomid = g1
+
+ # get min/max grid coordinates for overlap region
+ min_i, min_j, max_i, max_j = _hfield_overlap_range(
+ geom_dataid,
+ geom_rbound,
+ geom_margin,
+ hfield_nrow,
+ hfield_ncol,
+ hfield_size,
+ geom_xpos_in,
+ geom_xmat_in,
+ hfieldid,
+ geomid,
+ worldid,
+ )
+
+ # get hfield dimensions for triangle index calculation
+ dataid = geom_dataid[hfieldid]
+ ncol = hfield_ncol[dataid]
+
+ # loop through grid cells and add pairs for all triangles
+ for j in range(min_j, max_j + 1):
+ for i in range(min_i, max_i + 1):
+ # each grid cell contains two triangles
+ base_idx = ((j * (ncol - 1)) + i) * 2
+
+ # add both triangles from this cell
+ for t in range(2):
+ if i == 0 and j == 0 and t == 0:
+ # reuse the initial pair for the 1st triangle
+ new_pairid = pairid
+ else:
+ # for the rest create a new pair
+ new_pairid = wp.atomic_add(ncollision_out, 0, 1)
+
+ if new_pairid >= nconmax_in:
+ return
+
+ collision_pair_out[new_pairid] = pair
+ collision_hftri_index_out[new_pairid] = base_idx + t
+ collision_pairid_out[new_pairid] = pair_id
+ collision_worldid_out[new_pairid] = worldid
+
+
+def hfield_midphase(m: Model, d: Data):
+ """Midphase collision detection for heightfield triangles with other geoms.
+
+ Processes collision pairs from the broadphase where one geom is a heightfield and expands
+ them into multiple collision pairs, one for each potentially colliding triangle. The
+ function directly writes to the same collision buffers used by _add_geom_pair.
+
+ Args:
+ m: Model containing geometry and heightfield data
+ - geom_type: Array of geometry types
+ - geom_dataid: Array of geometry data IDs
+ - hfield_nrow: Array of heightfield rows
+ - hfield_ncol: Array of heightfield columns
+ - hfield_size: Array of heightfield sizes
+ - geom_rbound: Array of geometry bounding radii
+ - geom_margin: Array of geometry margins
+ d: Data containing current state and collision information
+ - nconmax: Maximum number of contacts
+ - geom_xpos: Array of geometry positions
+ - geom_xmat: Array of geometry orientation matrices
+ - collision_pair: Array of collision pairs
+ - collision_hftri_index: Array of heightfield triangle indices
+ - collision_pairid: Array of collision pair IDs
+ - collision_worldid: Array of collision world IDs
+ - ncollision: Number of collisions
+ """
+ # launch the midphase kernel to expand height field collision pairs
+ # write directly to the same buffers that _add_geom_pair writes to
+ wp.launch(
+ kernel=_hfield_midphase,
+ dim=d.nconmax, # launch threads to process all potential pairs
+ inputs=[
+ m.geom_type,
+ m.geom_dataid,
+ m.geom_rbound,
+ m.geom_margin,
+ m.hfield_nrow,
+ m.hfield_ncol,
+ m.hfield_size,
+ d.nconmax,
+ d.geom_xpos,
+ d.geom_xmat,
+ d.collision_pair,
+ d.collision_hftri_index,
+ d.collision_pairid,
+ d.collision_worldid,
+ ],
+ outputs=[
+ d.collision_pair,
+ d.collision_hftri_index,
+ d.collision_pairid,
+ d.collision_worldid,
+ d.ncollision,
+ ],
+ )
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py
new file mode 100644
index 00000000..b0e6de66
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py
@@ -0,0 +1,2949 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_triangle_prism
+from mujoco.mjx.third_party.mujoco_warp._src.math import closest_segment_point
+from mujoco.mjx.third_party.mujoco_warp._src.math import closest_segment_to_segment_points
+from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame
+from mujoco.mjx.third_party.mujoco_warp._src.math import normalize_with_norm
+from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index
+from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINMU
+from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec5
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
+
+wp.set_module_options({"enable_backward": False})
+
+
+class vec8f(wp.types.vector(length=8, dtype=wp.float32)):
+ pass
+
+
+class mat43f(wp.types.matrix(shape=(4, 3), dtype=wp.float32)):
+ pass
+
+
+class mat83f(wp.types.matrix(shape=(8, 3), dtype=wp.float32)):
+ pass
+
+
+@wp.struct
+class Geom:
+ pos: wp.vec3
+ rot: wp.mat33
+ normal: wp.vec3
+ size: wp.vec3
+ hfprism: wp.mat33
+ vertadr: int
+ vertnum: int
+ vert: wp.array(dtype=wp.vec3)
+ graphadr: int
+ graph: wp.array(dtype=int)
+ mesh_polynum: int
+ mesh_polyadr: int
+ mesh_polynormal: wp.array(dtype=wp.vec3)
+ mesh_polyvertadr: wp.array(dtype=int)
+ mesh_polyvertnum: wp.array(dtype=int)
+ mesh_polyvert: wp.array(dtype=int)
+ mesh_polymapadr: wp.array(dtype=int)
+ mesh_polymapnum: wp.array(dtype=int)
+ mesh_polymap: wp.array(dtype=int)
+ index: int
+
+
+@wp.func
+def _geom(
+ # Model:
+ geom_type: wp.array(dtype=int),
+ geom_dataid: wp.array(dtype=int),
+ geom_size: wp.array2d(dtype=wp.vec3),
+ hfield_adr: wp.array(dtype=int),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ hfield_data: wp.array(dtype=float),
+ mesh_vertadr: wp.array(dtype=int),
+ mesh_vertnum: wp.array(dtype=int),
+ mesh_vert: wp.array(dtype=wp.vec3),
+ mesh_graphadr: wp.array(dtype=int),
+ mesh_graph: wp.array(dtype=int),
+ mesh_polynum: wp.array(dtype=int),
+ mesh_polyadr: wp.array(dtype=int),
+ mesh_polynormal: wp.array(dtype=wp.vec3),
+ mesh_polyvertadr: wp.array(dtype=int),
+ mesh_polyvertnum: wp.array(dtype=int),
+ mesh_polyvert: wp.array(dtype=int),
+ mesh_polymapadr: wp.array(dtype=int),
+ mesh_polymapnum: wp.array(dtype=int),
+ mesh_polymap: wp.array(dtype=int),
+ # Data in:
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ # In:
+ worldid: int,
+ gid: int,
+ hftri_index: int,
+) -> Geom:
+ geom = Geom()
+ geom.pos = geom_xpos_in[worldid, gid]
+ rot = geom_xmat_in[worldid, gid]
+ geom.rot = rot
+ geom.size = geom_size[worldid, gid]
+ geom.normal = wp.vec3(rot[0, 2], rot[1, 2], rot[2, 2]) # plane
+ dataid = geom_dataid[gid]
+
+ # If geom is MESH, get mesh verts
+ if dataid >= 0 and geom_type[gid] == int(GeomType.MESH.value):
+ geom.vertadr = mesh_vertadr[dataid]
+ geom.vertnum = mesh_vertnum[dataid]
+ geom.graphadr = mesh_graphadr[dataid]
+ geom.mesh_polynum = mesh_polynum[dataid]
+ geom.mesh_polyadr = mesh_polyadr[dataid]
+ else:
+ geom.vertadr = -1
+ geom.vertnum = -1
+ geom.graphadr = -1
+ geom.mesh_polynum = -1
+ geom.mesh_polyadr = -1
+
+ if geom_type[gid] == int(GeomType.MESH.value):
+ geom.vert = mesh_vert
+ geom.graph = mesh_graph
+ geom.mesh_polynormal = mesh_polynormal
+ geom.mesh_polyvertadr = mesh_polyvertadr
+ geom.mesh_polyvertnum = mesh_polyvertnum
+ geom.mesh_polyvert = mesh_polyvert
+ geom.mesh_polymapadr = mesh_polymapadr
+ geom.mesh_polymapnum = mesh_polymapnum
+ geom.mesh_polymap = mesh_polymap
+
+ # If geom is HFIELD triangle, compute triangle prism verts
+ if geom_type[gid] == int(GeomType.HFIELD.value):
+ geom.hfprism = hfield_triangle_prism(
+ geom_dataid, hfield_adr, hfield_nrow, hfield_ncol, hfield_size, hfield_data, gid, hftri_index
+ )
+
+ geom.index = -1
+ return geom
+
+
+@wp.func
+def write_contact(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ dist_in: float,
+ pos_in: wp.vec3,
+ frame_in: wp.mat33,
+ margin_in: float,
+ gap_in: float,
+ condim_in: int,
+ friction_in: vec5,
+ solref_in: wp.vec2f,
+ solreffriction_in: wp.vec2f,
+ solimp_in: vec5,
+ geoms_in: wp.vec2i,
+ worldid_in: int,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ active = (dist_in - margin_in) < 0
+ if active:
+ cid = wp.atomic_add(ncon_out, 0, 1)
+ if cid < nconmax_in:
+ contact_dist_out[cid] = dist_in
+ contact_pos_out[cid] = pos_in
+ contact_frame_out[cid] = frame_in
+ contact_geom_out[cid] = geoms_in
+ contact_worldid_out[cid] = worldid_in
+ includemargin = margin_in - gap_in
+ contact_includemargin_out[cid] = includemargin
+ contact_dim_out[cid] = condim_in
+ contact_friction_out[cid] = friction_in
+ contact_solref_out[cid] = solref_in
+ contact_solreffriction_out[cid] = solreffriction_in
+ contact_solimp_out[cid] = solimp_in
+
+
+@wp.func
+def _plane_sphere(plane_normal: wp.vec3, plane_pos: wp.vec3, sphere_pos: wp.vec3, sphere_radius: float):
+ dist = wp.dot(sphere_pos - plane_pos, plane_normal) - sphere_radius
+ pos = sphere_pos - plane_normal * (sphere_radius + 0.5 * dist)
+ return dist, pos
+
+
+@wp.func
+def plane_sphere(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ plane: Geom,
+ sphere: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ dist, pos = _plane_sphere(plane.normal, plane.pos, sphere.pos, sphere.size[0])
+
+ write_contact(
+ nconmax_in,
+ dist,
+ pos,
+ make_frame(plane.normal),
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def _sphere_sphere(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ pos1: wp.vec3,
+ radius1: float,
+ pos2: wp.vec3,
+ radius2: float,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ dir = pos2 - pos1
+ dist = wp.length(dir)
+ if dist == 0.0:
+ n = wp.vec3(1.0, 0.0, 0.0)
+ else:
+ n = dir / dist
+ dist = dist - (radius1 + radius2)
+ pos = pos1 + n * (radius1 + 0.5 * dist)
+
+ write_contact(
+ nconmax_in,
+ dist,
+ pos,
+ make_frame(n),
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def _sphere_sphere_ext(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ pos1: wp.vec3,
+ radius1: float,
+ pos2: wp.vec3,
+ radius2: float,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ mat1: wp.mat33,
+ mat2: wp.mat33,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ dir = pos2 - pos1
+ dist = wp.length(dir)
+ if dist == 0.0:
+ # Use cross product of z axes like MuJoCo
+ axis1 = wp.vec3(mat1[0, 2], mat1[1, 2], mat1[2, 2])
+ axis2 = wp.vec3(mat2[0, 2], mat2[1, 2], mat2[2, 2])
+ n = wp.cross(axis1, axis2)
+ n = wp.normalize(n)
+ else:
+ n = dir / dist
+ dist = dist - (radius1 + radius2)
+ pos = pos1 + n * (radius1 + 0.5 * dist)
+
+ write_contact(
+ nconmax_in,
+ dist,
+ pos,
+ make_frame(n),
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def sphere_sphere(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ sphere1: Geom,
+ sphere2: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ _sphere_sphere(
+ nconmax_in,
+ sphere1.pos,
+ sphere1.size[0],
+ sphere2.pos,
+ sphere2.size[0],
+ worldid,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def sphere_capsule(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ sphere: Geom,
+ cap: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ """Calculates one contact between a sphere and a capsule."""
+ axis = wp.vec3(cap.rot[0, 2], cap.rot[1, 2], cap.rot[2, 2])
+ length = cap.size[1]
+ segment = axis * length
+
+ # Find closest point on capsule centerline to sphere center
+ pt = closest_segment_point(cap.pos - segment, cap.pos + segment, sphere.pos)
+
+ # Treat as sphere-sphere collision between sphere and closest point
+ _sphere_sphere(
+ nconmax_in,
+ sphere.pos,
+ sphere.size[0],
+ pt,
+ cap.size[0],
+ worldid,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def capsule_capsule(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ cap1: Geom,
+ cap2: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ axis1 = wp.vec3(cap1.rot[0, 2], cap1.rot[1, 2], cap1.rot[2, 2])
+ axis2 = wp.vec3(cap2.rot[0, 2], cap2.rot[1, 2], cap2.rot[2, 2])
+ length1 = cap1.size[1]
+ length2 = cap2.size[1]
+ seg1 = axis1 * length1
+ seg2 = axis2 * length2
+
+ pt1, pt2 = closest_segment_to_segment_points(
+ cap1.pos - seg1,
+ cap1.pos + seg1,
+ cap2.pos - seg2,
+ cap2.pos + seg2,
+ )
+
+ _sphere_sphere(
+ nconmax_in,
+ pt1,
+ cap1.size[0],
+ pt2,
+ cap2.size[0],
+ worldid,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def plane_capsule(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ plane: Geom,
+ cap: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ """Calculates two contacts between a capsule and a plane."""
+ n = plane.normal
+ axis = wp.vec3(cap.rot[0, 2], cap.rot[1, 2], cap.rot[2, 2])
+ # align contact frames with capsule axis
+ b, b_norm = normalize_with_norm(axis - n * wp.dot(n, axis))
+
+ if b_norm < 0.5:
+ if -0.5 < n[1] and n[1] < 0.5:
+ b = wp.vec3(0.0, 1.0, 0.0)
+ else:
+ b = wp.vec3(0.0, 0.0, 1.0)
+
+ c = wp.cross(n, b)
+ frame = wp.mat33(n[0], n[1], n[2], b[0], b[1], b[2], c[0], c[1], c[2])
+ segment = axis * cap.size[1]
+
+ dist1, pos1 = _plane_sphere(n, plane.pos, cap.pos + segment, cap.size[0])
+ write_contact(
+ nconmax_in,
+ dist1,
+ pos1,
+ frame,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+ dist2, pos2 = _plane_sphere(n, plane.pos, cap.pos - segment, cap.size[0])
+ write_contact(
+ nconmax_in,
+ dist2,
+ pos2,
+ frame,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def plane_ellipsoid(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ plane: Geom,
+ ellipsoid: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ sphere_support = -wp.normalize(wp.cw_mul(wp.transpose(ellipsoid.rot) @ plane.normal, ellipsoid.size))
+ pos = ellipsoid.pos + ellipsoid.rot @ wp.cw_mul(sphere_support, ellipsoid.size)
+ dist = wp.dot(plane.normal, pos - plane.pos)
+ pos = pos - plane.normal * dist * 0.5
+
+ write_contact(
+ nconmax_in,
+ dist,
+ pos,
+ make_frame(plane.normal),
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def plane_box(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ plane: Geom,
+ box: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ count = int(0)
+ corner = wp.vec3()
+ dist = wp.dot(box.pos - plane.pos, plane.normal)
+
+ # test all corners, pick bottom 4
+ for i in range(8):
+ # get corner in local coordinates
+ corner.x = wp.where(i & 1, box.size.x, -box.size.x)
+ corner.y = wp.where(i & 2, box.size.y, -box.size.y)
+ corner.z = wp.where(i & 4, box.size.z, -box.size.z)
+
+ # get corner in global coordinates relative to box center
+ corner = box.rot * corner
+
+ # compute distance to plane, skip if too far or pointing up
+ ldist = wp.dot(plane.normal, corner)
+ if dist + ldist > margin or ldist > 0:
+ continue
+
+ cdist = dist + ldist
+ frame = make_frame(plane.normal)
+ pos = corner + box.pos + (plane.normal * cdist / -2.0)
+ write_contact(
+ nconmax_in,
+ cdist,
+ pos,
+ frame,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+ count += 1
+ if count >= 4:
+ break
+
+
+_HUGE_VAL = 1e6
+
+
+@wp.func
+def plane_convex(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ plane: Geom,
+ convex: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ """Calculates contacts between a plane and a convex object."""
+
+ # get points in the convex frame
+ plane_pos = wp.transpose(convex.rot) @ (plane.pos - convex.pos)
+ n = wp.transpose(convex.rot) @ plane.normal
+
+ # Store indices in vec4
+ indices = wp.vec4i(-1, -1, -1, -1)
+
+ # exhaustive search over all vertices
+ if convex.graphadr == -1 or convex.vertnum < 10:
+ # Find support points
+ max_support = wp.float32(-_HUGE_VAL)
+ for i in range(convex.vertnum):
+ support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n)
+ max_support = wp.max(support, max_support)
+
+ threshold = wp.max(0.0, max_support - 1e-3)
+ # Find point a (first support point)
+ a_dist = wp.float32(-_HUGE_VAL)
+ for i in range(convex.vertnum):
+ support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n)
+ dist = wp.where(support > threshold, 0.0, -_HUGE_VAL)
+ if dist > a_dist:
+ indices[0] = i
+ a_dist = dist
+ a = convex.vert[convex.vertadr + indices[0]]
+
+ # Find point b (furthest from a)
+ b_dist = wp.float32(-_HUGE_VAL)
+ for i in range(convex.vertnum):
+ support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n)
+ dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL)
+ dist = wp.length_sq(a - convex.vert[convex.vertadr + i]) + dist_mask
+ if dist > b_dist:
+ indices[1] = i
+ b_dist = dist
+ b = convex.vert[convex.vertadr + indices[1]]
+
+ # Find point c (furthest along axis orthogonal to a-b)
+ ab = wp.cross(n, a - b)
+ c_dist = wp.float32(-_HUGE_VAL)
+ for i in range(convex.vertnum):
+ support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n)
+ dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL)
+ dist = wp.length_sq(ab - convex.vert[convex.vertadr + i]) + dist_mask
+ if dist > c_dist:
+ indices[2] = i
+ c_dist = dist
+ c = convex.vert[convex.vertadr + indices[2]]
+
+ # Find point d (furthest from other triangle edges)
+ ac = wp.cross(n, a - c)
+ bc = wp.cross(n, b - c)
+ d_dist = wp.float32(-_HUGE_VAL)
+ for i in range(convex.vertnum):
+ support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n)
+ dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL)
+ ap = ac - convex.vert[convex.vertadr + i]
+ bp = bc - convex.vert[convex.vertadr + i]
+ dist_ap = wp.abs(wp.dot(ap, ac)) + dist_mask
+ dist_bp = wp.abs(wp.dot(bp, bc)) + dist_mask
+ if dist_ap + dist_bp > d_dist:
+ indices[3] = i
+ d_dist = dist_ap + dist_bp
+
+ else:
+ numvert = convex.graph[convex.graphadr]
+ vert_edgeadr = convex.graphadr + 2
+ vert_globalid = convex.graphadr + 2 + numvert
+ edge_localid = convex.graphadr + 2 + 2 * numvert
+
+ # Find support points
+ max_support = wp.float32(-_HUGE_VAL)
+
+ # hillclimb until no change
+ prev = int(-1)
+ imax = int(0)
+
+ while True:
+ prev = int(imax)
+ i = int(convex.graph[vert_edgeadr + imax])
+ while convex.graph[edge_localid + i] >= 0:
+ subidx = convex.graph[edge_localid + i]
+ idx = convex.graph[vert_globalid + subidx]
+ support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n)
+ if support > max_support:
+ max_support = support
+ imax = int(subidx)
+ i += int(1)
+ if imax == prev:
+ break
+
+ threshold = wp.max(0.0, max_support - 1e-3)
+
+ a_dist = wp.float32(-_HUGE_VAL)
+ # hillclimb until no change
+ prev = int(-1)
+ imax = int(0)
+
+ while True:
+ prev = int(imax)
+ i = int(convex.graph[vert_edgeadr + imax])
+ while convex.graph[edge_localid + i] >= 0:
+ subidx = convex.graph[edge_localid + i]
+ idx = convex.graph[vert_globalid + subidx]
+ support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n)
+ dist = wp.where(support > threshold, 0.0, -_HUGE_VAL)
+ if dist > a_dist:
+ a_dist = dist
+ imax = int(subidx)
+ i += int(1)
+ if imax == prev:
+ break
+ imax = convex.graph[vert_globalid + imax]
+ a = convex.vert[convex.vertadr + imax]
+ indices[0] = imax
+
+ # Find point b (furthest from a)
+ b_dist = wp.float32(-_HUGE_VAL)
+ # hillclimb until no change
+ prev = int(-1)
+ imax = int(0)
+
+ while True:
+ prev = int(imax)
+ i = int(convex.graph[vert_edgeadr + imax])
+ while convex.graph[edge_localid + i] >= 0:
+ subidx = convex.graph[edge_localid + i]
+ idx = convex.graph[vert_globalid + subidx]
+ support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n)
+ dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL)
+ dist = wp.length_sq(a - convex.vert[convex.vertadr + idx]) + dist_mask
+ if dist > b_dist:
+ b_dist = dist
+ imax = int(subidx)
+ i += int(1)
+ if imax == prev:
+ break
+ imax = convex.graph[vert_globalid + imax]
+ b = convex.vert[convex.vertadr + imax]
+ indices[1] = imax
+
+ # Find point c (furthest along axis orthogonal to a-b)
+ ab = wp.cross(n, a - b)
+ c_dist = wp.float32(-_HUGE_VAL)
+ # hillclimb until no change
+ prev = int(-1)
+ imax = int(0)
+
+ while True:
+ prev = int(imax)
+ i = int(convex.graph[vert_edgeadr + imax])
+ while convex.graph[edge_localid + i] >= 0:
+ subidx = convex.graph[edge_localid + i]
+ idx = convex.graph[vert_globalid + subidx]
+ support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n)
+ dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL)
+ dist = wp.length_sq(ab - convex.vert[convex.vertadr + idx]) + dist_mask
+ if dist > c_dist:
+ c_dist = dist
+ imax = int(subidx)
+ i += int(1)
+ if imax == prev:
+ break
+ imax = convex.graph[vert_globalid + imax]
+ c = convex.vert[convex.vertadr + imax]
+ indices[2] = imax
+
+ # Find point d (furthest from other triangle edges)
+ ac = wp.cross(n, a - c)
+ bc = wp.cross(n, b - c)
+ d_dist = wp.float32(-_HUGE_VAL)
+ # hillclimb until no change
+ prev = int(-1)
+ imax = int(0)
+
+ while True:
+ prev = int(imax)
+ i = int(convex.graph[vert_edgeadr + imax])
+ while convex.graph[edge_localid + i] >= 0:
+ subidx = convex.graph[edge_localid + i]
+ idx = convex.graph[vert_globalid + subidx]
+ support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n)
+ dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL)
+ ap = ac - convex.vert[convex.vertadr + idx]
+ bp = bc - convex.vert[convex.vertadr + idx]
+ dist_ap = wp.abs(wp.dot(ap, ac)) + dist_mask
+ dist_bp = wp.abs(wp.dot(bp, bc)) + dist_mask
+ if dist_ap + dist_bp > d_dist:
+ d_dist = dist_ap + dist_bp
+ imax = int(subidx)
+ i += int(1)
+ if imax == prev:
+ break
+ imax = convex.graph[vert_globalid + imax]
+ indices[3] = imax
+
+ # Write contacts
+ frame = make_frame(plane.normal)
+ for i in range(3, -1, -1):
+ idx = indices[i]
+ count = int(0)
+ for j in range(i + 1):
+ if indices[j] == idx:
+ count = count + 1
+
+ # Check if the index is unique (appears exactly once)
+ if count == 1:
+ pos = convex.vert[convex.vertadr + idx]
+ pos = convex.pos + convex.rot @ pos
+ support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n)
+ dist = -support
+ pos = pos - 0.5 * dist * plane.normal
+ write_contact(
+ nconmax_in,
+ dist,
+ pos,
+ frame,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def sphere_cylinder(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ sphere: Geom,
+ cylinder: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ axis = wp.vec3(
+ cylinder.rot[0, 2],
+ cylinder.rot[1, 2],
+ cylinder.rot[2, 2],
+ )
+
+ vec = sphere.pos - cylinder.pos
+ x = wp.dot(vec, axis)
+
+ a_proj = axis * x
+ p_proj = vec - a_proj
+ p_proj_sqr = wp.dot(p_proj, p_proj)
+
+ collide_side = wp.abs(x) < cylinder.size[1]
+ collide_cap = p_proj_sqr < (cylinder.size[0] * cylinder.size[0])
+
+ if collide_side and collide_cap:
+ dist_cap = cylinder.size[1] - wp.abs(x)
+ dist_radius = cylinder.size[0] - wp.sqrt(p_proj_sqr)
+
+ if dist_cap < dist_radius:
+ collide_side = False
+ else:
+ collide_cap = False
+
+ # Side collision
+ if collide_side:
+ pos_target = cylinder.pos + a_proj
+
+ _sphere_sphere_ext(
+ nconmax_in,
+ sphere.pos,
+ sphere.size[0],
+ pos_target,
+ cylinder.size[0],
+ worldid,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ sphere.rot,
+ cylinder.rot,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+ return
+
+ # Cap collision
+ if collide_cap:
+ if x > 0.0:
+ # top cap
+ pos_cap = cylinder.pos + axis * cylinder.size[1]
+ plane_normal = axis
+ else:
+ # bottom cap
+ pos_cap = cylinder.pos - axis * cylinder.size[1]
+ plane_normal = -axis
+
+ dist, pos_contact = _plane_sphere(plane_normal, pos_cap, sphere.pos, sphere.size[0])
+ plane_normal = -plane_normal # Flip normal after position calculation
+
+ write_contact(
+ nconmax_in,
+ dist,
+ pos_contact,
+ make_frame(plane_normal),
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+ return
+
+ # Corner collision
+ inv_len = 1.0 / wp.sqrt(p_proj_sqr)
+ p_proj = p_proj * (cylinder.size[0] * inv_len)
+
+ cap_offset = axis * (wp.sign(x) * cylinder.size[1])
+ pos_corner = cylinder.pos + cap_offset + p_proj
+
+ _sphere_sphere_ext(
+ nconmax_in,
+ sphere.pos,
+ sphere.size[0],
+ pos_corner,
+ 0.0,
+ worldid,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ sphere.rot,
+ cylinder.rot,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def plane_cylinder(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ plane: Geom,
+ cylinder: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ """Calculates contacts between a cylinder and a plane."""
+ # Extract plane normal and cylinder axis
+ n = plane.normal
+ axis = wp.vec3(cylinder.rot[0, 2], cylinder.rot[1, 2], cylinder.rot[2, 2])
+
+ # Project, make sure axis points toward plane
+ prjaxis = wp.dot(n, axis)
+ if prjaxis > 0:
+ axis = -axis
+ prjaxis = -prjaxis
+
+ # Compute normal distance from plane to cylinder center
+ dist0 = wp.dot(cylinder.pos - plane.pos, n)
+
+ # Remove component of -normal along cylinder axis
+ vec = axis * prjaxis - n
+ len_sqr = wp.dot(vec, vec)
+
+ # If vector is nondegenerate, normalize and scale by radius
+ # Otherwise use cylinder's x-axis scaled by radius
+ vec = wp.where(
+ len_sqr >= 1e-12,
+ vec * (cylinder.size[0] / wp.sqrt(len_sqr)),
+ wp.vec3(cylinder.rot[0, 0], cylinder.rot[1, 0], cylinder.rot[2, 0]) * cylinder.size[0],
+ )
+
+ # Project scaled vector on normal
+ prjvec = wp.dot(vec, n)
+
+ # Scale cylinder axis by half-length
+ axis = axis * cylinder.size[1]
+ prjaxis = prjaxis * cylinder.size[1]
+
+ frame = make_frame(n)
+
+ # First contact point (end cap closer to plane)
+ dist1 = dist0 + prjaxis + prjvec
+ if dist1 <= margin:
+ pos1 = cylinder.pos + vec + axis - n * (dist1 * 0.5)
+ write_contact(
+ nconmax_in,
+ dist1,
+ pos1,
+ frame,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+ else:
+ # If nearest point is above margin, no contacts
+ return
+
+ # Second contact point (end cap farther from plane)
+ dist2 = dist0 - prjaxis + prjvec
+ if dist2 <= margin:
+ pos2 = cylinder.pos + vec - axis - n * (dist2 * 0.5)
+ write_contact(
+ nconmax_in,
+ dist2,
+ pos2,
+ make_frame(plane.normal),
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+ # Try triangle contact points on side closer to plane
+ prjvec1 = -prjvec * 0.5
+ dist3 = dist0 + prjaxis + prjvec1
+ if dist3 <= margin:
+ # Compute sideways vector scaled by radius*sqrt(3)/2
+ vec1 = wp.cross(vec, axis)
+ vec1 = wp.normalize(vec1) * (cylinder.size[0] * wp.sqrt(3.0) * 0.5)
+
+ # Add contact point A - adjust to closest side
+ pos3 = cylinder.pos + vec1 + axis - vec * 0.5 - n * (dist3 * 0.5)
+ write_contact(
+ nconmax_in,
+ dist3,
+ pos3,
+ frame,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+ # Add contact point B - adjust to closest side
+ pos4 = cylinder.pos - vec1 + axis - vec * 0.5 - n * (dist3 * 0.5)
+ write_contact(
+ nconmax_in,
+ dist3,
+ pos4,
+ frame,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def contact_params(
+ # Model:
+ geom_condim: wp.array(dtype=int),
+ geom_priority: wp.array(dtype=int),
+ geom_solmix: wp.array2d(dtype=float),
+ geom_solref: wp.array2d(dtype=wp.vec2),
+ geom_solimp: wp.array2d(dtype=vec5),
+ geom_friction: wp.array2d(dtype=wp.vec3),
+ geom_margin: wp.array2d(dtype=float),
+ geom_gap: wp.array2d(dtype=float),
+ pair_dim: wp.array(dtype=int),
+ pair_solref: wp.array2d(dtype=wp.vec2),
+ pair_solreffriction: wp.array2d(dtype=wp.vec2),
+ pair_solimp: wp.array2d(dtype=vec5),
+ pair_margin: wp.array2d(dtype=float),
+ pair_gap: wp.array2d(dtype=float),
+ pair_friction: wp.array2d(dtype=vec5),
+ # Data in:
+ collision_pair_in: wp.array(dtype=wp.vec2i),
+ collision_pairid_in: wp.array(dtype=int),
+ # In:
+ cid: int,
+ worldid: int,
+):
+ geoms = collision_pair_in[cid]
+ pairid = collision_pairid_in[cid]
+
+ if pairid > -1:
+ margin = pair_margin[worldid, pairid]
+ gap = pair_gap[worldid, pairid]
+ condim = pair_dim[pairid]
+ friction = pair_friction[worldid, pairid]
+ solref = pair_solref[worldid, pairid]
+ solreffriction = pair_solreffriction[worldid, pairid]
+ solimp = pair_solimp[worldid, pairid]
+ else:
+ g1 = geoms[0]
+ g2 = geoms[1]
+
+ p1 = geom_priority[g1]
+ p2 = geom_priority[g2]
+
+ solmix1 = geom_solmix[worldid, g1]
+ solmix2 = geom_solmix[worldid, g2]
+
+ mix = solmix1 / (solmix1 + solmix2)
+ mix = wp.where((solmix1 < MJ_MINVAL) and (solmix2 < MJ_MINVAL), 0.5, mix)
+ mix = wp.where((solmix1 < MJ_MINVAL) and (solmix2 >= MJ_MINVAL), 0.0, mix)
+ mix = wp.where((solmix1 >= MJ_MINVAL) and (solmix2 < MJ_MINVAL), 1.0, mix)
+ mix = wp.where(p1 == p2, mix, wp.where(p1 > p2, 1.0, 0.0))
+
+ margin = wp.max(geom_margin[worldid, g1], geom_margin[worldid, g2])
+ gap = wp.max(geom_gap[worldid, g1], geom_gap[worldid, g2])
+
+ condim1 = geom_condim[g1]
+ condim2 = geom_condim[g2]
+ condim = wp.where(p1 == p2, wp.max(condim1, condim2), wp.where(p1 > p2, condim1, condim2))
+
+ max_geom_friction = wp.max(geom_friction[worldid, g1], geom_friction[worldid, g2])
+ friction = vec5(
+ wp.max(MJ_MINMU, max_geom_friction[0]),
+ wp.max(MJ_MINMU, max_geom_friction[0]),
+ wp.max(MJ_MINMU, max_geom_friction[1]),
+ wp.max(MJ_MINMU, max_geom_friction[2]),
+ wp.max(MJ_MINMU, max_geom_friction[2]),
+ )
+
+ if geom_solref[worldid, g1].x > 0.0 and geom_solref[worldid, g2].x > 0.0:
+ solref = mix * geom_solref[worldid, g1] + (1.0 - mix) * geom_solref[worldid, g2]
+ else:
+ solref = wp.min(geom_solref[worldid, g1], geom_solref[worldid, g2])
+
+ solreffriction = wp.vec2(0.0, 0.0)
+
+ solimp = mix * geom_solimp[worldid, g1] + (1.0 - mix) * geom_solimp[worldid, g2]
+
+ return geoms, margin, gap, condim, friction, solref, solreffriction, solimp
+
+
+@wp.func
+def _sphere_box(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ sphere_pos: wp.vec3,
+ sphere_size: float,
+ box_pos: wp.vec3,
+ box_rot: wp.mat33,
+ box_size: wp.vec3,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ center = wp.transpose(box_rot) @ (sphere_pos - box_pos)
+
+ clamped = wp.max(-box_size, wp.min(box_size, center))
+ clamped_dir, dist = normalize_with_norm(clamped - center)
+
+ if dist - sphere_size > margin:
+ return
+
+ # sphere center inside box
+ if dist <= MJ_MINVAL:
+ closest = 2.0 * (box_size[0] + box_size[1] + box_size[2])
+ k = wp.int32(0)
+ for i in range(6):
+ face_dist = wp.abs(wp.where(i % 2, 1.0, -1.0) * box_size[i / 2] - center[i / 2])
+ if closest > face_dist:
+ closest = face_dist
+ k = i
+
+ nearest = wp.vec3(0.0)
+ nearest[k / 2] = wp.where(k % 2, -1.0, 1.0)
+ pos = center + nearest * (sphere_size - closest) / 2.0
+ contact_normal = box_rot @ nearest
+ contact_dist = -closest - sphere_size
+
+ else:
+ deepest = center + clamped_dir * sphere_size
+ pos = 0.5 * (clamped + deepest)
+ contact_normal = box_rot @ clamped_dir
+ contact_dist = dist - sphere_size
+
+ contact_pos = box_pos + box_rot @ pos
+ write_contact(
+ nconmax_in,
+ contact_dist,
+ contact_pos,
+ make_frame(contact_normal),
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def sphere_box(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ sphere: Geom,
+ box: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ _sphere_box(
+ nconmax_in,
+ sphere.pos,
+ sphere.size[0],
+ box.pos,
+ box.rot,
+ box.size,
+ worldid,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def capsule_box(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ cap: Geom,
+ box: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ """Calculates contacts between a capsule and a box."""
+ # Based on the mjc implementation
+ boxmatT = wp.transpose(box.rot)
+ pos = boxmatT @ (cap.pos - box.pos)
+ axis = boxmatT @ wp.vec3(cap.rot[0, 2], cap.rot[1, 2], cap.rot[2, 2])
+ halfaxis = axis * cap.size[1] # halfaxis is the capsule direction
+ axisdir = wp.int32(halfaxis[0] > 0.0) + 2 * wp.int32(halfaxis[1] > 0.0) + 4 * wp.int32(halfaxis[2] > 0.0)
+
+ bestdistmax = margin + 2.0 * (cap.size[0] + cap.size[1] + box.size[0] + box.size[1] + box.size[2])
+
+ # keep track of closest point
+ bestdist = wp.float32(bestdistmax)
+ bestsegmentpos = wp.float32(-12)
+
+ # cltype: encoded collision configuration
+ # cltype / 3 == 0 : lower corner is closest to the capsule
+ # == 2 : upper corner is closest to the capsule
+ # == 1 : middle of the edge is closest to the capsule
+ # cltype % 3 == 0 : lower corner is closest to the box
+ # == 2 : upper corner is closest to the box
+ # == 1 : middle of the capsule is closest to the box
+ cltype = wp.int32(-4)
+
+ # clface: index of the closest face of the box to the capsule
+ # -1: no face is closest (edge or corner is closest)
+ # 0, 1, 2: index of the axis perpendicular to the closest face
+ clface = wp.int32(-12)
+
+ # first: consider cases where a face of the box is closest
+ for i in range(-1, 2, 2):
+ axisTip = pos + wp.float32(i) * halfaxis
+ boxPoint = wp.vec3(axisTip)
+
+ n_out = wp.int32(0)
+ ax_out = wp.int32(-1)
+
+ for j in range(3):
+ if boxPoint[j] < -box.size[j]:
+ n_out += 1
+ ax_out = j
+ boxPoint[j] = -box.size[j]
+ elif boxPoint[j] > box.size[j]:
+ n_out += 1
+ ax_out = j
+ boxPoint[j] = box.size[j]
+
+ if n_out > 1:
+ continue
+
+ dist = wp.length_sq(boxPoint - axisTip)
+
+ if dist < bestdist:
+ bestdist = dist
+ bestsegmentpos = wp.float32(i)
+ cltype = -2 + i
+ clface = ax_out
+
+ # second: consider cases where an edge of the box is closest
+ clcorner = wp.int32(-123) # which corner is the closest
+ cledge = wp.int32(-123) # which axis
+ bestboxpos = wp.float32(0.0)
+
+ for i in range(8):
+ for j in range(3):
+ if i & (1 << j) != 0:
+ continue
+
+ c2 = wp.int32(-123)
+
+ # box_pt is the starting point (corner) on the box
+ box_pt = wp.cw_mul(
+ wp.vec3(
+ wp.where(i & 1, 1.0, -1.0),
+ wp.where(i & 2, 1.0, -1.0),
+ wp.where(i & 4, 1.0, -1.0),
+ ),
+ box.size,
+ )
+ box_pt[j] = 0.0
+
+ # find closest point between capsule and the edge
+ dif = box_pt - pos
+
+ u = -box.size[j] * dif[j]
+ v = wp.dot(halfaxis, dif)
+ ma = box.size[j] * box.size[j]
+ mb = -box.size[j] * halfaxis[j]
+ mc = cap.size[1] * cap.size[1]
+ det = ma * mc - mb * mb
+ if wp.abs(det) < MJ_MINVAL:
+ continue
+
+ idet = 1.0 / det
+ # sX : X=1 means middle of segment. X=0 or 2 one or the other end
+
+ x1 = wp.float32((mc * u - mb * v) * idet)
+ x2 = wp.float32((ma * v - mb * u) * idet)
+
+ s1 = wp.int32(1)
+ s2 = wp.int32(1)
+
+ if x1 > 1:
+ x1 = 1.0
+ s1 = 2
+ x2 = (v - mb) / mc
+ elif x1 < -1:
+ x1 = -1.0
+ s1 = 0
+ x2 = (v + mb) / mc
+
+ x2_over = x2 > 1.0
+ if x2_over or x2 < -1.0:
+ if x2_over:
+ x2 = 1.0
+ s2 = 2
+ x1 = (u - mb) / ma
+ else:
+ x2 = -1.0
+ s2 = 0
+ x1 = (u + mb) / ma
+
+ if x1 > 1:
+ x1 = 1.0
+ s1 = 2
+ elif x1 < -1:
+ x1 = -1.0
+ s1 = 0
+
+ dif -= halfaxis * x2
+ dif[j] += box.size[j] * x1
+
+ # encode relative positions of the closest points
+ ct = s1 * 3 + s2
+
+ dif_sq = wp.length_sq(dif)
+ if dif_sq < bestdist - MJ_MINVAL:
+ bestdist = dif_sq
+ bestsegmentpos = x2
+ bestboxpos = x1
+ # ct<6 means closest point on box is at lower end or middle of edge
+ c2 = ct / 6
+
+ clcorner = i + (1 << j) * c2 # index of closest box corner
+ cledge = j # axis index of closest box edge
+ cltype = ct # encoded collision configuration
+
+ best = wp.float32(0.0)
+
+ p = wp.vec2(pos.x, pos.y)
+ dd = wp.vec2(halfaxis.x, halfaxis.y)
+ s = wp.vec2(box.size.x, box.size.y)
+ secondpos = wp.float32(-4.0)
+
+ uu = dd.x * s.y
+ vv = dd.y * s.x
+ w_neg = dd.x * p.y - dd.y * p.x < 0
+
+ best = wp.float32(-1.0)
+
+ ee1 = uu - vv
+ ee2 = uu + vv
+
+ if wp.abs(ee1) > best:
+ best = wp.abs(ee1)
+ c1 = wp.where((ee1 < 0) == w_neg, 0, 3)
+
+ if wp.abs(ee2) > best:
+ best = wp.abs(ee2)
+ c1 = wp.where((ee2 > 0) == w_neg, 1, 2)
+
+ if cltype == -4: # invalid type
+ return
+
+ if cltype >= 0 and cltype / 3 != 1: # closest to a corner of the box
+ c1 = axisdir ^ clcorner
+ # Calculate relative orientation between capsule and corner
+ # There are two possible configurations:
+ # 1. Capsule axis points toward/away from corner
+ # 2. Capsule axis aligns with a face or edge
+ if c1 != 0 and c1 != 7: # create second contact point
+ if c1 == 1 or c1 == 2 or c1 == 4:
+ mul = 1
+ else:
+ mul = -1
+ c1 = 7 - c1
+
+ # "de" and "dp" distance from first closest point on the capsule to both ends of it
+ # mul is a direction along the capsule's axis
+
+ if c1 == 1:
+ ax = 0
+ ax1 = 1
+ ax2 = 2
+ elif c1 == 2:
+ ax = 1
+ ax1 = 2
+ ax2 = 0
+ elif c1 == 4:
+ ax = 2
+ ax1 = 0
+ ax2 = 1
+
+ if axis[ax] * axis[ax] > 0.5: # second point along the edge of the box
+ m = 2.0 * box.size[ax] / wp.abs(halfaxis[ax])
+ secondpos = min(1.0 - wp.float32(mul) * bestsegmentpos, m)
+ else: # second point along a face of the box
+ # check for overshoot again
+ m = 2.0 * min(
+ box.size[ax1] / wp.abs(halfaxis[ax1]),
+ box.size[ax2] / wp.abs(halfaxis[ax2]),
+ )
+ secondpos = -min(1.0 + wp.float32(mul) * bestsegmentpos, m)
+ secondpos *= wp.float32(mul)
+
+ elif cltype >= 0 and cltype / 3 == 1: # we are on box's edge
+ # Calculate relative orientation between capsule and edge
+ # Two possible configurations:
+ # - T configuration: c1 = 2^n (no additional contacts)
+ # - X configuration: c1 != 2^n (potential additional contacts)
+ c1 = axisdir ^ clcorner
+ c1 &= 7 - (1 << cledge) # mask out edge axis to determine configuration
+
+ if c1 == 1 or c1 == 2 or c1 == 4: # create second contact point
+ if cledge == 0:
+ ax1 = 1
+ ax2 = 2
+ if cledge == 1:
+ ax1 = 2
+ ax2 = 0
+ if cledge == 2:
+ ax1 = 0
+ ax2 = 1
+ ax = cledge
+
+ # find which face the capsule has a lower angle, and switch the axis
+ if wp.abs(axis[ax1]) > wp.abs(axis[ax2]):
+ ax1 = ax2
+ ax2 = 3 - ax - ax1
+
+ # mul determines direction along capsule axis for second contact point
+ if c1 & (1 << ax2):
+ mul = 1
+ secondpos = 1.0 - bestsegmentpos
+ else:
+ mul = -1
+ secondpos = 1.0 + bestsegmentpos
+
+ # now find out whether we point towards the opposite side or towards one of the sides
+ # and also find the farthest point along the capsule that is above the box
+
+ e1 = 2.0 * box.size[ax2] / wp.abs(halfaxis[ax2])
+ secondpos = min(e1, secondpos)
+
+ if ((axisdir & (1 << ax)) != 0) == ((c1 & (1 << ax2)) != 0):
+ e2 = 1.0 - bestboxpos
+ else:
+ e2 = 1.0 + bestboxpos
+
+ e1 = box.size[ax] * e2 / wp.abs(halfaxis[ax])
+
+ secondpos = min(e1, secondpos)
+ secondpos *= wp.float32(mul)
+
+ elif cltype < 0:
+ # similarly we handle the case when one capsule's end is closest to a face of the box
+ # and find where is the other end pointing to and clamping to the farthest point
+ # of the capsule that's above the box
+ # if the closest point is inside the box there's no need for a second point
+
+ if clface != -1: # create second contact point
+ mul = wp.where(cltype == -3, 1, -1)
+ secondpos = 2.0
+
+ tmp1 = pos - halfaxis * wp.float32(mul)
+
+ for i in range(3):
+ if i != clface:
+ ha_r = wp.float32(mul) / halfaxis[i]
+ e1 = (box.size[i] - tmp1[i]) * ha_r
+ if 0 < e1 and e1 < secondpos:
+ secondpos = e1
+
+ e1 = (-box.size[i] - tmp1[i]) * ha_r
+ if 0 < e1 and e1 < secondpos:
+ secondpos = e1
+
+ secondpos *= wp.float32(mul)
+
+ # create sphere in original orientation at first contact point
+ s1_pos_l = pos + halfaxis * bestsegmentpos
+ s1_pos_g = box.rot @ s1_pos_l + box.pos
+
+ # collide with sphere
+ _sphere_box(
+ nconmax_in,
+ s1_pos_g,
+ cap.size[0],
+ box.pos,
+ box.rot,
+ box.size,
+ worldid,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+ if secondpos > -3: # secondpos was modified
+ s2_pos_l = pos + halfaxis * (secondpos + bestsegmentpos)
+ s2_pos_g = box.rot @ s2_pos_l + box.pos
+ _sphere_box(
+ nconmax_in,
+ s2_pos_g,
+ cap.size[0],
+ box.pos,
+ box.rot,
+ box.size,
+ worldid,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@wp.func
+def _compute_rotmore(face_idx: int) -> wp.mat33:
+ rotmore = wp.mat33(0.0)
+
+ if face_idx == 0:
+ rotmore[0, 2] = -1.0
+ rotmore[1, 1] = +1.0
+ rotmore[2, 0] = +1.0
+ elif face_idx == 1:
+ rotmore[0, 0] = +1.0
+ rotmore[1, 2] = -1.0
+ rotmore[2, 1] = +1.0
+ elif face_idx == 2:
+ rotmore[0, 0] = +1.0
+ rotmore[1, 1] = +1.0
+ rotmore[2, 2] = +1.0
+ elif face_idx == 3:
+ rotmore[0, 2] = +1.0
+ rotmore[1, 1] = +1.0
+ rotmore[2, 0] = -1.0
+ elif face_idx == 4:
+ rotmore[0, 0] = +1.0
+ rotmore[1, 2] = +1.0
+ rotmore[2, 1] = -1.0
+ elif face_idx == 5:
+ rotmore[0, 0] = -1.0
+ rotmore[1, 1] = +1.0
+ rotmore[2, 2] = -1.0
+
+ return rotmore
+
+
+@wp.func
+def box_box(
+ # Data in:
+ nconmax_in: int,
+ # In:
+ box1: Geom,
+ box2: Geom,
+ worldid: int,
+ margin: float,
+ gap: float,
+ condim: int,
+ friction: vec5,
+ solref: wp.vec2f,
+ solreffriction: wp.vec2f,
+ solimp: vec5,
+ geoms: wp.vec2i,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ # Compute transforms between box's frames
+
+ pos21 = wp.transpose(box1.rot) @ (box2.pos - box1.pos)
+ pos12 = wp.transpose(box2.rot) @ (box1.pos - box2.pos)
+
+ rot21 = wp.transpose(box1.rot) @ box2.rot
+ rot12 = wp.transpose(rot21)
+
+ rot21abs = wp.matrix_from_rows(wp.abs(rot21[0]), wp.abs(rot21[1]), wp.abs(rot21[2]))
+ rot12abs = wp.transpose(rot21abs)
+
+ plen2 = rot21abs @ box2.size
+ plen1 = rot12abs @ box1.size
+
+ # Compute axis of maximum separation
+ s_sum_3 = 3.0 * (box1.size + box2.size)
+ separation = wp.float32(margin + s_sum_3[0] + s_sum_3[1] + s_sum_3[2])
+ axis_code = wp.int32(-1)
+
+ # First test: consider boxes' face normals
+ for i in range(3):
+ c1 = -wp.abs(pos21[i]) + box1.size[i] + plen2[i]
+
+ c2 = -wp.abs(pos12[i]) + box2.size[i] + plen1[i]
+
+ if c1 < -margin or c2 < -margin:
+ return
+
+ if c1 < separation:
+ separation = c1
+ axis_code = i + 3 * wp.int32(pos21[i] < 0) + 0 # Face of box1
+ if c2 < separation:
+ separation = c2
+ axis_code = i + 3 * wp.int32(pos12[i] < 0) + 6 # Face of box2
+
+ clnorm = wp.vec3(0.0)
+ inv = wp.bool(False)
+ cle1 = wp.int32(0)
+ cle2 = wp.int32(0)
+
+ # Second test: consider cross products of boxes' edges
+ for i in range(3):
+ for j in range(3):
+ # Compute cross product of box edges (potential separating axis)
+ if i == 0:
+ cross_axis = wp.vec3(0.0, -rot12[j, 2], rot12[j, 1])
+ elif i == 1:
+ cross_axis = wp.vec3(rot12[j, 2], 0.0, -rot12[j, 0])
+ else:
+ cross_axis = wp.vec3(-rot12[j, 1], rot12[j, 0], 0.0)
+
+ cross_length = wp.length(cross_axis)
+ if cross_length < MJ_MINVAL:
+ continue
+
+ cross_axis /= cross_length
+
+ box_dist = wp.dot(pos21, cross_axis)
+ c3 = wp.float32(0.0)
+
+ # Project box half-sizes onto the potential separating axis
+ for k in range(3):
+ if k != i:
+ c3 += box1.size[k] * wp.abs(cross_axis[k])
+ if k != j:
+ c3 += box2.size[k] * rot21abs[i, 3 - k - j] / cross_length
+
+ c3 -= wp.abs(box_dist)
+
+ # Early exit: no collision if separated along this axis
+ if c3 < -margin:
+ return
+
+ # Track minimum separation and which edge-edge pair it occurs on
+ if c3 < separation * (1.0 - 1e-12):
+ separation = c3
+ # Determine which corners/edges are closest
+ cle1 = 0
+ cle2 = 0
+
+ for k in range(3):
+ if k != i and (int(cross_axis[k] > 0) ^ int(box_dist < 0)):
+ cle1 += 1 << k
+ if k != j:
+ if int(rot21[i, 3 - k - j] > 0) ^ int(box_dist < 0) ^ int((k - j + 3) % 3 == 1):
+ cle2 += 1 << k
+
+ axis_code = 12 + i * 3 + j
+ clnorm = cross_axis
+ inv = box_dist < 0
+
+ # No axis with separation < margin found
+ if axis_code == -1:
+ return
+
+ points = mat83f()
+ depth = vec8f()
+ max_con_pair = 8
+ # 8 contacts should suffice for most configurations
+
+ if axis_code < 12:
+ # Handle face-vertex collision
+ face_idx = axis_code % 6
+ box_idx = axis_code / 6
+ rotmore = _compute_rotmore(face_idx)
+
+ r = rotmore @ wp.where(box_idx, rot12, rot21)
+ p = rotmore @ wp.where(box_idx, pos12, pos21)
+ ss = wp.abs(rotmore @ wp.where(box_idx, box2.size, box1.size))
+ s = wp.where(box_idx, box1.size, box2.size)
+ rt = wp.transpose(r)
+
+ lx, ly, hz = ss[0], ss[1], ss[2]
+ p[2] -= hz
+
+ clcorner = wp.int32(0) # corner of non-face box with least axis separation
+
+ for i in range(3):
+ if r[2, i] < 0:
+ clcorner += 1 << i
+
+ lp = p
+ for i in range(wp.static(3)):
+ lp += rt[i] * s[i] * wp.where(clcorner & 1 << i, 1.0, -1.0)
+
+ m = wp.int32(1)
+ dirs = wp.int32(0)
+
+ cn1 = wp.vec3(0.0)
+ cn2 = wp.vec3(0.0)
+
+ for i in range(3):
+ if wp.abs(r[2, i]) < 0.5:
+ if not dirs:
+ cn1 = rt[i] * s[i] * wp.where(clcorner & (1 << i), -2.0, 2.0)
+ else:
+ cn2 = rt[i] * s[i] * wp.where(clcorner & (1 << i), -2.0, 2.0)
+
+ dirs += 1
+
+ k = dirs * dirs
+
+ # Find potential contact points
+
+ n = wp.int32(0)
+
+ for i in range(k):
+ for q in range(2):
+ # lines_a and lines_b (lines between corners) computed on the fly
+ lav = lp + wp.where(i < 2, wp.vec3(0.0), wp.where(i == 2, cn1, cn2))
+ lbv = wp.where(i == 0 or i == 3, cn1, cn2)
+
+ if wp.abs(lbv[q]) > MJ_MINVAL:
+ br = 1.0 / lbv[q]
+ for j in range(-1, 2, 2):
+ l = ss[q] * wp.float32(j)
+ c1 = (l - lav[q]) * br
+ if c1 < 0 or c1 > 1:
+ continue
+ c2 = lav[1 - q] + lbv[1 - q] * c1
+ if wp.abs(c2) > ss[1 - q]:
+ continue
+
+ points[n] = lav + c1 * lbv
+ n += 1
+
+ if dirs == 2:
+ ax = cn1[0]
+ bx = cn2[0]
+ ay = cn1[1]
+ by = cn2[1]
+ C = 1.0 / (ax * by - bx * ay)
+
+ for i in range(4):
+ llx = wp.where(i / 2, lx, -lx)
+ lly = wp.where(i % 2, ly, -ly)
+
+ x = llx - lp[0]
+ y = lly - lp[1]
+
+ u = (x * by - y * bx) * C
+ v = (y * ax - x * ay) * C
+
+ if u > 0 and v > 0 and u < 1 and v < 1:
+ points[n] = wp.vec3(llx, lly, lp[2] + u * cn1[2] + v * cn2[2])
+ n += 1
+
+ for i in range(1 << dirs):
+ tmpv = lp + wp.float32(i & 1) * cn1 + wp.float32((i & 2) != 0) * cn2
+ if tmpv[0] > -lx and tmpv[0] < lx and tmpv[1] > -ly and tmpv[1] < ly:
+ points[n] = tmpv
+ n += 1
+
+ m = n
+ n = wp.int32(0)
+
+ for i in range(m):
+ if points[i][2] > margin:
+ continue
+ if i != n:
+ points[n] = points[i]
+
+ points[n, 2] *= 0.5
+ depth[n] = points[n, 2]
+ n += 1
+
+ # Set up contact frame
+ rw = wp.where(box_idx, box2.rot, box1.rot) @ wp.transpose(rotmore)
+ pw = wp.where(box_idx, box2.pos, box1.pos)
+ normal = wp.where(box_idx, -1.0, 1.0) * wp.transpose(rw)[2]
+
+ else:
+ # Handle edge-edge collision
+ edge1 = (axis_code - 12) / 3
+ edge2 = (axis_code - 12) % 3
+
+ # Set up non-contacting edges ax1, ax2 for box2 and pax1, pax2 for box 1
+ ax1 = wp.int(1 - (edge2 & 1))
+ ax2 = wp.int(2 - (edge2 & 2))
+
+ pax1 = wp.int(1 - (edge1 & 1))
+ pax2 = wp.int(2 - (edge1 & 2))
+
+ if rot21abs[edge1, ax1] < rot21abs[edge1, ax2]:
+ ax1, ax2 = ax2, ax1
+
+ if rot12abs[edge2, pax1] < rot12abs[edge2, pax2]:
+ pax1, pax2 = pax2, pax1
+
+ rotmore = _compute_rotmore(wp.where(cle1 & (1 << pax2), pax2, pax2 + 3))
+
+ # Transform coordinates for edge-edge contact calculation
+ p = rotmore @ pos21
+ rnorm = rotmore @ clnorm
+ r = rotmore @ rot21
+ rt = wp.transpose(r)
+ s = wp.abs(wp.transpose(rotmore) @ box1.size)
+
+ lx, ly, hz = s[0], s[1], s[2]
+ p[2] -= hz
+
+ # Calculate closest box2 face
+
+ points[0] = (
+ p
+ + rt[ax1] * box2.size[ax1] * wp.where(cle2 & (1 << ax1), 1.0, -1.0)
+ + rt[ax2] * box2.size[ax2] * wp.where(cle2 & (1 << ax2), 1.0, -1.0)
+ )
+ points[1] = points[0] - rt[edge2] * box2.size[edge2]
+ points[0] += rt[edge2] * box2.size[edge2]
+
+ points[2] = (
+ p
+ + rt[ax1] * box2.size[ax1] * wp.where(cle2 & (1 << ax1), -1.0, 1.0)
+ + rt[ax2] * box2.size[ax2] * wp.where(cle2 & (1 << ax2), 1.0, -1.0)
+ )
+
+ points[3] = points[2] - rt[edge2] * box2.size[edge2]
+ points[2] += rt[edge2] * box2.size[edge2]
+
+ n = 4
+
+ # Set up coordinate axes for contact face of box2
+ axi_lp = points[0]
+ axi_cn1 = points[1] - points[0]
+ axi_cn2 = points[2] - points[0]
+
+ # Check if contact normal is valid
+ if wp.abs(rnorm[2]) < MJ_MINVAL:
+ return # Shouldn't happen
+
+ # Calculate inverse normal for projection
+ innorm = wp.where(inv, -1.0, 1.0) / rnorm[2]
+
+ pu = mat43f()
+
+ # Project points onto contact plane
+ for i in range(4):
+ pu[i] = points[i]
+ c_scl = points[i, 2] * wp.where(inv, -1.0, 1.0) * innorm
+ points[i] -= rnorm * c_scl
+
+ pts_lp = points[0]
+ pts_cn1 = points[1] - points[0]
+ pts_cn2 = points[2] - points[0]
+
+ n = wp.int32(0)
+
+ for i in range(4):
+ for q in range(2):
+ la = pts_lp[q] + wp.where(i < 2, 0.0, wp.where(i == 2, pts_cn1[q], pts_cn2[q]))
+ lb = wp.where(i == 0 or i == 3, pts_cn1[q], pts_cn2[q])
+ lc = pts_lp[1 - q] + wp.where(i < 2, 0.0, wp.where(i == 2, pts_cn1[1 - q], pts_cn2[1 - q]))
+ ld = wp.where(i == 0 or i == 3, pts_cn1[1 - q], pts_cn2[1 - q])
+
+ # linesu_a and linesu_b (lines between corners) computed on the fly
+ lua = axi_lp + wp.where(i < 2, wp.vec3(0.0), wp.where(i == 2, axi_cn1, axi_cn2))
+ lub = wp.where(i == 0 or i == 3, axi_cn1, axi_cn2)
+
+ if wp.abs(lb) > MJ_MINVAL:
+ br = 1.0 / lb
+ for j in range(-1, 2, 2):
+ if n == max_con_pair:
+ break
+ l = s[q] * wp.float32(j)
+ c1 = (l - la) * br
+ if c1 < 0 or c1 > 1:
+ continue
+ c2 = lc + ld * c1
+ if wp.abs(c2) > s[1 - q]:
+ continue
+ if (lua[2] + lub[2] * c1) * innorm > margin:
+ continue
+
+ points[n] = lua * 0.5 + c1 * lub * 0.5
+ points[n, q] += 0.5 * l
+ points[n, 1 - q] += 0.5 * c2
+ depth[n] = points[n, 2] * innorm * 2.0
+ n += 1
+
+ nl = n
+
+ ax = pts_cn1[0]
+ bx = pts_cn2[0]
+ ay = pts_cn1[1]
+ by = pts_cn2[1]
+ C = 1.0 / (ax * by - bx * ay)
+
+ for i in range(4):
+ if n == max_con_pair:
+ break
+ llx = wp.where(i / 2, lx, -lx)
+ lly = wp.where(i % 2, ly, -ly)
+
+ x = llx - pts_lp[0]
+ y = lly - pts_lp[1]
+
+ u = (x * by - y * bx) * C
+ v = (y * ax - x * ay) * C
+
+ if nl == 0:
+ if (u < 0 or u > 0) and (v < 0 or v > 1):
+ continue
+ elif u < 0 or v < 0 or u > 1 or v > 1:
+ continue
+
+ u = wp.clamp(u, 0.0, 1.0)
+ v = wp.clamp(v, 0.0, 1.0)
+ w = 1.0 - u - v
+ vtmp = pu[0] * w + pu[1] * u + pu[2] * v
+
+ points[n] = wp.vec3(llx, lly, 0.0)
+
+ vtmp2 = points[n] - vtmp
+ tc1 = wp.length_sq(vtmp2)
+ if vtmp[2] > 0 and tc1 > margin * margin:
+ continue
+
+ points[n] = 0.5 * (points[n] + vtmp)
+
+ depth[n] = wp.sqrt(tc1) * wp.where(vtmp[2] < 0, -1.0, 1.0)
+ n += 1
+
+ nf = n
+
+ for i in range(4):
+ if n >= max_con_pair:
+ break
+ x = pu[i, 0]
+ y = pu[i, 1]
+ if nl == 0 and nf != 0:
+ if (x < -lx or x > lx) and (y < -ly or y > ly):
+ continue
+ elif x < -lx or x > lx or y < -ly or y > ly:
+ continue
+
+ c1 = wp.float32(0)
+
+ for j in range(2):
+ if pu[i, j] < -s[j]:
+ c1 += (pu[i, j] + s[j]) * (pu[i, j] + s[j])
+ elif pu[i, j] > s[j]:
+ c1 += (pu[i, j] - s[j]) * (pu[i, j] - s[j])
+
+ c1 += pu[i, 2] * innorm * pu[i, 2] * innorm
+
+ if pu[i, 2] > 0 and c1 > margin * margin:
+ continue
+
+ tmp_p = wp.vec3(pu[i, 0], pu[i, 1], 0.0)
+
+ for j in range(2):
+ if pu[i, j] < -s[j]:
+ tmp_p[j] = -s[j] * 0.5
+ elif pu[i, j] > s[j]:
+ tmp_p[j] = +s[j] * 0.5
+
+ tmp_p += pu[i]
+ points[n] = tmp_p * 0.5
+
+ depth[n] = wp.sqrt(c1) * wp.where(pu[i, 2] < 0, -1.0, 1.0)
+ n += 1
+
+ # Set up contact data for all points
+ rw = box1.rot @ wp.transpose(rotmore)
+ pw = box1.pos
+ normal = wp.where(inv, -1.0, 1.0) * rw @ rnorm
+
+ frame = make_frame(normal)
+ coff = wp.atomic_add(ncon_out, 0, n)
+
+ for i in range(min(nconmax_in - coff, n)):
+ points[i, 2] += hz
+ pos = rw @ points[i] + pw
+
+ cid = coff + i
+
+ contact_dist_out[cid] = depth[i]
+ contact_pos_out[cid] = pos
+ contact_frame_out[cid] = frame
+ contact_geom_out[cid] = geoms
+ contact_worldid_out[cid] = worldid
+ contact_includemargin_out[cid] = margin - gap
+ contact_dim_out[cid] = condim
+ contact_friction_out[cid] = friction
+ contact_solref_out[cid] = solref
+ contact_solreffriction_out[cid] = solreffriction
+ contact_solimp_out[cid] = solimp
+
+
+_PRIMITIVE_COLLISIONS = {
+ (GeomType.PLANE.value, GeomType.SPHERE.value): plane_sphere,
+ (GeomType.PLANE.value, GeomType.CAPSULE.value): plane_capsule,
+ (GeomType.PLANE.value, GeomType.ELLIPSOID.value): plane_ellipsoid,
+ (GeomType.PLANE.value, GeomType.CYLINDER.value): plane_cylinder,
+ (GeomType.PLANE.value, GeomType.BOX.value): plane_box,
+ (GeomType.PLANE.value, GeomType.MESH.value): plane_convex,
+ (GeomType.SPHERE.value, GeomType.SPHERE.value): sphere_sphere,
+ (GeomType.SPHERE.value, GeomType.CAPSULE.value): sphere_capsule,
+ (GeomType.SPHERE.value, GeomType.CYLINDER.value): sphere_cylinder,
+ (GeomType.SPHERE.value, GeomType.BOX.value): sphere_box,
+ (GeomType.CAPSULE.value, GeomType.CAPSULE.value): capsule_capsule,
+ (GeomType.CAPSULE.value, GeomType.BOX.value): capsule_box,
+ (GeomType.BOX.value, GeomType.BOX.value): box_box,
+}
+
+
+# TODO(team): _check_collisions shared utility
+def _check_primitive_collisions():
+ prev_idx = -1
+ for types in _PRIMITIVE_COLLISIONS.keys():
+ idx = upper_trid_index(len(GeomType), types[0], types[1])
+ if types[1] < types[0] or idx <= prev_idx:
+ return False
+ prev_idx = idx
+ return True
+
+
+assert _check_primitive_collisions(), "_PRIMITIVE_COLLISIONS is in invalid order"
+
+_primitive_collisions_types = []
+_primitive_collisions_func = []
+
+
+def _primitive_narrowphase_builder(m: Model):
+ for types, func in _PRIMITIVE_COLLISIONS.items():
+ idx = upper_trid_index(len(GeomType), types[0], types[1])
+ if m.geom_pair_type_count[idx] and types not in _primitive_collisions_types:
+ _primitive_collisions_types.append(types)
+ _primitive_collisions_func.append(func)
+
+ @wp.kernel
+ def _primitive_narrowphase(
+ # Model:
+ geom_type: wp.array(dtype=int),
+ geom_condim: wp.array(dtype=int),
+ geom_dataid: wp.array(dtype=int),
+ geom_priority: wp.array(dtype=int),
+ geom_solmix: wp.array2d(dtype=float),
+ geom_solref: wp.array2d(dtype=wp.vec2),
+ geom_solimp: wp.array2d(dtype=vec5),
+ geom_size: wp.array2d(dtype=wp.vec3),
+ geom_friction: wp.array2d(dtype=wp.vec3),
+ geom_margin: wp.array2d(dtype=float),
+ geom_gap: wp.array2d(dtype=float),
+ hfield_adr: wp.array(dtype=int),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ hfield_data: wp.array(dtype=float),
+ mesh_vertadr: wp.array(dtype=int),
+ mesh_vertnum: wp.array(dtype=int),
+ mesh_vert: wp.array(dtype=wp.vec3),
+ mesh_graphadr: wp.array(dtype=int),
+ mesh_graph: wp.array(dtype=int),
+ mesh_polynum: wp.array(dtype=int),
+ mesh_polyadr: wp.array(dtype=int),
+ mesh_polynormal: wp.array(dtype=wp.vec3),
+ mesh_polyvertadr: wp.array(dtype=int),
+ mesh_polyvertnum: wp.array(dtype=int),
+ mesh_polyvert: wp.array(dtype=int),
+ mesh_polymapadr: wp.array(dtype=int),
+ mesh_polymapnum: wp.array(dtype=int),
+ mesh_polymap: wp.array(dtype=int),
+ pair_dim: wp.array(dtype=int),
+ pair_solref: wp.array2d(dtype=wp.vec2),
+ pair_solreffriction: wp.array2d(dtype=wp.vec2),
+ pair_solimp: wp.array2d(dtype=vec5),
+ pair_margin: wp.array2d(dtype=float),
+ pair_gap: wp.array2d(dtype=float),
+ pair_friction: wp.array2d(dtype=vec5),
+ # Data in:
+ nconmax_in: int,
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ collision_pair_in: wp.array(dtype=wp.vec2i),
+ collision_hftri_index_in: wp.array(dtype=int),
+ collision_pairid_in: wp.array(dtype=int),
+ collision_worldid_in: wp.array(dtype=int),
+ ncollision_in: wp.array(dtype=int),
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+ ):
+ tid = wp.tid()
+
+ if tid >= ncollision_in[0]:
+ return
+
+ geoms = collision_pair_in[tid]
+ g1 = geoms[0]
+ g2 = geoms[1]
+
+ type1 = geom_type[g1]
+ type2 = geom_type[g2]
+
+ worldid = collision_worldid_in[tid]
+
+ _, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params(
+ geom_condim,
+ geom_priority,
+ geom_solmix,
+ geom_solref,
+ geom_solimp,
+ geom_friction,
+ geom_margin,
+ geom_gap,
+ pair_dim,
+ pair_solref,
+ pair_solreffriction,
+ pair_solimp,
+ pair_margin,
+ pair_gap,
+ pair_friction,
+ collision_pair_in,
+ collision_pairid_in,
+ tid,
+ worldid,
+ )
+
+ hftri_index = collision_hftri_index_in[tid]
+
+ geom1 = _geom(
+ geom_type,
+ geom_dataid,
+ geom_size,
+ hfield_adr,
+ hfield_nrow,
+ hfield_ncol,
+ hfield_size,
+ hfield_data,
+ mesh_vertadr,
+ mesh_vertnum,
+ mesh_vert,
+ mesh_graphadr,
+ mesh_graph,
+ mesh_polynum,
+ mesh_polyadr,
+ mesh_polynormal,
+ mesh_polyvertadr,
+ mesh_polyvertnum,
+ mesh_polyvert,
+ mesh_polymapadr,
+ mesh_polymapnum,
+ mesh_polymap,
+ geom_xpos_in,
+ geom_xmat_in,
+ worldid,
+ g1,
+ hftri_index,
+ )
+
+ geom2 = _geom(
+ geom_type,
+ geom_dataid,
+ geom_size,
+ hfield_adr,
+ hfield_nrow,
+ hfield_ncol,
+ hfield_size,
+ hfield_data,
+ mesh_vertadr,
+ mesh_vertnum,
+ mesh_vert,
+ mesh_graphadr,
+ mesh_graph,
+ mesh_polynum,
+ mesh_polyadr,
+ mesh_polynormal,
+ mesh_polyvertadr,
+ mesh_polyvertnum,
+ mesh_polyvert,
+ mesh_polymapadr,
+ mesh_polymapnum,
+ mesh_polymap,
+ geom_xpos_in,
+ geom_xmat_in,
+ worldid,
+ g2,
+ hftri_index,
+ )
+
+ for i in range(wp.static(len(_primitive_collisions_func))):
+ collision_type1 = wp.static(_primitive_collisions_types[i][0])
+ collision_type2 = wp.static(_primitive_collisions_types[i][1])
+
+ if collision_type1 == type1 and collision_type2 == type2:
+ wp.static(_primitive_collisions_func[i])(
+ nconmax_in,
+ geom1,
+ geom2,
+ worldid,
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+ return _primitive_narrowphase
+
+
+@event_scope
+def primitive_narrowphase(m: Model, d: Data):
+ """Runs collision detection on primitive geom pairs discovered during broadphase.
+
+ This function processes collision pairs involving primitive shapes that were
+ identified during the broadphase stage. It computes detailed contact information
+ such as distance, position, and frame, and populates the `d.contact` array.
+
+ The primitive geom types handled are PLANE, SPHERE, CAPSULE, CYLINDER, BOX.
+
+ It also handles collisions between planes and convex hulls.
+
+ To improve performance, it dynamically builds and launches a kernel tailored to
+ the specific primitive collision types present in the model, avoiding
+ unnecessary checks for non-existent collision pairs.
+ """
+ # we need to figure out how to keep the overhead of this small - not launching anything
+ # for pair types without collisions, as well as updating the launch dimensions.
+ wp.launch(
+ _primitive_narrowphase_builder(m),
+ dim=d.nconmax,
+ inputs=[
+ m.geom_type,
+ m.geom_condim,
+ m.geom_dataid,
+ m.geom_priority,
+ m.geom_solmix,
+ m.geom_solref,
+ m.geom_solimp,
+ m.geom_size,
+ m.geom_friction,
+ m.geom_margin,
+ m.geom_gap,
+ m.hfield_adr,
+ m.hfield_nrow,
+ m.hfield_ncol,
+ m.hfield_size,
+ m.hfield_data,
+ m.mesh_vertadr,
+ m.mesh_vertnum,
+ m.mesh_vert,
+ m.mesh_graphadr,
+ m.mesh_graph,
+ m.mesh_polynum,
+ m.mesh_polyadr,
+ m.mesh_polynormal,
+ m.mesh_polyvertadr,
+ m.mesh_polyvertnum,
+ m.mesh_polyvert,
+ m.mesh_polymapadr,
+ m.mesh_polymapnum,
+ m.mesh_polymap,
+ m.pair_dim,
+ m.pair_solref,
+ m.pair_solreffriction,
+ m.pair_solimp,
+ m.pair_margin,
+ m.pair_gap,
+ m.pair_friction,
+ d.nconmax,
+ d.geom_xpos,
+ d.geom_xmat,
+ d.collision_pair,
+ d.collision_hftri_index,
+ d.collision_pairid,
+ d.collision_worldid,
+ d.ncollision,
+ ],
+ outputs=[
+ d.ncon,
+ d.contact.dist,
+ d.contact.pos,
+ d.contact.frame,
+ d.contact.includemargin,
+ d.contact.friction,
+ d.contact.solref,
+ d.contact.solreffriction,
+ d.contact.solimp,
+ d.contact.dim,
+ d.contact.geom,
+ d.contact.worldid,
+ ],
+ )
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py
new file mode 100644
index 00000000..88457e91
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py
@@ -0,0 +1,599 @@
+# Copyright 2025 The Physics-Next Project Developers
+#
+# 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.
+# ==============================================================================
+
+from typing import Tuple
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src import math
+from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import _geom
+from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import contact_params
+from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import write_contact
+from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec5
+from mujoco.mjx.third_party.mujoco_warp._src.util_misc import halton
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
+
+
+@wp.struct
+class OptimizationParams:
+ rel_mat: wp.mat33
+ rel_pos: wp.vec3
+ attr1: wp.vec3
+ attr2: wp.vec3
+
+
+@wp.struct
+class AABB:
+ min: wp.vec3
+ max: wp.vec3
+
+
+@wp.func
+def transform_aabb(aabb_pos: wp.vec3, aabb_size: wp.vec3, pos: wp.vec3, ori: wp.mat33) -> AABB:
+ aabb = AABB()
+ aabb.max = wp.vec3(-1000000000.0, -1000000000.0, -1000000000.0)
+ aabb.min = wp.vec3(1000000000.0, 1000000000.0, 1000000000.0)
+
+ for i in range(8):
+ vec = wp.vec3(
+ aabb_size.x * (1.0 if (i & 1) else -1.0),
+ aabb_size.y * (1.0 if (i & 2) else -1.0),
+ aabb_size.z * (1.0 if (i & 4) else -1.0),
+ )
+
+ frame_vec = ori * (vec + aabb_pos) + pos
+
+ aabb.min = wp.min(aabb.min, frame_vec)
+ aabb.max = wp.max(aabb.max, frame_vec)
+
+ return aabb
+
+
+@wp.func
+def sphere(p: wp.vec3, size: wp.vec3) -> float:
+ return wp.length(p) - size[0]
+
+
+@wp.func
+def ellipsoid(p: wp.vec3, size: wp.vec3) -> float:
+ scaled_p = wp.vec3(p[0] / size[0], p[1] / size[1], p[2] / size[2])
+ k0 = wp.length(scaled_p)
+ k1 = wp.length(wp.vec3(p[0] / (size[0] ** 2.0), p[1] / (size[1] ** 2.0), p[2] / (size[2] ** 2.0)))
+ if k1 != 0.0:
+ denom = k1
+ else:
+ denom = 1e-12
+ return k0 * (k0 - 1.0) / denom
+
+
+@wp.func
+def grad_sphere(p: wp.vec3) -> wp.vec3:
+ c = wp.length(p)
+ if c > 1e-9:
+ return p / c
+ else:
+ wp.vec3(0.0)
+
+
+@wp.func
+def grad_ellipsoid(p: wp.vec3, size: wp.vec3) -> wp.vec3:
+ a = wp.vec3(p[0] / size[0], p[1] / size[1], p[2] / size[2])
+
+ b = wp.vec3(a[0] / size[0], a[1] / size[1], a[2] / size[2])
+ k0 = wp.length(a)
+ k1 = wp.length(b)
+ invK0 = 1.0 / k0
+ invK1 = 1.0 / k1
+
+ gk0 = b * invK0
+ gk1 = wp.vec3(
+ b[0] * invK1 / (size[0] * size[0]),
+ b[1] * invK1 / (size[1] * size[1]),
+ b[2] * invK1 / (size[2] * size[2]),
+ )
+ df_dk0 = (2.0 * k0 - 1.0) * invK1
+ df_dk1 = k0 * (k0 - 1.0) * invK1 * invK1
+
+ raw_grad = gk0 * df_dk0 - gk1 * df_dk1
+ return raw_grad / wp.length(raw_grad)
+
+
+@wp.func
+def user_sdf(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> float:
+ wp.printf("ERROR: user_sdf function must be implemented by user code\n")
+ return 0.0
+
+
+@wp.func
+def user_sdf_grad(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> wp.vec3:
+ wp.printf("ERROR: user_sdf_grad function must be implemented by user code\n")
+ return wp.vec3(0.0)
+
+
+@wp.func
+def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int) -> float:
+ if type == int(GeomType.SPHERE.value):
+ return sphere(p, attr)
+ elif type == int(GeomType.ELLIPSOID.value):
+ return ellipsoid(p, attr)
+ elif type == int(GeomType.SDF.value):
+ return user_sdf(p, attr, sdf_type)
+ wp.printf("ERROR: SDF type not implemented\n")
+ return 0.0
+
+
+@wp.func
+def sdf_grad(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int) -> wp.vec3:
+ if type == int(GeomType.SPHERE.value):
+ return grad_sphere(p)
+ elif type == int(GeomType.ELLIPSOID.value):
+ return grad_ellipsoid(p, attr)
+ elif type == int(GeomType.SDF.value):
+ return user_sdf_grad(p, attr, sdf_type)
+ wp.printf("ERROR: SDF grad type not implemented\n")
+ return wp.vec3(0.0)
+
+
+@wp.func
+def clearance(
+ type1: int, p1: wp.vec3, p2: wp.vec3, s1: wp.vec3, s2: wp.vec3, sdf_type1: int, sdf_type2: int, sfd_intersection: bool
+) -> float:
+ sdf1 = sdf(type1, p1, s1, sdf_type1)
+ sdf2 = sdf(int(GeomType.SDF.value), p2, s2, sdf_type2)
+ if sfd_intersection:
+ return wp.max(sdf1, sdf2)
+ else:
+ return sdf1 + sdf2 + wp.abs(wp.max(sdf1, sdf2))
+
+
+@wp.func
+def compute_grad(
+ type1: int, p1: wp.vec3, p2: wp.vec3, params: OptimizationParams, sdf_type1: int, sdf_type2: int, sfd_intersection: bool
+) -> wp.vec3:
+ A = sdf(type1, p1, params.attr1, sdf_type1)
+ B = sdf(int(GeomType.SDF.value), p2, params.attr2, sdf_type2)
+ grad1 = sdf_grad(type1, p1, params.attr1, sdf_type1)
+ grad2 = sdf_grad(int(GeomType.SDF.value), p2, params.attr2, sdf_type2)
+ grad1_transformed = params.rel_mat * grad1
+ if sfd_intersection:
+ if A > B:
+ return grad1_transformed
+ else:
+ return grad2
+ else:
+ gradient = grad2 + grad1_transformed
+ max_val = wp.max(A, B)
+ if A > B:
+ max_grad = grad1_transformed
+ else:
+ max_grad = grad2
+ sign = wp.sign(max_val)
+ gradient += max_grad * sign
+ return gradient
+
+
+@wp.func
+def gradient_step(
+ type1: int, x: wp.vec3, params: OptimizationParams, sdf_type1: int, sdf_type2: int, niter: int, sfd_intersection: bool
+) -> Tuple[float, wp.vec3]:
+ amin = 1e-4
+ rho = 0.5
+ c = 0.1
+ dist = float(1e10)
+
+ for _ in range(niter):
+ alpha = float(2.0)
+ x2 = wp.vec3(x[0], x[1], x[2])
+ x1 = params.rel_mat * x2 + params.rel_pos
+ grad = compute_grad(type1, x1, x2, params, sdf_type1, sdf_type2, sfd_intersection)
+ dist0 = clearance(type1, x1, x, params.attr1, params.attr2, sdf_type1, sdf_type2, sfd_intersection)
+ grad_dot = wp.dot(grad, grad)
+
+ if grad_dot < 1e-12:
+ return dist0, x
+
+ wolfe = -c * alpha * grad_dot
+ while True:
+ alpha *= rho
+ wolfe *= rho
+
+ x = x2 - grad * alpha
+ x1 = params.rel_mat * x + params.rel_pos
+ dist = clearance(type1, x1, x, params.attr1, params.attr2, sdf_type1, sdf_type2, sfd_intersection)
+
+ if alpha <= amin or (dist - dist0) <= wolfe:
+ break
+ if dist > dist0:
+ return dist, x
+ return dist, x
+
+
+@wp.func
+def gradient_descent(
+ # In:
+ type1: int,
+ x0_initial: wp.vec3,
+ attr1: wp.vec3,
+ attr2: wp.vec3,
+ pos1: wp.vec3,
+ rot1: wp.mat33,
+ pos2: wp.vec3,
+ rot2: wp.mat33,
+ sdf_type1: int,
+ sdf_type2: int,
+ sdf_iterations: int,
+) -> Tuple[float, wp.vec3, wp.vec3]:
+ params = OptimizationParams()
+ params.rel_mat = wp.transpose(rot1) * rot2
+ params.rel_pos = wp.transpose(rot1) * (pos2 - pos1)
+ params.attr1 = attr1
+ params.attr2 = attr2
+
+ # Collision phase (10 iterations, sfd_intersection=False)
+ dist, x = gradient_step(type1, x0_initial, params, sdf_type1, sdf_type2, sdf_iterations, False)
+
+ # Intersection phase (1 iteration, sfd_intersection=True)
+ dist, x = gradient_step(type1, x, params, sdf_type1, sdf_type2, 1, True)
+
+ # Midsurface calculation
+ x_1 = params.rel_mat * x + params.rel_pos
+
+ grad1 = sdf_grad(type1, x_1, params.attr1, sdf_type1)
+ grad1 = wp.transpose(params.rel_mat) * grad1
+ grad1 = wp.normalize(grad1)
+
+ grad2 = sdf_grad(int(GeomType.SDF.value), x, params.attr2, sdf_type2)
+ grad2 = wp.normalize(grad2)
+
+ n = grad1 - grad2
+ n = wp.normalize(n)
+ pos = rot2 * x + pos2
+ n = rot2 * n
+ pos3 = pos - n * dist / 2.0
+ return dist, pos3, n
+
+
+@wp.kernel
+def _sdf_narrowphase(
+ # Model:
+ geom_type: wp.array(dtype=int),
+ geom_condim: wp.array(dtype=int),
+ geom_dataid: wp.array(dtype=int),
+ geom_priority: wp.array(dtype=int),
+ geom_solmix: wp.array2d(dtype=float),
+ geom_solref: wp.array2d(dtype=wp.vec2),
+ geom_solimp: wp.array2d(dtype=vec5),
+ geom_size: wp.array2d(dtype=wp.vec3),
+ geom_aabb: wp.array2d(dtype=wp.vec3),
+ geom_pos: wp.array2d(dtype=wp.vec3),
+ geom_quat: wp.array2d(dtype=wp.quat),
+ geom_friction: wp.array2d(dtype=wp.vec3),
+ geom_margin: wp.array2d(dtype=float),
+ geom_gap: wp.array2d(dtype=float),
+ hfield_adr: wp.array(dtype=int),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ hfield_data: wp.array(dtype=float),
+ mesh_vertadr: wp.array(dtype=int),
+ mesh_vertnum: wp.array(dtype=int),
+ mesh_vert: wp.array(dtype=wp.vec3),
+ mesh_graphadr: wp.array(dtype=int),
+ mesh_graph: wp.array(dtype=int),
+ mesh_polynum: wp.array(dtype=int),
+ mesh_polyadr: wp.array(dtype=int),
+ mesh_polynormal: wp.array(dtype=wp.vec3),
+ mesh_polyvertadr: wp.array(dtype=int),
+ mesh_polyvertnum: wp.array(dtype=int),
+ mesh_polyvert: wp.array(dtype=int),
+ mesh_polymapadr: wp.array(dtype=int),
+ mesh_polymapnum: wp.array(dtype=int),
+ mesh_polymap: wp.array(dtype=int),
+ pair_dim: wp.array(dtype=int),
+ pair_solref: wp.array2d(dtype=wp.vec2),
+ pair_solreffriction: wp.array2d(dtype=wp.vec2),
+ pair_solimp: wp.array2d(dtype=vec5),
+ pair_margin: wp.array2d(dtype=float),
+ pair_gap: wp.array2d(dtype=float),
+ pair_friction: wp.array2d(dtype=vec5),
+ # In:
+ plugin: wp.array(dtype=int),
+ plugin_attr: wp.array(dtype=wp.vec3f),
+ geom_plugin_index: wp.array(dtype=int),
+ # Data in:
+ nconmax_in: int,
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ collision_pair_in: wp.array(dtype=wp.vec2i),
+ collision_hftri_index_in: wp.array(dtype=int),
+ collision_pairid_in: wp.array(dtype=int),
+ collision_worldid_in: wp.array(dtype=int),
+ ncollision_in: wp.array(dtype=int),
+ # In:
+ sdf_initpoints: int,
+ sdf_iterations: int,
+ # Data out:
+ ncon_out: wp.array(dtype=int),
+ contact_dist_out: wp.array(dtype=float),
+ contact_pos_out: wp.array(dtype=wp.vec3),
+ contact_frame_out: wp.array(dtype=wp.mat33),
+ contact_includemargin_out: wp.array(dtype=float),
+ contact_friction_out: wp.array(dtype=vec5),
+ contact_solref_out: wp.array(dtype=wp.vec2),
+ contact_solreffriction_out: wp.array(dtype=wp.vec2),
+ contact_solimp_out: wp.array(dtype=vec5),
+ contact_dim_out: wp.array(dtype=int),
+ contact_geom_out: wp.array(dtype=wp.vec2i),
+ contact_worldid_out: wp.array(dtype=int),
+):
+ tid = wp.tid()
+
+ if tid >= ncollision_in[0]:
+ return
+
+ geoms = collision_pair_in[tid]
+
+ g2 = geoms[1]
+ type2 = geom_type[g2]
+ if type2 != int(GeomType.SDF.value):
+ return
+
+ worldid = collision_worldid_in[tid]
+
+ _, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params(
+ geom_condim,
+ geom_priority,
+ geom_solmix,
+ geom_solref,
+ geom_solimp,
+ geom_friction,
+ geom_margin,
+ geom_gap,
+ pair_dim,
+ pair_solref,
+ pair_solreffriction,
+ pair_solimp,
+ pair_margin,
+ pair_gap,
+ pair_friction,
+ collision_pair_in,
+ collision_pairid_in,
+ tid,
+ worldid,
+ )
+ g1 = geoms[0]
+
+ hftri_index = collision_hftri_index_in[tid]
+
+ geom1 = _geom(
+ geom_type,
+ geom_dataid,
+ geom_size,
+ hfield_adr,
+ hfield_nrow,
+ hfield_ncol,
+ hfield_size,
+ hfield_data,
+ mesh_vertadr,
+ mesh_vertnum,
+ mesh_vert,
+ mesh_graphadr,
+ mesh_graph,
+ mesh_polynum,
+ mesh_polyadr,
+ mesh_polynormal,
+ mesh_polyvertadr,
+ mesh_polyvertnum,
+ mesh_polyvert,
+ mesh_polymapadr,
+ mesh_polymapnum,
+ mesh_polymap,
+ geom_xpos_in,
+ geom_xmat_in,
+ worldid,
+ g1,
+ hftri_index,
+ )
+ geom2 = _geom(
+ geom_type,
+ geom_dataid,
+ geom_size,
+ hfield_adr,
+ hfield_nrow,
+ hfield_ncol,
+ hfield_size,
+ hfield_data,
+ mesh_vertadr,
+ mesh_vertnum,
+ mesh_vert,
+ mesh_graphadr,
+ mesh_graph,
+ mesh_polynum,
+ mesh_polyadr,
+ mesh_polynormal,
+ mesh_polyvertadr,
+ mesh_polyvertnum,
+ mesh_polyvert,
+ mesh_polymapadr,
+ mesh_polymapnum,
+ mesh_polymap,
+ geom_xpos_in,
+ geom_xmat_in,
+ worldid,
+ g2,
+ hftri_index,
+ )
+
+ type1 = geom_type[g1]
+ g1_plugin = geom_plugin_index[g1]
+ g2_plugin = geom_plugin_index[g2]
+
+ g2_to_g1_rot = wp.transpose(geom2.rot) * geom1.rot
+ g2_to_g1_pos = wp.transpose(geom2.rot) * (geom1.pos - geom2.pos)
+
+ aabb_pos = geom_aabb[g1, 0]
+ aabb_size = geom_aabb[g1, 1]
+ aabb1 = transform_aabb(aabb_pos, aabb_size, g2_to_g1_pos, g2_to_g1_rot)
+
+ aabb_pos = geom_aabb[g2, 0]
+ aabb_size = geom_aabb[g2, 1]
+ aabb2 = transform_aabb(aabb_pos, aabb_size, wp.vec3(0.0), wp.mat33(1.0))
+
+ aabb_intersection = AABB()
+ aabb_intersection.min = wp.max(aabb1.min, aabb2.min)
+ aabb_intersection.max = wp.min(aabb1.max, aabb2.max)
+
+ geom_pos2 = geom_pos[worldid, g2]
+ quat2 = geom_quat[worldid, g2]
+ geom_mat2 = math.quat_to_mat(quat2)
+ rot2 = math.mul(geom2.rot, math.transpose(geom_mat2))
+ pos2 = wp.sub(geom2.pos, math.mul(rot2, geom_pos2))
+
+ if type1 == int(GeomType.SDF.value):
+ geom_pos1 = geom_pos[worldid, g1]
+ quat1 = geom_quat[worldid, g1]
+ geom_mat1 = math.quat_to_mat(quat1)
+ rot1 = math.mul(geom1.rot, math.transpose(geom_mat1))
+ pos1 = wp.sub(geom1.pos, math.mul(rot1, geom_pos1))
+ attr1 = plugin_attr[g1_plugin]
+ g1_plugin_id = plugin[g1_plugin]
+ else:
+ pos1 = geom1.pos
+ rot1 = geom1.rot
+ attr1 = geom1.size
+ g1_plugin_id = -1
+
+ for i in range(sdf_initpoints):
+ x_g2 = wp.vec3(
+ aabb_intersection.min[0] + (aabb_intersection.max[0] - aabb_intersection.min[0]) * halton(i, 2),
+ aabb_intersection.min[1] + (aabb_intersection.max[1] - aabb_intersection.min[1]) * halton(i, 3),
+ aabb_intersection.min[2] + (aabb_intersection.max[2] - aabb_intersection.min[2]) * halton(i, 5),
+ )
+
+ x = geom2.rot * x_g2 + geom2.pos
+ x0_initial = wp.transpose(rot2) * (x - pos2)
+
+ dist, pos, n = gradient_descent(
+ type1, x0_initial, attr1, plugin_attr[g2_plugin], pos1, rot1, pos2, rot2, g1_plugin_id, plugin[g2_plugin], sdf_iterations
+ )
+
+ write_contact(
+ nconmax_in,
+ dist,
+ pos,
+ make_frame(n),
+ margin,
+ gap,
+ condim,
+ friction,
+ solref,
+ solreffriction,
+ solimp,
+ geoms,
+ worldid,
+ ncon_out,
+ contact_dist_out,
+ contact_pos_out,
+ contact_frame_out,
+ contact_includemargin_out,
+ contact_friction_out,
+ contact_solref_out,
+ contact_solreffriction_out,
+ contact_solimp_out,
+ contact_dim_out,
+ contact_geom_out,
+ contact_worldid_out,
+ )
+
+
+@event_scope
+def sdf_narrowphase(m: Model, d: Data):
+ wp.launch(
+ _sdf_narrowphase,
+ dim=d.nconmax,
+ inputs=[
+ m.geom_type,
+ m.geom_condim,
+ m.geom_dataid,
+ m.geom_priority,
+ m.geom_solmix,
+ m.geom_solref,
+ m.geom_solimp,
+ m.geom_size,
+ m.geom_aabb,
+ m.geom_pos,
+ m.geom_quat,
+ m.geom_friction,
+ m.geom_margin,
+ m.geom_gap,
+ m.hfield_adr,
+ m.hfield_nrow,
+ m.hfield_ncol,
+ m.hfield_size,
+ m.hfield_data,
+ m.mesh_vertadr,
+ m.mesh_vertnum,
+ m.mesh_vert,
+ m.mesh_graphadr,
+ m.mesh_graph,
+ m.mesh_polynum,
+ m.mesh_polyadr,
+ m.mesh_polynormal,
+ m.mesh_polyvertadr,
+ m.mesh_polyvertnum,
+ m.mesh_polyvert,
+ m.mesh_polymapadr,
+ m.mesh_polymapnum,
+ m.mesh_polymap,
+ m.pair_dim,
+ m.pair_solref,
+ m.pair_solreffriction,
+ m.pair_solimp,
+ m.pair_margin,
+ m.pair_gap,
+ m.pair_friction,
+ m.plugin,
+ m.plugin_attr,
+ m.geom_plugin_index,
+ d.nconmax,
+ d.geom_xpos,
+ d.geom_xmat,
+ d.collision_pair,
+ d.collision_hftri_index,
+ d.collision_pairid,
+ d.collision_worldid,
+ d.ncollision,
+ m.opt.sdf_initpoints,
+ m.opt.sdf_iterations,
+ ],
+ outputs=[
+ d.ncon,
+ d.contact.dist,
+ d.contact.pos,
+ d.contact.frame,
+ d.contact.includemargin,
+ d.contact.friction,
+ d.contact.solref,
+ d.contact.solreffriction,
+ d.contact.solimp,
+ d.contact.dim,
+ d.contact.geom,
+ d.contact.worldid,
+ ],
+ )
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py
new file mode 100644
index 00000000..9ab71173
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py
@@ -0,0 +1,1908 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src import math
+from mujoco.mjx.third_party.mujoco_warp._src import support
+from mujoco.mjx.third_party.mujoco_warp._src import types
+from mujoco.mjx.third_party.mujoco_warp._src.types import ConstraintType
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec5
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec11
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
+
+wp.config.enable_backward = False
+
+
+@wp.kernel
+def zero_constraint_counts(
+ # Data out:
+ ne_out: wp.array(dtype=int),
+ ne_connect_out: wp.array(dtype=int),
+ ne_weld_out: wp.array(dtype=int),
+ ne_jnt_out: wp.array(dtype=int),
+ ne_ten_out: wp.array(dtype=int),
+ nf_out: wp.array(dtype=int),
+ nl_out: wp.array(dtype=int),
+ nefc_out: wp.array(dtype=int),
+):
+ worldid = wp.tid()
+
+ # Zero all constraint counters
+ ne_out[worldid] = 0
+ ne_connect_out[worldid] = 0
+ ne_weld_out[worldid] = 0
+ ne_jnt_out[worldid] = 0
+ ne_ten_out[worldid] = 0
+ nf_out[worldid] = 0
+ nl_out[worldid] = 0
+ nefc_out[worldid] = 0
+
+
+@wp.func
+def _update_efc_row(
+ # In:
+ worldid: int,
+ timestep: float,
+ refsafe: int,
+ efcid: int,
+ pos_aref: float,
+ pos_imp: float,
+ invweight: float,
+ solref: wp.vec2,
+ solimp: vec5,
+ margin: float,
+ vel: float,
+ frictionloss: float,
+ type: int,
+ id: int,
+ # Data out:
+ efc_type_out: wp.array2d(dtype=int),
+ efc_id_out: wp.array2d(dtype=int),
+ efc_pos_out: wp.array2d(dtype=float),
+ efc_margin_out: wp.array2d(dtype=float),
+ efc_D_out: wp.array2d(dtype=float),
+ efc_vel_out: wp.array2d(dtype=float),
+ efc_aref_out: wp.array2d(dtype=float),
+ efc_frictionloss_out: wp.array2d(dtype=float),
+):
+ # Calculate kbi
+ timeconst = solref[0]
+ dampratio = solref[1]
+ dmin = solimp[0]
+ dmax = solimp[1]
+ width = solimp[2]
+ mid = solimp[3]
+ power = solimp[4]
+
+ # TODO(team): wp.static?
+ if not refsafe:
+ timeconst = wp.max(timeconst, 2.0 * timestep)
+
+ dmin = wp.clamp(dmin, types.MJ_MINIMP, types.MJ_MAXIMP)
+ dmax = wp.clamp(dmax, types.MJ_MINIMP, types.MJ_MAXIMP)
+ width = wp.max(types.MJ_MINVAL, width)
+ mid = wp.clamp(mid, types.MJ_MINIMP, types.MJ_MAXIMP)
+ power = wp.max(1.0, power)
+
+ # See https://mujoco.readthedocs.io/en/latest/modeling.html#solver-parameters
+ k = 1.0 / (dmax * dmax * timeconst * timeconst * dampratio * dampratio)
+ b = 2.0 / (dmax * timeconst)
+ k = wp.where(solref[0] <= 0, -solref[0] / (dmax * dmax), k)
+ b = wp.where(solref[1] <= 0, -solref[1] / dmax, b)
+
+ imp_x = wp.abs(pos_imp) / width
+ imp_a = (1.0 / wp.pow(mid, power - 1.0)) * wp.pow(imp_x, power)
+ imp_b = 1.0 - (1.0 / wp.pow(1.0 - mid, power - 1.0)) * wp.pow(1.0 - imp_x, power)
+ imp_y = wp.where(imp_x < mid, imp_a, imp_b)
+ imp = dmin + imp_y * (dmax - dmin)
+ imp = wp.clamp(imp, dmin, dmax)
+ imp = wp.where(imp_x > 1.0, dmax, imp)
+
+ # Update constraints
+ efc_D_out[worldid, efcid] = 1.0 / wp.max(invweight * (1.0 - imp) / imp, types.MJ_MINVAL)
+ efc_vel_out[worldid, efcid] = vel
+ efc_aref_out[worldid, efcid] = -k * imp * pos_aref - b * vel
+ efc_pos_out[worldid, efcid] = pos_aref + margin
+ efc_margin_out[worldid, efcid] = margin
+ efc_frictionloss_out[worldid, efcid] = frictionloss
+ efc_type_out[worldid, efcid] = type
+ efc_id_out[worldid, efcid] = id
+
+
+@wp.kernel
+def _efc_equality_connect(
+ # Model:
+ nv: int,
+ nsite: int,
+ opt_timestep: wp.array(dtype=float),
+ body_parentid: wp.array(dtype=int),
+ body_rootid: wp.array(dtype=int),
+ body_invweight0: wp.array2d(dtype=wp.vec2),
+ dof_bodyid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ eq_obj1id: wp.array(dtype=int),
+ eq_obj2id: wp.array(dtype=int),
+ eq_objtype: wp.array(dtype=int),
+ eq_solref: wp.array2d(dtype=wp.vec2),
+ eq_solimp: wp.array2d(dtype=vec5),
+ eq_data: wp.array2d(dtype=vec11),
+ eq_connect_adr: wp.array(dtype=int),
+ # Data in:
+ njmax_in: int,
+ qvel_in: wp.array2d(dtype=float),
+ eq_active_in: wp.array2d(dtype=bool),
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xmat_in: wp.array2d(dtype=wp.mat33),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ refsafe_in: int,
+ # Data out:
+ ne_connect_out: wp.array(dtype=int),
+ nefc_out: wp.array(dtype=int),
+ efc_type_out: wp.array2d(dtype=int),
+ efc_id_out: wp.array2d(dtype=int),
+ efc_J_out: wp.array3d(dtype=float),
+ efc_pos_out: wp.array2d(dtype=float),
+ efc_margin_out: wp.array2d(dtype=float),
+ efc_D_out: wp.array2d(dtype=float),
+ efc_vel_out: wp.array2d(dtype=float),
+ efc_aref_out: wp.array2d(dtype=float),
+ efc_frictionloss_out: wp.array2d(dtype=float),
+):
+ """Calculates constraint rows for connect equality constraints."""
+
+ worldid, i_eq_connect_adr = wp.tid()
+ timestep = opt_timestep[worldid]
+ i_eq = eq_connect_adr[i_eq_connect_adr]
+
+ if not eq_active_in[worldid, i_eq]:
+ return
+
+ wp.atomic_add(ne_connect_out, worldid, 3)
+ efcid = wp.atomic_add(nefc_out, worldid, 3)
+
+ if efcid + 3 >= njmax_in:
+ return
+
+ data = eq_data[worldid, i_eq]
+ anchor1 = wp.vec3f(data[0], data[1], data[2])
+ anchor2 = wp.vec3f(data[3], data[4], data[5])
+
+ obj1id = eq_obj1id[i_eq]
+ obj2id = eq_obj2id[i_eq]
+
+ if nsite and eq_objtype[i_eq] == wp.static(types.ObjType.SITE.value):
+ # body1id stores the index of site_bodyid.
+ body1id = site_bodyid[obj1id]
+ body2id = site_bodyid[obj2id]
+ pos1 = site_xpos_in[worldid, obj1id]
+ pos2 = site_xpos_in[worldid, obj2id]
+ else:
+ body1id = obj1id
+ body2id = obj2id
+ pos1 = xpos_in[worldid, body1id] + xmat_in[worldid, body1id] @ anchor1
+ pos2 = xpos_in[worldid, body2id] + xmat_in[worldid, body2id] @ anchor2
+
+ # error is difference in global positions
+ pos = pos1 - pos2
+
+ # compute Jacobian difference (opposite of contact: 0 - 1)
+ Jqvel = wp.vec3f(0.0, 0.0, 0.0)
+ for dofid in range(nv): # TODO: parallelize
+ jacp1, _ = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ pos1,
+ body1id,
+ dofid,
+ worldid,
+ )
+ jacp2, _ = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ pos2,
+ body2id,
+ dofid,
+ worldid,
+ )
+ j1mj2 = jacp1 - jacp2
+ efc_J_out[worldid, efcid + 0, dofid] = j1mj2[0]
+ efc_J_out[worldid, efcid + 1, dofid] = j1mj2[1]
+ efc_J_out[worldid, efcid + 2, dofid] = j1mj2[2]
+ Jqvel += j1mj2 * qvel_in[worldid, dofid]
+
+ invweight = body_invweight0[worldid, body1id][0] + body_invweight0[worldid, body2id][0]
+ pos_imp = wp.length(pos)
+
+ solref = eq_solref[worldid, i_eq]
+ solimp = eq_solimp[worldid, i_eq]
+
+ for i in range(3):
+ efcidi = efcid + i
+
+ _update_efc_row(
+ worldid,
+ timestep,
+ refsafe_in,
+ efcidi,
+ pos[i],
+ pos_imp,
+ invweight,
+ solref,
+ solimp,
+ 0.0,
+ Jqvel[i],
+ 0.0,
+ ConstraintType.EQUALITY.value,
+ i_eq,
+ efc_type_out,
+ efc_id_out,
+ efc_pos_out,
+ efc_margin_out,
+ efc_D_out,
+ efc_vel_out,
+ efc_aref_out,
+ efc_frictionloss_out,
+ )
+
+
+@wp.kernel
+def _efc_equality_joint(
+ # Model:
+ nv: int,
+ opt_timestep: wp.array(dtype=float),
+ qpos0: wp.array2d(dtype=float),
+ jnt_qposadr: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ dof_invweight0: wp.array2d(dtype=float),
+ eq_obj1id: wp.array(dtype=int),
+ eq_obj2id: wp.array(dtype=int),
+ eq_solref: wp.array2d(dtype=wp.vec2),
+ eq_solimp: wp.array2d(dtype=vec5),
+ eq_data: wp.array2d(dtype=vec11),
+ eq_jnt_adr: wp.array(dtype=int),
+ # Data in:
+ njmax_in: int,
+ qpos_in: wp.array2d(dtype=float),
+ qvel_in: wp.array2d(dtype=float),
+ eq_active_in: wp.array2d(dtype=bool),
+ # In:
+ refsafe_in: int,
+ # Data out:
+ ne_jnt_out: wp.array(dtype=int),
+ nefc_out: wp.array(dtype=int),
+ efc_type_out: wp.array2d(dtype=int),
+ efc_id_out: wp.array2d(dtype=int),
+ efc_J_out: wp.array3d(dtype=float),
+ efc_pos_out: wp.array2d(dtype=float),
+ efc_margin_out: wp.array2d(dtype=float),
+ efc_D_out: wp.array2d(dtype=float),
+ efc_vel_out: wp.array2d(dtype=float),
+ efc_aref_out: wp.array2d(dtype=float),
+ efc_frictionloss_out: wp.array2d(dtype=float),
+):
+ worldid, i_eq_joint_adr = wp.tid()
+ timestep = opt_timestep[worldid]
+ i_eq = eq_jnt_adr[i_eq_joint_adr]
+ if not eq_active_in[worldid, i_eq]:
+ return
+
+ wp.atomic_add(ne_jnt_out, worldid, 1)
+ efcid = wp.atomic_add(nefc_out, worldid, 1)
+
+ if efcid >= njmax_in:
+ return
+
+ for i in range(nv):
+ efc_J_out[worldid, efcid, i] = 0.0
+
+ jntid_1 = eq_obj1id[i_eq]
+ jntid_2 = eq_obj2id[i_eq]
+ data = eq_data[worldid, i_eq]
+ dofadr1 = jnt_dofadr[jntid_1]
+ qposadr1 = jnt_qposadr[jntid_1]
+ efc_J_out[worldid, efcid, dofadr1] = 1.0
+
+ if jntid_2 > -1:
+ # Two joint constraint
+ qposadr2 = jnt_qposadr[jntid_2]
+ dofadr2 = jnt_dofadr[jntid_2]
+ dif = qpos_in[worldid, qposadr2] - qpos0[worldid, qposadr2]
+
+ # Horner's method for polynomials
+ rhs = data[0] + dif * (data[1] + dif * (data[2] + dif * (data[3] + dif * data[4])))
+ deriv_2 = data[1] + dif * (2.0 * data[2] + dif * (3.0 * data[3] + dif * 4.0 * data[4]))
+
+ pos = qpos_in[worldid, qposadr1] - qpos0[worldid, qposadr1] - rhs
+ Jqvel = qvel_in[worldid, dofadr1] - qvel_in[worldid, dofadr2] * deriv_2
+ invweight = dof_invweight0[worldid, dofadr1] + dof_invweight0[worldid, dofadr2]
+
+ efc_J_out[worldid, efcid, dofadr2] = -deriv_2
+ else:
+ # Single joint constraint
+ pos = qpos_in[worldid, qposadr1] - qpos0[worldid, qposadr1] - data[0]
+ Jqvel = qvel_in[worldid, dofadr1]
+ invweight = dof_invweight0[worldid, dofadr1]
+
+ # Update constraint parameters
+ _update_efc_row(
+ worldid,
+ timestep,
+ refsafe_in,
+ efcid,
+ pos,
+ pos,
+ invweight,
+ eq_solref[worldid, i_eq],
+ eq_solimp[worldid, i_eq],
+ 0.0,
+ Jqvel,
+ 0.0,
+ ConstraintType.EQUALITY.value,
+ i_eq,
+ efc_type_out,
+ efc_id_out,
+ efc_pos_out,
+ efc_margin_out,
+ efc_D_out,
+ efc_vel_out,
+ efc_aref_out,
+ efc_frictionloss_out,
+ )
+
+
+@wp.kernel
+def _efc_equality_tendon(
+ # Model:
+ nv: int,
+ opt_timestep: wp.array(dtype=float),
+ eq_obj1id: wp.array(dtype=int),
+ eq_obj2id: wp.array(dtype=int),
+ eq_solref: wp.array2d(dtype=wp.vec2),
+ eq_solimp: wp.array2d(dtype=vec5),
+ eq_data: wp.array2d(dtype=vec11),
+ eq_ten_adr: wp.array(dtype=int),
+ tendon_length0: wp.array2d(dtype=float),
+ tendon_invweight0: wp.array2d(dtype=float),
+ # Data in:
+ njmax_in: int,
+ qvel_in: wp.array2d(dtype=float),
+ eq_active_in: wp.array2d(dtype=bool),
+ ten_length_in: wp.array2d(dtype=float),
+ ten_J_in: wp.array3d(dtype=float),
+ # In:
+ refsafe_in: int,
+ # Data out:
+ ne_ten_out: wp.array(dtype=int),
+ nefc_out: wp.array(dtype=int),
+ efc_type_out: wp.array2d(dtype=int),
+ efc_id_out: wp.array2d(dtype=int),
+ efc_J_out: wp.array3d(dtype=float),
+ efc_pos_out: wp.array2d(dtype=float),
+ efc_margin_out: wp.array2d(dtype=float),
+ efc_D_out: wp.array2d(dtype=float),
+ efc_vel_out: wp.array2d(dtype=float),
+ efc_aref_out: wp.array2d(dtype=float),
+ efc_frictionloss_out: wp.array2d(dtype=float),
+):
+ worldid, tenid = wp.tid()
+ timestep = opt_timestep[worldid]
+ eqid = eq_ten_adr[tenid]
+
+ if not eq_active_in[worldid, eqid]:
+ return
+
+ wp.atomic_add(ne_ten_out, worldid, 1)
+ efcid = wp.atomic_add(nefc_out, worldid, 1)
+
+ if efcid >= njmax_in:
+ return
+
+ obj1id = eq_obj1id[eqid]
+ obj2id = eq_obj2id[eqid]
+
+ data = eq_data[worldid, eqid]
+ solref = eq_solref[worldid, eqid]
+ solimp = eq_solimp[worldid, eqid]
+ pos1 = ten_length_in[worldid, obj1id] - tendon_length0[worldid, obj1id]
+ jac1 = ten_J_in[worldid, obj1id]
+
+ if obj2id > -1:
+ invweight = tendon_invweight0[worldid, obj1id] + tendon_invweight0[worldid, obj2id]
+
+ pos2 = ten_length_in[worldid, obj2id] - tendon_length0[worldid, obj2id]
+ jac2 = ten_J_in[worldid, obj2id]
+
+ dif = pos2
+ dif2 = dif * dif
+ dif3 = dif2 * dif
+ dif4 = dif3 * dif
+
+ pos = pos1 - (data[0] + data[1] * dif + data[2] * dif2 + data[3] * dif3 + data[4] * dif4)
+ deriv = data[1] + 2.0 * data[2] * dif + 3.0 * data[3] * dif2 + 4.0 * data[4] * dif3
+ else:
+ invweight = tendon_invweight0[worldid, obj1id]
+ pos = pos1 - data[0]
+ deriv = 0.0
+
+ Jqvel = float(0.0)
+ for i in range(nv):
+ if deriv != 0.0:
+ J = jac1[i] + jac2[i] * -deriv
+ else:
+ J = jac1[i]
+ efc_J_out[worldid, efcid, i] = J
+ Jqvel += J * qvel_in[worldid, i]
+
+ _update_efc_row(
+ worldid,
+ timestep,
+ refsafe_in,
+ efcid,
+ pos,
+ pos,
+ invweight,
+ solref,
+ solimp,
+ 0.0,
+ Jqvel,
+ 0.0,
+ ConstraintType.EQUALITY.value,
+ eqid,
+ efc_type_out,
+ efc_id_out,
+ efc_pos_out,
+ efc_margin_out,
+ efc_D_out,
+ efc_vel_out,
+ efc_aref_out,
+ efc_frictionloss_out,
+ )
+
+
+@wp.kernel
+def _efc_friction_dof(
+ # Model:
+ nv: int,
+ opt_timestep: wp.array(dtype=float),
+ dof_invweight0: wp.array2d(dtype=float),
+ dof_frictionloss: wp.array2d(dtype=float),
+ dof_solimp: wp.array2d(dtype=vec5),
+ dof_solref: wp.array2d(dtype=wp.vec2),
+ # Data in:
+ njmax_in: int,
+ qvel_in: wp.array2d(dtype=float),
+ # In:
+ refsafe_in: int,
+ # Data out:
+ nf_out: wp.array(dtype=int),
+ nefc_out: wp.array(dtype=int),
+ efc_type_out: wp.array2d(dtype=int),
+ efc_id_out: wp.array2d(dtype=int),
+ efc_J_out: wp.array3d(dtype=float),
+ efc_pos_out: wp.array2d(dtype=float),
+ efc_margin_out: wp.array2d(dtype=float),
+ efc_D_out: wp.array2d(dtype=float),
+ efc_vel_out: wp.array2d(dtype=float),
+ efc_aref_out: wp.array2d(dtype=float),
+ efc_frictionloss_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ timestep = opt_timestep[worldid]
+
+ if dof_frictionloss[worldid, dofid] <= 0.0:
+ return
+
+ efcid = wp.atomic_add(nefc_out, worldid, 1)
+
+ if efcid >= njmax_in:
+ return
+
+ wp.atomic_add(nf_out, worldid, 1)
+
+ for i in range(nv):
+ efc_J_out[worldid, efcid, i] = 0.0
+
+ efc_J_out[worldid, efcid, dofid] = 1.0
+ Jqvel = qvel_in[worldid, dofid]
+
+ _update_efc_row(
+ worldid,
+ timestep,
+ refsafe_in,
+ efcid,
+ 0.0,
+ 0.0,
+ dof_invweight0[worldid, dofid],
+ dof_solref[worldid, dofid],
+ dof_solimp[worldid, dofid],
+ 0.0,
+ Jqvel,
+ dof_frictionloss[worldid, dofid],
+ ConstraintType.FRICTION_DOF.value,
+ dofid,
+ efc_type_out,
+ efc_id_out,
+ efc_pos_out,
+ efc_margin_out,
+ efc_D_out,
+ efc_vel_out,
+ efc_aref_out,
+ efc_frictionloss_out,
+ )
+
+
+@wp.kernel
+def _efc_friction_tendon(
+ # Model:
+ nv: int,
+ opt_timestep: wp.array(dtype=float),
+ tendon_solref_fri: wp.array2d(dtype=wp.vec2),
+ tendon_solimp_fri: wp.array2d(dtype=vec5),
+ tendon_frictionloss: wp.array2d(dtype=float),
+ tendon_invweight0: wp.array2d(dtype=float),
+ # Data in:
+ qvel_in: wp.array2d(dtype=float),
+ ten_J_in: wp.array3d(dtype=float),
+ # In:
+ refsafe_in: int,
+ # Data out:
+ nf_out: wp.array(dtype=int),
+ nefc_out: wp.array(dtype=int),
+ efc_type_out: wp.array2d(dtype=int),
+ efc_id_out: wp.array2d(dtype=int),
+ efc_J_out: wp.array3d(dtype=float),
+ efc_pos_out: wp.array2d(dtype=float),
+ efc_margin_out: wp.array2d(dtype=float),
+ efc_D_out: wp.array2d(dtype=float),
+ efc_vel_out: wp.array2d(dtype=float),
+ efc_aref_out: wp.array2d(dtype=float),
+ efc_frictionloss_out: wp.array2d(dtype=float),
+):
+ worldid, tenid = wp.tid()
+ timestep = opt_timestep[worldid]
+
+ frictionloss = tendon_frictionloss[worldid, tenid]
+ if frictionloss <= 0.0:
+ return
+
+ efcid = wp.atomic_add(nefc_out, worldid, 1)
+ wp.atomic_add(nf_out, worldid, 1)
+
+ Jqvel = float(0.0)
+
+ # TODO(team): parallelize
+ for i in range(nv):
+ J = ten_J_in[worldid, tenid, i]
+ efc_J_out[worldid, efcid, i] = J
+ Jqvel += J * qvel_in[worldid, i]
+
+ _update_efc_row(
+ worldid,
+ timestep,
+ refsafe_in,
+ efcid,
+ 0.0,
+ 0.0,
+ tendon_invweight0[worldid, tenid],
+ tendon_solref_fri[worldid, tenid],
+ tendon_solimp_fri[worldid, tenid],
+ 0.0,
+ Jqvel,
+ frictionloss,
+ ConstraintType.FRICTION_TENDON.value,
+ tenid,
+ efc_type_out,
+ efc_id_out,
+ efc_pos_out,
+ efc_margin_out,
+ efc_D_out,
+ efc_vel_out,
+ efc_aref_out,
+ efc_frictionloss_out,
+ )
+
+
+@wp.kernel
+def _efc_equality_weld(
+ # Model:
+ nv: int,
+ nsite: int,
+ opt_timestep: wp.array(dtype=float),
+ body_parentid: wp.array(dtype=int),
+ body_rootid: wp.array(dtype=int),
+ body_invweight0: wp.array2d(dtype=wp.vec2),
+ dof_bodyid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ site_quat: wp.array2d(dtype=wp.quat),
+ eq_obj1id: wp.array(dtype=int),
+ eq_obj2id: wp.array(dtype=int),
+ eq_objtype: wp.array(dtype=int),
+ eq_solref: wp.array2d(dtype=wp.vec2),
+ eq_solimp: wp.array2d(dtype=vec5),
+ eq_data: wp.array2d(dtype=vec11),
+ eq_wld_adr: wp.array(dtype=int),
+ # Data in:
+ njmax_in: int,
+ qvel_in: wp.array2d(dtype=float),
+ eq_active_in: wp.array2d(dtype=bool),
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xquat_in: wp.array2d(dtype=wp.quat),
+ xmat_in: wp.array2d(dtype=wp.mat33),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ refsafe_in: int,
+ # Data out:
+ ne_weld_out: wp.array(dtype=int),
+ nefc_out: wp.array(dtype=int),
+ efc_type_out: wp.array2d(dtype=int),
+ efc_id_out: wp.array2d(dtype=int),
+ efc_J_out: wp.array3d(dtype=float),
+ efc_pos_out: wp.array2d(dtype=float),
+ efc_margin_out: wp.array2d(dtype=float),
+ efc_D_out: wp.array2d(dtype=float),
+ efc_vel_out: wp.array2d(dtype=float),
+ efc_aref_out: wp.array2d(dtype=float),
+ efc_frictionloss_out: wp.array2d(dtype=float),
+):
+ worldid, i_eq_weld_adr = wp.tid()
+ timestep = opt_timestep[worldid]
+ i_eq = eq_wld_adr[i_eq_weld_adr]
+ if not eq_active_in[worldid, i_eq]:
+ return
+
+ wp.atomic_add(ne_weld_out, worldid, 6)
+ efcid = wp.atomic_add(nefc_out, worldid, 6)
+
+ if efcid + 6 >= njmax_in:
+ return
+
+ is_site = eq_objtype[i_eq] == wp.static(types.ObjType.SITE.value) and nsite > 0
+
+ obj1id = eq_obj1id[i_eq]
+ obj2id = eq_obj2id[i_eq]
+
+ data = eq_data[worldid, i_eq]
+ anchor1 = wp.vec3(data[0], data[1], data[2])
+ anchor2 = wp.vec3(data[3], data[4], data[5])
+ relpose = wp.quat(data[6], data[7], data[8], data[9])
+ torquescale = data[10]
+
+ if is_site:
+ # body1id stores the index of site_bodyid.
+ body1id = site_bodyid[obj1id]
+ body2id = site_bodyid[obj2id]
+ pos1 = site_xpos_in[worldid, obj1id]
+ pos2 = site_xpos_in[worldid, obj2id]
+
+ quat = math.mul_quat(xquat_in[worldid, body1id], site_quat[worldid, obj1id])
+ quat1 = math.quat_inv(math.mul_quat(xquat_in[worldid, body2id], site_quat[worldid, obj2id]))
+
+ else:
+ body1id = obj1id
+ body2id = obj2id
+ pos1 = xpos_in[worldid, body1id] + xmat_in[worldid, body1id] @ anchor2
+ pos2 = xpos_in[worldid, body2id] + xmat_in[worldid, body2id] @ anchor1
+
+ quat = math.mul_quat(xquat_in[worldid, body1id], relpose)
+ quat1 = math.quat_inv(xquat_in[worldid, body2id])
+
+ # compute Jacobian difference (opposite of contact: 0 - 1)
+ Jqvelp = wp.vec3f(0.0, 0.0, 0.0)
+ Jqvelr = wp.vec3f(0.0, 0.0, 0.0)
+
+ for dofid in range(nv): # TODO: parallelize
+ jacp1, jacr1 = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ pos1,
+ body1id,
+ dofid,
+ worldid,
+ )
+ jacp2, jacr2 = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ pos2,
+ body2id,
+ dofid,
+ worldid,
+ )
+
+ jacdifp = jacp1 - jacp2
+ for i in range(wp.static(3)):
+ efc_J_out[worldid, efcid + i, dofid] = jacdifp[i]
+
+ jacdifr = (jacr1 - jacr2) * torquescale
+ jacdifrq = math.mul_quat(math.quat_mul_axis(quat1, jacdifr), quat)
+ jacdifr = 0.5 * wp.vec3(jacdifrq[1], jacdifrq[2], jacdifrq[3])
+
+ for i in range(wp.static(3)):
+ efc_J_out[worldid, efcid + 3 + i, dofid] = jacdifr[i]
+
+ Jqvelp += jacdifp * qvel_in[worldid, dofid]
+ Jqvelr += jacdifr * qvel_in[worldid, dofid]
+
+ # error is difference in global position and orientation
+ cpos = pos1 - pos2
+
+ crotq = math.mul_quat(quat1, quat) # copy axis components
+ crot = wp.vec3(crotq[1], crotq[2], crotq[3]) * torquescale
+
+ invweight_t = body_invweight0[worldid, body1id][0] + body_invweight0[worldid, body2id][0]
+
+ pos_imp = wp.sqrt(wp.length_sq(cpos) + wp.length_sq(crot))
+
+ solref = eq_solref[worldid, i_eq]
+ solimp = eq_solimp[worldid, i_eq]
+
+ for i in range(3):
+ _update_efc_row(
+ worldid,
+ timestep,
+ refsafe_in,
+ efcid + i,
+ cpos[i],
+ pos_imp,
+ invweight_t,
+ solref,
+ solimp,
+ 0.0,
+ Jqvelp[i],
+ 0.0,
+ ConstraintType.EQUALITY.value,
+ i_eq,
+ efc_type_out,
+ efc_id_out,
+ efc_pos_out,
+ efc_margin_out,
+ efc_D_out,
+ efc_vel_out,
+ efc_aref_out,
+ efc_frictionloss_out,
+ )
+
+ invweight_r = body_invweight0[worldid, body1id][1] + body_invweight0[worldid, body2id][1]
+
+ for i in range(3):
+ _update_efc_row(
+ worldid,
+ timestep,
+ refsafe_in,
+ efcid + 3 + i,
+ crot[i],
+ pos_imp,
+ invweight_r,
+ solref,
+ solimp,
+ 0.0,
+ Jqvelr[i],
+ 0.0,
+ ConstraintType.EQUALITY.value,
+ i_eq,
+ efc_type_out,
+ efc_id_out,
+ efc_pos_out,
+ efc_margin_out,
+ efc_D_out,
+ efc_vel_out,
+ efc_aref_out,
+ efc_frictionloss_out,
+ )
+
+
+@wp.kernel
+def _efc_limit_slide_hinge(
+ # Model:
+ nv: int,
+ opt_timestep: wp.array(dtype=float),
+ jnt_qposadr: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ jnt_solref: wp.array2d(dtype=wp.vec2),
+ jnt_solimp: wp.array2d(dtype=vec5),
+ jnt_range: wp.array2d(dtype=wp.vec2),
+ jnt_margin: wp.array2d(dtype=float),
+ jnt_limited_slide_hinge_adr: wp.array(dtype=int),
+ dof_invweight0: wp.array2d(dtype=float),
+ # Data in:
+ njmax_in: int,
+ qpos_in: wp.array2d(dtype=float),
+ qvel_in: wp.array2d(dtype=float),
+ # In:
+ refsafe_in: int,
+ # Data out:
+ nl_out: wp.array(dtype=int),
+ nefc_out: wp.array(dtype=int),
+ efc_type_out: wp.array2d(dtype=int),
+ efc_id_out: wp.array2d(dtype=int),
+ efc_J_out: wp.array3d(dtype=float),
+ efc_pos_out: wp.array2d(dtype=float),
+ efc_margin_out: wp.array2d(dtype=float),
+ efc_D_out: wp.array2d(dtype=float),
+ efc_vel_out: wp.array2d(dtype=float),
+ efc_aref_out: wp.array2d(dtype=float),
+ efc_frictionloss_out: wp.array2d(dtype=float),
+):
+ worldid, jntlimitedid = wp.tid()
+ timestep = opt_timestep[worldid]
+ jntid = jnt_limited_slide_hinge_adr[jntlimitedid]
+ jntrange = jnt_range[worldid, jntid]
+
+ qpos = qpos_in[worldid, jnt_qposadr[jntid]]
+ jntmargin = jnt_margin[worldid, jntid]
+ dist_min, dist_max = qpos - jntrange[0], jntrange[1] - qpos
+ pos = wp.min(dist_min, dist_max) - jntmargin
+ active = pos < 0
+
+ if active:
+ wp.atomic_add(nl_out, worldid, 1)
+ efcid = wp.atomic_add(nefc_out, worldid, 1)
+
+ if efcid >= njmax_in:
+ return
+
+ for i in range(nv):
+ efc_J_out[worldid, efcid, i] = 0.0
+
+ dofadr = jnt_dofadr[jntid]
+
+ J = float(dist_min < dist_max) * 2.0 - 1.0
+ efc_J_out[worldid, efcid, dofadr] = J
+ Jqvel = J * qvel_in[worldid, dofadr]
+
+ _update_efc_row(
+ worldid,
+ timestep,
+ refsafe_in,
+ efcid,
+ pos,
+ pos,
+ dof_invweight0[worldid, dofadr],
+ jnt_solref[worldid, jntid],
+ jnt_solimp[worldid, jntid],
+ jntmargin,
+ Jqvel,
+ 0.0,
+ ConstraintType.LIMIT_JOINT.value,
+ jntid,
+ efc_type_out,
+ efc_id_out,
+ efc_pos_out,
+ efc_margin_out,
+ efc_D_out,
+ efc_vel_out,
+ efc_aref_out,
+ efc_frictionloss_out,
+ )
+
+
+@wp.kernel
+def _efc_limit_ball(
+ # Model:
+ nv: int,
+ opt_timestep: wp.array(dtype=float),
+ jnt_qposadr: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ jnt_solref: wp.array2d(dtype=wp.vec2),
+ jnt_solimp: wp.array2d(dtype=vec5),
+ jnt_range: wp.array2d(dtype=wp.vec2),
+ jnt_margin: wp.array2d(dtype=float),
+ jnt_limited_ball_adr: wp.array(dtype=int),
+ dof_invweight0: wp.array2d(dtype=float),
+ # Data in:
+ njmax_in: int,
+ qpos_in: wp.array2d(dtype=float),
+ qvel_in: wp.array2d(dtype=float),
+ # In:
+ refsafe_in: int,
+ # Data out:
+ nl_out: wp.array(dtype=int),
+ nefc_out: wp.array(dtype=int),
+ efc_type_out: wp.array2d(dtype=int),
+ efc_id_out: wp.array2d(dtype=int),
+ efc_J_out: wp.array3d(dtype=float),
+ efc_pos_out: wp.array2d(dtype=float),
+ efc_margin_out: wp.array2d(dtype=float),
+ efc_D_out: wp.array2d(dtype=float),
+ efc_vel_out: wp.array2d(dtype=float),
+ efc_aref_out: wp.array2d(dtype=float),
+ efc_frictionloss_out: wp.array2d(dtype=float),
+):
+ worldid, jntlimitedid = wp.tid()
+ timestep = opt_timestep[worldid]
+ jntid = jnt_limited_ball_adr[jntlimitedid]
+ qposadr = jnt_qposadr[jntid]
+
+ qpos = qpos_in[worldid]
+ jnt_quat = wp.quat(qpos[qposadr + 0], qpos[qposadr + 1], qpos[qposadr + 2], qpos[qposadr + 3])
+ jnt_quat = wp.normalize(jnt_quat)
+ axis_angle = math.quat_to_vel(jnt_quat)
+ jntrange = jnt_range[worldid, jntid]
+ axis, angle = math.normalize_with_norm(axis_angle)
+ jntmargin = jnt_margin[worldid, jntid]
+
+ pos = wp.max(jntrange[0], jntrange[1]) - angle - jntmargin
+ active = pos < 0
+
+ if active:
+ wp.atomic_add(nl_out, worldid, 1)
+ efcid = wp.atomic_add(nefc_out, worldid, 1)
+
+ if efcid >= njmax_in:
+ return
+
+ for i in range(nv):
+ efc_J_out[worldid, efcid, i] = 0.0
+
+ dofadr = jnt_dofadr[jntid]
+
+ efc_J_out[worldid, efcid, dofadr + 0] = -axis[0]
+ efc_J_out[worldid, efcid, dofadr + 1] = -axis[1]
+ efc_J_out[worldid, efcid, dofadr + 2] = -axis[2]
+
+ Jqvel = -axis[0] * qvel_in[worldid, dofadr + 0]
+ Jqvel -= axis[1] * qvel_in[worldid, dofadr + 1]
+ Jqvel -= axis[2] * qvel_in[worldid, dofadr + 2]
+
+ _update_efc_row(
+ worldid,
+ timestep,
+ refsafe_in,
+ efcid,
+ pos,
+ pos,
+ dof_invweight0[worldid, dofadr],
+ jnt_solref[worldid, jntid],
+ jnt_solimp[worldid, jntid],
+ jntmargin,
+ Jqvel,
+ 0.0,
+ ConstraintType.LIMIT_JOINT.value,
+ jntid,
+ efc_type_out,
+ efc_id_out,
+ efc_pos_out,
+ efc_margin_out,
+ efc_D_out,
+ efc_vel_out,
+ efc_aref_out,
+ efc_frictionloss_out,
+ )
+
+
+@wp.kernel
+def _efc_limit_tendon(
+ # Model:
+ nv: int,
+ opt_timestep: wp.array(dtype=float),
+ jnt_dofadr: wp.array(dtype=int),
+ tendon_adr: wp.array(dtype=int),
+ tendon_num: wp.array(dtype=int),
+ tendon_limited_adr: wp.array(dtype=int),
+ tendon_solref_lim: wp.array2d(dtype=wp.vec2),
+ tendon_solimp_lim: wp.array2d(dtype=vec5),
+ tendon_range: wp.array2d(dtype=wp.vec2),
+ tendon_margin: wp.array2d(dtype=float),
+ tendon_invweight0: wp.array2d(dtype=float),
+ wrap_objid: wp.array(dtype=int),
+ wrap_type: wp.array(dtype=int),
+ # Data in:
+ njmax_in: int,
+ qvel_in: wp.array2d(dtype=float),
+ ten_length_in: wp.array2d(dtype=float),
+ ten_J_in: wp.array3d(dtype=float),
+ # In:
+ refsafe_in: int,
+ # Data out:
+ nl_out: wp.array(dtype=int),
+ nefc_out: wp.array(dtype=int),
+ efc_type_out: wp.array2d(dtype=int),
+ efc_id_out: wp.array2d(dtype=int),
+ efc_J_out: wp.array3d(dtype=float),
+ efc_pos_out: wp.array2d(dtype=float),
+ efc_margin_out: wp.array2d(dtype=float),
+ efc_D_out: wp.array2d(dtype=float),
+ efc_vel_out: wp.array2d(dtype=float),
+ efc_aref_out: wp.array2d(dtype=float),
+ efc_frictionloss_out: wp.array2d(dtype=float),
+):
+ worldid, tenlimitedid = wp.tid()
+ timestep = opt_timestep[worldid]
+ tenid = tendon_limited_adr[tenlimitedid]
+
+ tenrange = tendon_range[worldid, tenid]
+ length = ten_length_in[worldid, tenid]
+ dist_min, dist_max = length - tenrange[0], tenrange[1] - length
+ tenmargin = tendon_margin[worldid, tenid]
+ pos = wp.min(dist_min, dist_max) - tenmargin
+ active = pos < 0
+
+ if active:
+ wp.atomic_add(nl_out, worldid, 1)
+ efcid = wp.atomic_add(nefc_out, worldid, 1)
+
+ if efcid >= njmax_in:
+ return
+
+ for i in range(nv):
+ efc_J_out[worldid, efcid, i] = 0.0
+
+ Jqvel = float(0.0)
+ scl = float(dist_min < dist_max) * 2.0 - 1.0
+
+ adr = tendon_adr[tenid]
+ if wrap_type[adr] == wp.static(types.WrapType.JOINT.value):
+ ten_num = tendon_num[tenid]
+ for i in range(ten_num):
+ dofadr = jnt_dofadr[wrap_objid[adr + i]]
+ J = scl * ten_J_in[worldid, tenid, dofadr]
+ efc_J_out[worldid, efcid, dofadr] = J
+ Jqvel += J * qvel_in[worldid, dofadr]
+ else:
+ for i in range(nv):
+ J = scl * ten_J_in[worldid, tenid, i]
+ efc_J_out[worldid, efcid, i] = J
+ Jqvel += J * qvel_in[worldid, i]
+
+ _update_efc_row(
+ worldid,
+ timestep,
+ refsafe_in,
+ efcid,
+ pos,
+ pos,
+ tendon_invweight0[worldid, tenid],
+ tendon_solref_lim[worldid, tenid],
+ tendon_solimp_lim[worldid, tenid],
+ tenmargin,
+ Jqvel,
+ 0.0,
+ ConstraintType.LIMIT_TENDON.value,
+ tenid,
+ efc_type_out,
+ efc_id_out,
+ efc_pos_out,
+ efc_margin_out,
+ efc_D_out,
+ efc_vel_out,
+ efc_aref_out,
+ efc_frictionloss_out,
+ )
+
+
+@wp.kernel
+def _efc_contact_pyramidal(
+ # Model:
+ nv: int,
+ opt_timestep: wp.array(dtype=float),
+ opt_impratio: wp.array(dtype=float),
+ body_parentid: wp.array(dtype=int),
+ body_rootid: wp.array(dtype=int),
+ body_invweight0: wp.array2d(dtype=wp.vec2),
+ dof_bodyid: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ # Data in:
+ njmax_in: int,
+ ncon_in: wp.array(dtype=int),
+ qvel_in: wp.array2d(dtype=float),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ refsafe_in: int,
+ dist_in: wp.array(dtype=float),
+ condim_in: wp.array(dtype=int),
+ includemargin_in: wp.array(dtype=float),
+ worldid_in: wp.array(dtype=int),
+ geom_in: wp.array(dtype=wp.vec2i),
+ pos_in: wp.array(dtype=wp.vec3),
+ frame_in: wp.array(dtype=wp.mat33),
+ friction_in: wp.array(dtype=vec5),
+ solref_in: wp.array(dtype=wp.vec2),
+ solimp_in: wp.array(dtype=vec5),
+ # Data out:
+ nefc_out: wp.array(dtype=int),
+ contact_efc_address_out: wp.array2d(dtype=int),
+ efc_type_out: wp.array2d(dtype=int),
+ efc_id_out: wp.array2d(dtype=int),
+ efc_J_out: wp.array3d(dtype=float),
+ efc_pos_out: wp.array2d(dtype=float),
+ efc_margin_out: wp.array2d(dtype=float),
+ efc_D_out: wp.array2d(dtype=float),
+ efc_vel_out: wp.array2d(dtype=float),
+ efc_aref_out: wp.array2d(dtype=float),
+ efc_frictionloss_out: wp.array2d(dtype=float),
+):
+ conid, dimid = wp.tid()
+
+ if conid >= ncon_in[0]:
+ return
+
+ condim = condim_in[conid]
+
+ if condim == 1 and dimid > 0:
+ return
+ elif condim > 1 and dimid >= 2 * (condim - 1):
+ return
+
+ includemargin = includemargin_in[conid]
+ pos = dist_in[conid] - includemargin
+ active = pos < 0
+
+ if active:
+ worldid = worldid_in[conid]
+ efcid = wp.atomic_add(nefc_out, worldid, 1)
+
+ if efcid >= njmax_in:
+ return
+
+ timestep = opt_timestep[worldid]
+ impratio = opt_impratio[worldid]
+ contact_efc_address_out[conid, dimid] = efcid
+
+ geom = geom_in[conid]
+ body1 = geom_bodyid[geom[0]]
+ body2 = geom_bodyid[geom[1]]
+
+ con_pos = pos_in[conid]
+ frame = frame_in[conid]
+
+ # pyramidal has common invweight across all edges
+ invweight = body_invweight0[worldid, body1][0] + body_invweight0[worldid, body2][0]
+
+ if condim > 1:
+ dimid2 = dimid / 2 + 1
+
+ friction = friction_in[conid]
+ fri0 = friction[0]
+ frii = friction[dimid2 - 1]
+ invweight = invweight + fri0 * fri0 * invweight
+ invweight = invweight * 2.0 * fri0 * fri0 / impratio
+
+ Jqvel = float(0.0)
+ for i in range(nv):
+ J = float(0.0)
+ Ji = float(0.0)
+ jac1p, jac1r = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ con_pos,
+ body1,
+ i,
+ worldid,
+ )
+ jac2p, jac2r = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ con_pos,
+ body2,
+ i,
+ worldid,
+ )
+ jacp_dif = jac2p - jac1p
+ for xyz in range(3):
+ J += frame[0, xyz] * jacp_dif[xyz]
+
+ if condim > 1:
+ if dimid2 < 3:
+ Ji += frame[dimid2, xyz] * jacp_dif[xyz]
+ else:
+ Ji += frame[dimid2 - 3, xyz] * (jac2r[xyz] - jac1r[xyz])
+
+ if condim > 1:
+ if dimid % 2 == 0:
+ J += Ji * frii
+ else:
+ J -= Ji * frii
+
+ efc_J_out[worldid, efcid, i] = J
+ Jqvel += J * qvel_in[worldid, i]
+
+ if condim == 1:
+ efc_type = int(ConstraintType.CONTACT_FRICTIONLESS.value)
+ else:
+ efc_type = int(ConstraintType.CONTACT_PYRAMIDAL.value)
+
+ _update_efc_row(
+ worldid,
+ timestep,
+ refsafe_in,
+ efcid,
+ pos,
+ pos,
+ invweight,
+ solref_in[conid],
+ solimp_in[conid],
+ includemargin,
+ Jqvel,
+ 0.0,
+ efc_type,
+ conid,
+ efc_type_out,
+ efc_id_out,
+ efc_pos_out,
+ efc_margin_out,
+ efc_D_out,
+ efc_vel_out,
+ efc_aref_out,
+ efc_frictionloss_out,
+ )
+
+
+@wp.kernel
+def _efc_contact_elliptic(
+ # Model:
+ nv: int,
+ opt_timestep: wp.array(dtype=float),
+ opt_impratio: wp.array(dtype=float),
+ body_parentid: wp.array(dtype=int),
+ body_rootid: wp.array(dtype=int),
+ body_invweight0: wp.array2d(dtype=wp.vec2),
+ dof_bodyid: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ # Data in:
+ njmax_in: int,
+ ncon_in: wp.array(dtype=int),
+ qvel_in: wp.array2d(dtype=float),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ refsafe_in: int,
+ dist_in: wp.array(dtype=float),
+ condim_in: wp.array(dtype=int),
+ includemargin_in: wp.array(dtype=float),
+ worldid_in: wp.array(dtype=int),
+ geom_in: wp.array(dtype=wp.vec2i),
+ pos_in: wp.array(dtype=wp.vec3),
+ frame_in: wp.array(dtype=wp.mat33),
+ friction_in: wp.array(dtype=vec5),
+ solref_in: wp.array(dtype=wp.vec2),
+ solreffriction_in: wp.array(dtype=wp.vec2),
+ solimp_in: wp.array(dtype=vec5),
+ # Data out:
+ nefc_out: wp.array(dtype=int),
+ contact_efc_address_out: wp.array2d(dtype=int),
+ efc_type_out: wp.array2d(dtype=int),
+ efc_id_out: wp.array2d(dtype=int),
+ efc_J_out: wp.array3d(dtype=float),
+ efc_pos_out: wp.array2d(dtype=float),
+ efc_margin_out: wp.array2d(dtype=float),
+ efc_D_out: wp.array2d(dtype=float),
+ efc_vel_out: wp.array2d(dtype=float),
+ efc_aref_out: wp.array2d(dtype=float),
+ efc_frictionloss_out: wp.array2d(dtype=float),
+):
+ conid, dimid = wp.tid()
+
+ if conid >= ncon_in[0]:
+ return
+
+ condim = condim_in[conid]
+
+ if dimid > condim - 1:
+ return
+
+ includemargin = includemargin_in[conid]
+ pos = dist_in[conid] - includemargin
+ active = pos < 0.0
+
+ if active:
+ worldid = worldid_in[conid]
+ efcid = wp.atomic_add(nefc_out, worldid, 1)
+
+ if efcid >= njmax_in:
+ return
+
+ timestep = opt_timestep[worldid]
+ impratio = opt_impratio[worldid]
+ contact_efc_address_out[conid, dimid] = efcid
+
+ geom = geom_in[conid]
+ body1 = geom_bodyid[geom[0]]
+ body2 = geom_bodyid[geom[1]]
+
+ cpos = pos_in[conid]
+ frame = frame_in[conid]
+
+ # TODO(team): parallelize J and Jqvel computation?
+ Jqvel = float(0.0)
+ for i in range(nv):
+ J = float(0.0)
+ jac1p, jac1r = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ cpos,
+ body1,
+ i,
+ worldid,
+ )
+ jac2p, jac2r = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ cpos,
+ body2,
+ i,
+ worldid,
+ )
+ for xyz in range(3):
+ if dimid < 3:
+ jac_dif = jac2p[xyz] - jac1p[xyz]
+ J += frame[dimid, xyz] * jac_dif
+ else:
+ jac_dif = jac2r[xyz] - jac1r[xyz]
+ J += frame[dimid - 3, xyz] * jac_dif
+
+ efc_J_out[worldid, efcid, i] = J
+ Jqvel += J * qvel_in[worldid, i]
+
+ invweight = body_invweight0[worldid, body1][0] + body_invweight0[worldid, body2][0]
+
+ ref = solref_in[conid]
+ pos_aref = pos
+
+ if dimid > 0:
+ solreffriction = solreffriction_in[conid]
+
+ # non-normal directions use solreffriction (if non-zero)
+ if solreffriction[0] or solreffriction[1]:
+ ref = solreffriction
+
+ # TODO(team): precompute 1 / impratio
+ invweight = invweight / impratio
+ friction = friction_in[conid]
+
+ if dimid > 1:
+ fri0 = friction[0]
+ frii = friction[dimid - 1]
+ fri = fri0 * fri0 / (frii * frii)
+ invweight *= fri
+
+ pos_aref = 0.0
+
+ if condim == 1:
+ efc_type = int(ConstraintType.CONTACT_FRICTIONLESS.value)
+ else:
+ efc_type = int(ConstraintType.CONTACT_ELLIPTIC.value)
+
+ _update_efc_row(
+ worldid,
+ timestep,
+ refsafe_in,
+ efcid,
+ pos_aref,
+ pos,
+ invweight,
+ ref,
+ solimp_in[conid],
+ includemargin,
+ Jqvel,
+ 0.0,
+ efc_type,
+ conid,
+ efc_type_out,
+ efc_id_out,
+ efc_pos_out,
+ efc_margin_out,
+ efc_D_out,
+ efc_vel_out,
+ efc_aref_out,
+ efc_frictionloss_out,
+ )
+
+
+@wp.kernel
+def _num_equality(
+ # Data in:
+ ne_connect_in: wp.array(dtype=int),
+ ne_weld_in: wp.array(dtype=int),
+ ne_jnt_in: wp.array(dtype=int),
+ ne_ten_in: wp.array(dtype=int),
+ # Data out:
+ ne_out: wp.array(dtype=int),
+):
+ worldid = wp.tid()
+ ne = ne_connect_in[worldid] + ne_weld_in[worldid] + ne_jnt_in[worldid] + ne_ten_in[worldid]
+ ne_out[worldid] = ne
+
+
+@event_scope
+def make_constraint(m: types.Model, d: types.Data):
+ """Creates constraint jacobians and other supporting data."""
+
+ wp.launch(
+ zero_constraint_counts,
+ dim=d.nworld,
+ inputs=[
+ d.ne,
+ d.ne_connect,
+ d.ne_weld,
+ d.ne_jnt,
+ d.ne_ten,
+ d.nf,
+ d.nl,
+ d.nefc,
+ ],
+ )
+
+ if not (m.opt.disableflags & types.DisableBit.CONSTRAINT.value):
+ refsafe = m.opt.disableflags & types.DisableBit.REFSAFE
+
+ if not (m.opt.disableflags & types.DisableBit.EQUALITY.value):
+ wp.launch(
+ _efc_equality_connect,
+ dim=(d.nworld, m.eq_connect_adr.size),
+ inputs=[
+ m.nv,
+ m.nsite,
+ m.opt.timestep,
+ m.body_parentid,
+ m.body_rootid,
+ m.body_invweight0,
+ m.dof_bodyid,
+ m.site_bodyid,
+ m.eq_obj1id,
+ m.eq_obj2id,
+ m.eq_objtype,
+ m.eq_solref,
+ m.eq_solimp,
+ m.eq_data,
+ m.eq_connect_adr,
+ d.njmax,
+ d.qvel,
+ d.eq_active,
+ d.xpos,
+ d.xmat,
+ d.site_xpos,
+ d.subtree_com,
+ d.cdof,
+ refsafe,
+ ],
+ outputs=[
+ d.ne_connect,
+ d.nefc,
+ d.efc.type,
+ d.efc.id,
+ d.efc.J,
+ d.efc.pos,
+ d.efc.margin,
+ d.efc.D,
+ d.efc.vel,
+ d.efc.aref,
+ d.efc.frictionloss,
+ ],
+ )
+ wp.launch(
+ _efc_equality_weld,
+ dim=(d.nworld, m.eq_wld_adr.size),
+ inputs=[
+ m.nv,
+ m.nsite,
+ m.opt.timestep,
+ m.body_parentid,
+ m.body_rootid,
+ m.body_invweight0,
+ m.dof_bodyid,
+ m.site_bodyid,
+ m.site_quat,
+ m.eq_obj1id,
+ m.eq_obj2id,
+ m.eq_objtype,
+ m.eq_solref,
+ m.eq_solimp,
+ m.eq_data,
+ m.eq_wld_adr,
+ d.njmax,
+ d.qvel,
+ d.eq_active,
+ d.xpos,
+ d.xquat,
+ d.xmat,
+ d.site_xpos,
+ d.subtree_com,
+ d.cdof,
+ refsafe,
+ ],
+ outputs=[
+ d.ne_weld,
+ d.nefc,
+ d.efc.type,
+ d.efc.id,
+ d.efc.J,
+ d.efc.pos,
+ d.efc.margin,
+ d.efc.D,
+ d.efc.vel,
+ d.efc.aref,
+ d.efc.frictionloss,
+ ],
+ )
+ wp.launch(
+ _efc_equality_joint,
+ dim=(d.nworld, m.eq_jnt_adr.size),
+ inputs=[
+ m.nv,
+ m.opt.timestep,
+ m.qpos0,
+ m.jnt_qposadr,
+ m.jnt_dofadr,
+ m.dof_invweight0,
+ m.eq_obj1id,
+ m.eq_obj2id,
+ m.eq_solref,
+ m.eq_solimp,
+ m.eq_data,
+ m.eq_jnt_adr,
+ d.njmax,
+ d.qpos,
+ d.qvel,
+ d.eq_active,
+ refsafe,
+ ],
+ outputs=[
+ d.ne_jnt,
+ d.nefc,
+ d.efc.type,
+ d.efc.id,
+ d.efc.J,
+ d.efc.pos,
+ d.efc.margin,
+ d.efc.D,
+ d.efc.vel,
+ d.efc.aref,
+ d.efc.frictionloss,
+ ],
+ )
+ wp.launch(
+ _efc_equality_tendon,
+ dim=(d.nworld, m.eq_ten_adr.size),
+ inputs=[
+ m.nv,
+ m.opt.timestep,
+ m.eq_obj1id,
+ m.eq_obj2id,
+ m.eq_solref,
+ m.eq_solimp,
+ m.eq_data,
+ m.eq_ten_adr,
+ m.tendon_length0,
+ m.tendon_invweight0,
+ d.njmax,
+ d.qvel,
+ d.eq_active,
+ d.ten_length,
+ d.ten_J,
+ refsafe,
+ ],
+ outputs=[
+ d.ne_ten,
+ d.nefc,
+ d.efc.type,
+ d.efc.id,
+ d.efc.J,
+ d.efc.pos,
+ d.efc.margin,
+ d.efc.D,
+ d.efc.vel,
+ d.efc.aref,
+ d.efc.frictionloss,
+ ],
+ )
+
+ wp.launch(
+ _num_equality,
+ dim=(d.nworld,),
+ inputs=[
+ d.ne_connect,
+ d.ne_weld,
+ d.ne_jnt,
+ d.ne_ten,
+ ],
+ outputs=[
+ d.ne,
+ ],
+ )
+
+ if not (m.opt.disableflags & types.DisableBit.FRICTIONLOSS.value):
+ wp.launch(
+ _efc_friction_dof,
+ dim=(d.nworld, m.nv),
+ inputs=[
+ m.nv,
+ m.opt.timestep,
+ m.dof_invweight0,
+ m.dof_frictionloss,
+ m.dof_solimp,
+ m.dof_solref,
+ d.njmax,
+ d.qvel,
+ refsafe,
+ ],
+ outputs=[
+ d.nf,
+ d.nefc,
+ d.efc.type,
+ d.efc.id,
+ d.efc.J,
+ d.efc.pos,
+ d.efc.margin,
+ d.efc.D,
+ d.efc.vel,
+ d.efc.aref,
+ d.efc.frictionloss,
+ ],
+ )
+
+ wp.launch(
+ _efc_friction_tendon,
+ dim=(d.nworld, m.ntendon),
+ inputs=[
+ m.nv,
+ m.opt.timestep,
+ m.tendon_solref_fri,
+ m.tendon_solimp_fri,
+ m.tendon_frictionloss,
+ m.tendon_invweight0,
+ d.qvel,
+ d.ten_J,
+ refsafe,
+ ],
+ outputs=[
+ d.nf,
+ d.nefc,
+ d.efc.type,
+ d.efc.id,
+ d.efc.J,
+ d.efc.pos,
+ d.efc.margin,
+ d.efc.D,
+ d.efc.vel,
+ d.efc.aref,
+ d.efc.frictionloss,
+ ],
+ )
+
+ # limit
+ if not (m.opt.disableflags & types.DisableBit.LIMIT.value):
+ limit_ball = m.jnt_limited_ball_adr.size > 0
+ if limit_ball:
+ wp.launch(
+ _efc_limit_ball,
+ dim=(d.nworld, m.jnt_limited_ball_adr.size),
+ inputs=[
+ m.nv,
+ m.opt.timestep,
+ m.jnt_qposadr,
+ m.jnt_dofadr,
+ m.jnt_solref,
+ m.jnt_solimp,
+ m.jnt_range,
+ m.jnt_margin,
+ m.jnt_limited_ball_adr,
+ m.dof_invweight0,
+ d.njmax,
+ d.qpos,
+ d.qvel,
+ refsafe,
+ ],
+ outputs=[
+ d.nl,
+ d.nefc,
+ d.efc.type,
+ d.efc.id,
+ d.efc.J,
+ d.efc.pos,
+ d.efc.margin,
+ d.efc.D,
+ d.efc.vel,
+ d.efc.aref,
+ d.efc.frictionloss,
+ ],
+ )
+
+ limit_slide_hinge = m.jnt_limited_slide_hinge_adr.size > 0
+ if limit_slide_hinge:
+ wp.launch(
+ _efc_limit_slide_hinge,
+ dim=(d.nworld, m.jnt_limited_slide_hinge_adr.size),
+ inputs=[
+ m.nv,
+ m.opt.timestep,
+ m.jnt_qposadr,
+ m.jnt_dofadr,
+ m.jnt_solref,
+ m.jnt_solimp,
+ m.jnt_range,
+ m.jnt_margin,
+ m.jnt_limited_slide_hinge_adr,
+ m.dof_invweight0,
+ d.njmax,
+ d.qpos,
+ d.qvel,
+ refsafe,
+ ],
+ outputs=[
+ d.nl,
+ d.nefc,
+ d.efc.type,
+ d.efc.id,
+ d.efc.J,
+ d.efc.pos,
+ d.efc.margin,
+ d.efc.D,
+ d.efc.vel,
+ d.efc.aref,
+ d.efc.frictionloss,
+ ],
+ )
+
+ limit_tendon = m.tendon_limited_adr.size > 0
+ if limit_tendon:
+ wp.launch(
+ _efc_limit_tendon,
+ dim=(d.nworld, m.tendon_limited_adr.size),
+ inputs=[
+ m.nv,
+ m.opt.timestep,
+ m.jnt_dofadr,
+ m.tendon_adr,
+ m.tendon_num,
+ m.tendon_limited_adr,
+ m.tendon_solref_lim,
+ m.tendon_solimp_lim,
+ m.tendon_range,
+ m.tendon_margin,
+ m.tendon_invweight0,
+ m.wrap_objid,
+ m.wrap_type,
+ d.njmax,
+ d.qvel,
+ d.ten_length,
+ d.ten_J,
+ refsafe,
+ ],
+ outputs=[
+ d.nl,
+ d.nefc,
+ d.efc.type,
+ d.efc.id,
+ d.efc.J,
+ d.efc.pos,
+ d.efc.margin,
+ d.efc.D,
+ d.efc.vel,
+ d.efc.aref,
+ d.efc.frictionloss,
+ ],
+ )
+
+ # contact
+ if not (m.opt.disableflags & types.DisableBit.CONTACT.value):
+ if m.opt.cone == types.ConeType.PYRAMIDAL.value:
+ wp.launch(
+ _efc_contact_pyramidal,
+ dim=(d.nconmax, 2 * (m.condim_max - 1) if m.condim_max > 1 else 1),
+ inputs=[
+ m.nv,
+ m.opt.timestep,
+ m.opt.impratio,
+ m.body_parentid,
+ m.body_rootid,
+ m.body_invweight0,
+ m.dof_bodyid,
+ m.geom_bodyid,
+ d.njmax,
+ d.ncon,
+ d.qvel,
+ d.subtree_com,
+ d.cdof,
+ refsafe,
+ d.contact.dist,
+ d.contact.dim,
+ d.contact.includemargin,
+ d.contact.worldid,
+ d.contact.geom,
+ d.contact.pos,
+ d.contact.frame,
+ d.contact.friction,
+ d.contact.solref,
+ d.contact.solimp,
+ ],
+ outputs=[
+ d.nefc,
+ d.contact.efc_address,
+ d.efc.type,
+ d.efc.id,
+ d.efc.J,
+ d.efc.pos,
+ d.efc.margin,
+ d.efc.D,
+ d.efc.vel,
+ d.efc.aref,
+ d.efc.frictionloss,
+ ],
+ )
+ elif m.opt.cone == types.ConeType.ELLIPTIC.value:
+ wp.launch(
+ _efc_contact_elliptic,
+ dim=(d.nconmax, m.condim_max),
+ inputs=[
+ m.nv,
+ m.opt.timestep,
+ m.opt.impratio,
+ m.body_parentid,
+ m.body_rootid,
+ m.body_invweight0,
+ m.dof_bodyid,
+ m.geom_bodyid,
+ d.njmax,
+ d.ncon,
+ d.qvel,
+ d.subtree_com,
+ d.cdof,
+ refsafe,
+ d.contact.dist,
+ d.contact.dim,
+ d.contact.includemargin,
+ d.contact.worldid,
+ d.contact.geom,
+ d.contact.pos,
+ d.contact.frame,
+ d.contact.friction,
+ d.contact.solref,
+ d.contact.solreffriction,
+ d.contact.solimp,
+ ],
+ outputs=[
+ d.nefc,
+ d.contact.efc_address,
+ d.efc.type,
+ d.efc.id,
+ d.efc.J,
+ d.efc.pos,
+ d.efc.margin,
+ d.efc.D,
+ d.efc.vel,
+ d.efc.aref,
+ d.efc.frictionloss,
+ ],
+ )
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint_test.py
new file mode 100644
index 00000000..28cffc8c
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint_test.py
@@ -0,0 +1,217 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 constraint functions."""
+
+import mujoco
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+from absl.testing import parameterized
+
+import mujoco_warp as mjwarp
+
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType
+
+# tolerance for difference between MuJoCo and MJWarp constraint calculations,
+# mostly due to float precision
+_TOLERANCE = 5e-5
+
+
+def _assert_eq(a, b, name):
+ tol = _TOLERANCE * 10 # avoid test noise
+ err_msg = f"mismatch: {name}"
+ np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+
+
+class ConstraintTest(parameterized.TestCase):
+ @parameterized.parameters(
+ (ConeType.PYRAMIDAL, 1, 1),
+ (ConeType.PYRAMIDAL, 1, 3),
+ (ConeType.PYRAMIDAL, 1, 4),
+ (ConeType.PYRAMIDAL, 1, 6),
+ (ConeType.PYRAMIDAL, 3, 3),
+ (ConeType.PYRAMIDAL, 3, 4),
+ (ConeType.PYRAMIDAL, 3, 6),
+ (ConeType.PYRAMIDAL, 4, 4),
+ (ConeType.PYRAMIDAL, 4, 6),
+ (ConeType.PYRAMIDAL, 6, 6),
+ (ConeType.ELLIPTIC, 1, 1),
+ (ConeType.ELLIPTIC, 1, 3),
+ (ConeType.ELLIPTIC, 1, 4),
+ (ConeType.ELLIPTIC, 1, 6),
+ (ConeType.ELLIPTIC, 3, 3),
+ (ConeType.ELLIPTIC, 3, 4),
+ (ConeType.ELLIPTIC, 3, 6),
+ (ConeType.ELLIPTIC, 4, 4),
+ (ConeType.ELLIPTIC, 4, 6),
+ (ConeType.ELLIPTIC, 6, 6),
+ )
+ def test_condim(self, cone, condim1, condim2):
+ """Test condim."""
+ xml = f"""
+
+
+
+
+
+
+
+
+
+
+
+
+ """
+
+ _, mjd, m, d = test_util.fixture(xml=xml, cone=cone)
+
+ for arr in (
+ d.efc.D,
+ d.efc.aref,
+ d.efc.pos,
+ d.efc.margin,
+ ):
+ arr.zero_()
+
+ # fill with nan to check whether we are not reading uninitialized values
+ d.efc.J.fill_(wp.nan)
+
+ mjwarp.make_constraint(m, d)
+
+ _assert_eq(d.efc.J.numpy()[0, : mjd.nefc, :].reshape(-1), mjd.efc_J, "efc_J")
+ _assert_eq(d.efc.D.numpy()[0, : mjd.nefc], mjd.efc_D, "efc_D")
+ _assert_eq(d.efc.aref.numpy()[0, : mjd.nefc], mjd.efc_aref, "efc_aref")
+ _assert_eq(d.efc.pos.numpy()[0, : mjd.nefc], mjd.efc_pos, "efc_pos")
+ _assert_eq(d.efc.margin.numpy()[0, : mjd.nefc], mjd.efc_margin, "efc_margin")
+
+ @parameterized.parameters(
+ mujoco.mjtCone.mjCONE_PYRAMIDAL,
+ mujoco.mjtCone.mjCONE_ELLIPTIC,
+ )
+ def test_constraints(self, cone):
+ """Test constraints."""
+ for key in range(3):
+ mjm, mjd, m, d = test_util.fixture("constraints.xml", sparse=False, cone=cone, keyframe=key)
+
+ for arr in (
+ d.efc.D,
+ d.efc.aref,
+ d.efc.pos,
+ d.efc.margin,
+ d.ne,
+ d.nefc,
+ d.nf,
+ d.nl,
+ ):
+ arr.zero_()
+
+ d.efc.J.fill_(wp.nan)
+
+ mjwarp.make_constraint(m, d)
+
+ _assert_eq(d.ne.numpy()[0], mjd.ne, "ne")
+ _assert_eq(d.nefc.numpy()[0], mjd.nefc, "nefc")
+ _assert_eq(d.nf.numpy()[0], mjd.nf, "nf")
+ _assert_eq(d.nl.numpy()[0], mjd.nl, "nl")
+ _assert_eq(d.efc.J.numpy()[0, : mjd.nefc, :].reshape(-1), mjd.efc_J, "efc_J")
+ _assert_eq(d.efc.D.numpy()[0, : mjd.nefc], mjd.efc_D, "efc_D")
+ _assert_eq(d.efc.vel.numpy()[0, : mjd.nefc], mjd.efc_vel, "efc_vel")
+ _assert_eq(d.efc.aref.numpy()[0, : mjd.nefc], mjd.efc_aref, "efc_aref")
+ _assert_eq(d.efc.pos.numpy()[0, : mjd.nefc], mjd.efc_pos, "efc_pos")
+ _assert_eq(d.efc.margin.numpy()[0, : mjd.nefc], mjd.efc_margin, "efc_margin")
+ _assert_eq(d.efc.type.numpy()[0, : mjd.nefc], mjd.efc_type, "efc_type")
+
+ def test_limit_tendon(self):
+ """Test limit tendon constraints."""
+ for keyframe in range(-1, 1):
+ _, mjd, m, d = test_util.fixture("tendon/tendon_limit.xml", sparse=False, keyframe=keyframe)
+
+ for arr in (d.nefc, d.nl, d.efc.J, d.efc.D, d.efc.aref, d.efc.pos, d.efc.margin):
+ arr.zero_()
+
+ mjwarp.make_constraint(m, d)
+
+ _assert_eq(d.nefc.numpy()[0], mjd.nefc, "nefc")
+ _assert_eq(d.nl.numpy()[0], mjd.nl, "nl")
+ _assert_eq(d.efc.J.numpy()[0, : mjd.nefc, :].reshape(-1), mjd.efc_J, "efc_J")
+ _assert_eq(d.efc.D.numpy()[0, : mjd.nefc], mjd.efc_D, "efc_D")
+ _assert_eq(d.efc.vel.numpy()[0, : mjd.nefc], mjd.efc_vel, "efc_vel")
+ _assert_eq(d.efc.aref.numpy()[0, : mjd.nefc], mjd.efc_aref, "efc_aref")
+ _assert_eq(d.efc.pos.numpy()[0, : mjd.nefc], mjd.efc_pos, "efc_pos")
+ _assert_eq(d.efc.margin.numpy()[0, : mjd.nefc], mjd.efc_margin, "efc_margin")
+ _assert_eq(d.efc.type.numpy()[0, : mjd.nefc], mjd.efc_type, "efc_type")
+
+ def test_equality_tendon(self):
+ """Test equality tendon constraints."""
+
+ _, mjd, m, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """
+ )
+
+ mjwarp.make_constraint(m, d)
+
+ _assert_eq(d.nefc.numpy()[0], mjd.nefc, "nefc")
+ _assert_eq(d.ne.numpy()[0], mjd.ne, "ne")
+ _assert_eq(d.efc.J.numpy()[0, : mjd.nefc, :].reshape(-1), mjd.efc_J, "efc_J")
+ _assert_eq(d.efc.D.numpy()[0, : mjd.nefc], mjd.efc_D, "efc_D")
+ _assert_eq(d.efc.vel.numpy()[0, : mjd.nefc], mjd.efc_vel, "efc_vel")
+ _assert_eq(d.efc.aref.numpy()[0, : mjd.nefc], mjd.efc_aref, "efc_aref")
+ _assert_eq(d.efc.pos.numpy()[0, : mjd.nefc], mjd.efc_pos, "efc_pos")
+ _assert_eq(d.efc.margin.numpy()[0, : mjd.nefc], mjd.efc_margin, "efc_margin")
+ _assert_eq(d.efc.type.numpy()[0, : mjd.nefc], mjd.efc_type, "efc_type")
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py
new file mode 100644
index 00000000..dcd70f4b
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py
@@ -0,0 +1,206 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src.support import mul_m
+from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import DynType
+from mujoco.mjx.third_party.mujoco_warp._src.types import GainType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec10f
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
+
+wp.set_module_options({"enable_backward": False})
+
+
+# TODO(team): improve performance with tile operations?
+@wp.kernel
+def _qderiv_actuator_passive(
+ # Model:
+ nu: int,
+ opt_timestep: wp.array(dtype=float),
+ opt_is_sparse: bool,
+ dof_damping: wp.array2d(dtype=float),
+ actuator_dyntype: wp.array(dtype=int),
+ actuator_gaintype: wp.array(dtype=int),
+ actuator_biastype: wp.array(dtype=int),
+ actuator_actadr: wp.array(dtype=int),
+ actuator_actnum: wp.array(dtype=int),
+ actuator_gainprm: wp.array2d(dtype=vec10f),
+ actuator_biasprm: wp.array2d(dtype=vec10f),
+ # Data in:
+ act_in: wp.array2d(dtype=float),
+ ctrl_in: wp.array2d(dtype=float),
+ actuator_moment_in: wp.array3d(dtype=float),
+ qM_in: wp.array3d(dtype=float),
+ # In:
+ qMi: wp.array(dtype=int),
+ qMj: wp.array(dtype=int),
+ actuation_enabled: bool,
+ passive_enabled: bool,
+ # Data out:
+ qM_integration_out: wp.array3d(dtype=float),
+):
+ worldid, elemid = wp.tid()
+ dofiid = qMi[elemid]
+ dofjid = qMj[elemid]
+
+ qderiv = float(0.0)
+ for actid in range(nu):
+ if actuation_enabled:
+ if actuator_gaintype[actid] == int(GainType.AFFINE.value):
+ gain = actuator_gainprm[worldid, actid][2]
+ else:
+ gain = 0.0
+
+ if actuator_biastype[actid] == int(BiasType.AFFINE.value):
+ bias = actuator_biasprm[worldid, actid][2]
+ else:
+ bias = 0.0
+
+ if actuator_dyntype[actid] != int(DynType.NONE.value):
+ act_first = actuator_actadr[actid]
+ act_last = act_first + actuator_actnum[actid] - 1
+ vel = bias + gain * act_in[worldid, act_last]
+ else:
+ vel = bias + gain * ctrl_in[worldid, actid]
+
+ qderiv += actuator_moment_in[worldid, actid, dofiid] * actuator_moment_in[worldid, actid, dofjid] * vel
+
+ if passive_enabled and dofiid == dofjid:
+ qderiv -= dof_damping[worldid, dofiid] / float(nu)
+
+ qderiv *= opt_timestep[worldid]
+
+ if opt_is_sparse:
+ qM_integration_out[worldid, 0, elemid] = qM_in[worldid, 0, elemid] - qderiv
+ else:
+ qM = qM_in[worldid, dofiid, dofjid] - qderiv
+ qM_integration_out[worldid, dofiid, dofjid] = qM
+ if dofiid != dofjid:
+ qM_integration_out[worldid, dofjid, dofiid] = qM
+
+
+# TODO(team): improve performance with tile operations?
+@wp.kernel
+def _qderiv_tendon_damping(
+ # Model:
+ ntendon: int,
+ opt_timestep: wp.array(dtype=float),
+ opt_is_sparse: bool,
+ tendon_damping: wp.array2d(dtype=float),
+ # Data in:
+ ten_J_in: wp.array3d(dtype=float),
+ # In:
+ qMi: wp.array(dtype=int),
+ qMj: wp.array(dtype=int),
+ # Data out:
+ qM_integration_out: wp.array3d(dtype=float),
+):
+ worldid, elemid = wp.tid()
+ dofiid = qMi[elemid]
+ dofjid = qMj[elemid]
+
+ qderiv = float(0.0)
+ for tenid in range(ntendon):
+ qderiv -= ten_J_in[worldid, tenid, dofiid] * ten_J_in[worldid, tenid, dofjid] * tendon_damping[worldid, tenid]
+ qderiv *= opt_timestep[worldid]
+
+ if opt_is_sparse:
+ qM_integration_out[worldid, 0, elemid] -= qderiv
+ else:
+ qM_integration_out[worldid, dofiid, dofjid] -= qderiv
+ if dofiid != dofjid:
+ qM_integration_out[worldid, dofjid, dofiid] -= qderiv
+
+
+@wp.kernel
+def _qfrc_forward(
+ # Data in:
+ qfrc_smooth_in: wp.array2d(dtype=float),
+ qfrc_constraint_in: wp.array2d(dtype=float),
+ # Data out:
+ qfrc_integration_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ qfrc_integration_out[worldid, dofid] = qfrc_smooth_in[worldid, dofid] + qfrc_constraint_in[worldid, dofid]
+
+
+@event_scope
+def deriv_smooth_vel(m: Model, d: Data, flg_forward: bool = True):
+ """Analytical derivative of smooth forces w.r.t. velocities.
+
+ Args:
+ m (Model): The model containing kinematic and dynamic information (device).
+ d (Data): The data object containing the current state and output arrays (device).
+ flg_forward (bool, optional): If True forward dynamics else inverse dynamics routine.
+ Default is True.
+ """
+ actuation_enabled = not (m.opt.disableflags & DisableBit.ACTUATION)
+ passive_enabled = not (m.opt.disableflags & DisableBit.PASSIVE)
+
+ qMi = m.qM_fullm_i if m.opt.is_sparse else m.dof_tri_row
+ qMj = m.qM_fullm_j if m.opt.is_sparse else m.dof_tri_col
+ if actuation_enabled or passive_enabled:
+ wp.launch(
+ _qderiv_actuator_passive,
+ dim=(d.nworld, qMi.size),
+ inputs=[
+ m.nu,
+ m.opt.timestep,
+ m.opt.is_sparse,
+ m.dof_damping,
+ m.actuator_dyntype,
+ m.actuator_gaintype,
+ m.actuator_biastype,
+ m.actuator_actadr,
+ m.actuator_actnum,
+ m.actuator_gainprm,
+ m.actuator_biasprm,
+ d.act,
+ d.ctrl,
+ d.actuator_moment,
+ d.qM,
+ qMi,
+ qMj,
+ actuation_enabled,
+ passive_enabled,
+ ],
+ outputs=[d.qM_integration],
+ )
+
+ if passive_enabled:
+ wp.launch(
+ _qderiv_tendon_damping,
+ dim=(d.nworld, qMi.size),
+ inputs=[m.ntendon, m.opt.timestep, m.opt.is_sparse, m.tendon_damping, d.ten_J, qMi, qMj],
+ outputs=[d.qM_integration],
+ )
+
+ if flg_forward:
+ wp.launch(
+ _qfrc_forward,
+ dim=(d.nworld, m.nv),
+ inputs=[d.qfrc_smooth, d.qfrc_constraint],
+ outputs=[d.qfrc_integration],
+ )
+ else:
+ # qfrc = qM @ qacc
+ mul_m(m, d, d.qfrc_integration, d.qacc, d.inverse_mul_m_skip, d.qM_integration)
+
+ # TODO(team): rne derivative
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py
new file mode 100644
index 00000000..a27c38a1
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py
@@ -0,0 +1,1037 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+from typing import Optional
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src import collision_driver
+from mujoco.mjx.third_party.mujoco_warp._src import constraint
+from mujoco.mjx.third_party.mujoco_warp._src import derivative
+from mujoco.mjx.third_party.mujoco_warp._src import math
+from mujoco.mjx.third_party.mujoco_warp._src import passive
+from mujoco.mjx.third_party.mujoco_warp._src import sensor
+from mujoco.mjx.third_party.mujoco_warp._src import smooth
+from mujoco.mjx.third_party.mujoco_warp._src import solver
+from mujoco.mjx.third_party.mujoco_warp._src import util_misc
+from mujoco.mjx.third_party.mujoco_warp._src.support import xfrc_accumulate
+from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
+from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import DynType
+from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import GainType
+from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType
+from mujoco.mjx.third_party.mujoco_warp._src.types import JointType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+from mujoco.mjx.third_party.mujoco_warp._src.types import TileSet
+from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec10f
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel
+
+wp.set_module_options({"enable_backward": False})
+
+# RK4 tableau
+_RK4_A = [
+ [0.5, 0.0, 0.0],
+ [0.0, 0.5, 0.0],
+ [0.0, 0.0, 1.0],
+]
+_RK4_B = [1.0 / 6.0, 1.0 / 3.0, 1.0 / 3.0, 1.0 / 6.0]
+
+
+@wp.kernel
+def _next_position(
+ # Model:
+ opt_timestep: wp.array(dtype=float),
+ jnt_type: wp.array(dtype=int),
+ jnt_qposadr: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ # Data in:
+ qpos_in: wp.array2d(dtype=float),
+ qvel_in: wp.array2d(dtype=float),
+ # In:
+ qvel_scale_in: float,
+ # Data out:
+ qpos_out: wp.array2d(dtype=float),
+):
+ worldid, jntid = wp.tid()
+ timestep = opt_timestep[worldid]
+
+ jnttype = jnt_type[jntid]
+ qpos_adr = jnt_qposadr[jntid]
+ dof_adr = jnt_dofadr[jntid]
+ qpos = qpos_in[worldid]
+ qpos_next = qpos_out[worldid]
+ qvel = qvel_in[worldid]
+
+ if jnttype == wp.static(JointType.FREE.value):
+ qpos_pos = wp.vec3(qpos[qpos_adr], qpos[qpos_adr + 1], qpos[qpos_adr + 2])
+ qvel_lin = wp.vec3(qvel[dof_adr], qvel[dof_adr + 1], qvel[dof_adr + 2]) * qvel_scale_in
+
+ qpos_new = qpos_pos + timestep * qvel_lin
+
+ qpos_quat = wp.quat(
+ qpos[qpos_adr + 3],
+ qpos[qpos_adr + 4],
+ qpos[qpos_adr + 5],
+ qpos[qpos_adr + 6],
+ )
+ qvel_ang = wp.vec3(qvel[dof_adr + 3], qvel[dof_adr + 4], qvel[dof_adr + 5]) * qvel_scale_in
+
+ qpos_quat_new = math.quat_integrate(qpos_quat, qvel_ang, timestep)
+
+ qpos_next[qpos_adr + 0] = qpos_new[0]
+ qpos_next[qpos_adr + 1] = qpos_new[1]
+ qpos_next[qpos_adr + 2] = qpos_new[2]
+ qpos_next[qpos_adr + 3] = qpos_quat_new[0]
+ qpos_next[qpos_adr + 4] = qpos_quat_new[1]
+ qpos_next[qpos_adr + 5] = qpos_quat_new[2]
+ qpos_next[qpos_adr + 6] = qpos_quat_new[3]
+
+ elif jnttype == wp.static(JointType.BALL.value):
+ qpos_quat = wp.quat(
+ qpos[qpos_adr + 0],
+ qpos[qpos_adr + 1],
+ qpos[qpos_adr + 2],
+ qpos[qpos_adr + 3],
+ )
+ qvel_ang = wp.vec3(qvel[dof_adr], qvel[dof_adr + 1], qvel[dof_adr + 2]) * qvel_scale_in
+
+ qpos_quat_new = math.quat_integrate(qpos_quat, qvel_ang, timestep)
+
+ qpos_next[qpos_adr + 0] = qpos_quat_new[0]
+ qpos_next[qpos_adr + 1] = qpos_quat_new[1]
+ qpos_next[qpos_adr + 2] = qpos_quat_new[2]
+ qpos_next[qpos_adr + 3] = qpos_quat_new[3]
+
+ else: # if jnt_type in (JointType.HINGE, JointType.SLIDE):
+ qpos_next[qpos_adr] = qpos[qpos_adr] + timestep * qvel[dof_adr] * qvel_scale_in
+
+
+@wp.kernel
+def _next_velocity(
+ # Model:
+ opt_timestep: wp.array(dtype=float),
+ # Data in:
+ qvel_in: wp.array2d(dtype=float),
+ qacc_in: wp.array2d(dtype=float),
+ # In:
+ qacc_scale_in: float,
+ # Data out:
+ qvel_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ timestep = opt_timestep[worldid]
+ qvel_out[worldid, dofid] = qvel_in[worldid, dofid] + qacc_scale_in * qacc_in[worldid, dofid] * timestep
+
+
+# TODO(team): kernel analyzer array slice?
+@wp.func
+def _next_act(
+ # Model:
+ opt_timestep: float, # kernel_analyzer: ignore
+ actuator_dyntype: int, # kernel_analyzer: ignore
+ actuator_dynprm: vec10f, # kernel_analyzer: ignore
+ actuator_actrange: wp.vec2, # kernel_analyzer: ignore
+ # Data In:
+ act_in: float, # kernel_analyzer: ignore
+ act_dot_in: float, # kernel_analyzer: ignore
+ # In:
+ act_dot_scale: float,
+ clamp: bool,
+) -> float:
+ # advance actuation
+ if actuator_dyntype == wp.static(DynType.FILTEREXACT.value):
+ tau = wp.max(MJ_MINVAL, actuator_dynprm[0])
+ act = act_in + act_dot_scale * act_dot_in * tau * (1.0 - wp.exp(-opt_timestep / tau))
+ else:
+ act = act_in + act_dot_scale * act_dot_in * opt_timestep
+
+ # clamp to actrange
+ if clamp:
+ act = wp.clamp(act, actuator_actrange[0], actuator_actrange[1])
+
+ return act
+
+
+@wp.kernel
+def _next_activation(
+ # Model:
+ opt_timestep: wp.array(dtype=float),
+ actuator_dyntype: wp.array(dtype=int),
+ actuator_actlimited: wp.array(dtype=bool),
+ actuator_dynprm: wp.array2d(dtype=vec10f),
+ actuator_actrange: wp.array2d(dtype=wp.vec2),
+ # Data in:
+ act_in: wp.array2d(dtype=float),
+ act_dot_in: wp.array2d(dtype=float),
+ # In:
+ act_dot_scale: float,
+ limit: bool,
+ # Data out:
+ act_out: wp.array2d(dtype=float),
+):
+ worldid, actid = wp.tid()
+ act = _next_act(
+ opt_timestep[worldid],
+ actuator_dyntype[actid],
+ actuator_dynprm[worldid, actid],
+ actuator_actrange[worldid, actid],
+ act_in[worldid, actid],
+ act_dot_in[worldid, actid],
+ act_dot_scale,
+ limit and actuator_actlimited[actid],
+ )
+ act_out[worldid, actid] = act
+
+
+@wp.kernel
+def _next_time(
+ # Model:
+ opt_timestep: wp.array(dtype=float),
+ # Data in:
+ nconmax_in: int,
+ njmax_in: int,
+ ncon_in: wp.array(dtype=int),
+ nefc_in: wp.array(dtype=int),
+ time_in: wp.array(dtype=float),
+ ncollision_in: wp.array(dtype=int),
+ # Data out:
+ time_out: wp.array(dtype=float),
+):
+ worldid = wp.tid()
+ time_out[worldid] = time_in[worldid] + opt_timestep[worldid]
+ nefc = nefc_in[worldid]
+
+ if nefc > njmax_in:
+ wp.printf("nefc overflow - please increase njmax to %u\n", nefc)
+
+ if worldid == 0:
+ ncollision = ncollision_in[0]
+ if ncollision > nconmax_in:
+ wp.printf("ncollision overflow - please increase nconmax to %u\n", ncollision)
+
+ if ncon_in[0] > nconmax_in:
+ wp.printf("ncon overflow - please increase nconmax to %u\n", ncon_in[0])
+
+
+def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None):
+ """Advance state and time given activation derivatives and acceleration."""
+
+ # TODO(team): can we assume static timesteps?
+
+ # advance activations
+ if m.na:
+ wp.launch(
+ _next_activation,
+ dim=(d.nworld, m.na),
+ inputs=[
+ m.opt.timestep,
+ m.actuator_dyntype,
+ m.actuator_actlimited,
+ m.actuator_dynprm,
+ m.actuator_actrange,
+ d.act,
+ d.act_dot,
+ 1.0,
+ True,
+ ],
+ outputs=[
+ d.act,
+ ],
+ )
+
+ wp.launch(
+ _next_velocity,
+ dim=(d.nworld, m.nv),
+ inputs=[
+ m.opt.timestep,
+ d.qvel,
+ qacc,
+ 1.0,
+ ],
+ outputs=[
+ d.qvel,
+ ],
+ )
+
+ # advance positions with qvel if given, d.qvel otherwise (semi-implicit)
+ if qvel is not None:
+ qvel_in = qvel
+ else:
+ qvel_in = d.qvel
+
+ wp.launch(
+ _next_position,
+ dim=(d.nworld, m.njnt),
+ inputs=[
+ m.opt.timestep,
+ m.jnt_type,
+ m.jnt_qposadr,
+ m.jnt_dofadr,
+ d.qpos,
+ qvel_in,
+ 1.0,
+ ],
+ outputs=[
+ d.qpos,
+ ],
+ )
+
+ wp.launch(
+ _next_time,
+ dim=(d.nworld,),
+ inputs=[
+ m.opt.timestep,
+ d.nconmax,
+ d.njmax,
+ d.ncon,
+ d.nefc,
+ d.time,
+ d.ncollision,
+ ],
+ outputs=[
+ d.time,
+ ],
+ )
+
+
+@wp.kernel
+def _euler_damp_qfrc_sparse(
+ # Model:
+ opt_timestep: wp.array(dtype=float),
+ dof_Madr: wp.array(dtype=int),
+ dof_damping: wp.array2d(dtype=float),
+ # Data in:
+ qfrc_smooth_in: wp.array2d(dtype=float),
+ qfrc_constraint_in: wp.array2d(dtype=float),
+ # Data out:
+ qfrc_integration_out: wp.array2d(dtype=float),
+ qM_integration_out: wp.array3d(dtype=float),
+):
+ worldid, tid = wp.tid()
+ timestep = opt_timestep[worldid]
+
+ adr = dof_Madr[tid]
+ qM_integration_out[worldid, 0, adr] += timestep * dof_damping[worldid, tid]
+ qfrc_integration_out[worldid, tid] = qfrc_smooth_in[worldid, tid] + qfrc_constraint_in[worldid, tid]
+
+
+def _euler_sparse(m: Model, d: Data):
+ wp.copy(d.qM_integration, d.qM)
+ wp.launch(
+ _euler_damp_qfrc_sparse,
+ dim=(d.nworld, m.nv),
+ inputs=[
+ m.opt.timestep,
+ m.dof_Madr,
+ m.dof_damping,
+ d.qfrc_smooth,
+ d.qfrc_constraint,
+ ],
+ outputs=[
+ d.qfrc_integration,
+ d.qM_integration,
+ ],
+ )
+ smooth.factor_solve_i(
+ m,
+ d,
+ d.qM_integration,
+ d.qLD_integration,
+ d.qLDiagInv_integration,
+ d.qacc_integration,
+ d.qfrc_integration,
+ )
+
+
+@cache_kernel
+def _tile_euler_dense(tile: TileSet):
+ @nested_kernel
+ def euler_dense(
+ # Model:
+ dof_damping: wp.array2d(dtype=float),
+ opt_timestep: wp.array(dtype=float),
+ # Data in:
+ qM_in: wp.array3d(dtype=float),
+ qfrc_smooth_in: wp.array2d(dtype=float),
+ qfrc_constraint_in: wp.array2d(dtype=float),
+ # In:
+ adr_in: wp.array(dtype=int),
+ # Data out:
+ qacc_integration_out: wp.array2d(dtype=float),
+ ):
+ worldid, nodeid = wp.tid()
+ timestep = opt_timestep[worldid]
+ TILE_SIZE = wp.static(tile.size)
+
+ dofid = adr_in[nodeid]
+ M_tile = wp.tile_load(qM_in[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid))
+ damping_tile = wp.tile_load(dof_damping[worldid], shape=(TILE_SIZE,), offset=(dofid,))
+ damping_scaled = damping_tile * timestep
+ qm_integration_tile = wp.tile_diag_add(M_tile, damping_scaled)
+
+ qfrc_smooth_tile = wp.tile_load(qfrc_smooth_in[worldid], shape=(TILE_SIZE,), offset=(dofid,))
+ qfrc_constraint_tile = wp.tile_load(qfrc_constraint_in[worldid], shape=(TILE_SIZE,), offset=(dofid,))
+
+ qfrc_tile = qfrc_smooth_tile + qfrc_constraint_tile
+
+ L_tile = wp.tile_cholesky(qm_integration_tile)
+ qacc_tile = wp.tile_cholesky_solve(L_tile, qfrc_tile)
+ wp.tile_store(qacc_integration_out[worldid], qacc_tile, offset=(dofid))
+
+ return euler_dense
+
+
+@event_scope
+def euler(m: Model, d: Data):
+ """Euler integrator, semi-implicit in velocity."""
+
+ # integrate damping implicitly
+ if not m.opt.disableflags & DisableBit.EULERDAMP.value:
+ if m.opt.is_sparse:
+ _euler_sparse(m, d)
+ else:
+ for tile in m.qM_tiles:
+ wp.launch_tiled(
+ _tile_euler_dense(tile),
+ dim=(d.nworld, tile.adr.size),
+ inputs=[m.dof_damping, m.opt.timestep, d.qM, d.qfrc_smooth, d.qfrc_constraint, tile.adr],
+ outputs=[d.qacc_integration],
+ block_dim=m.block_dim.euler_dense,
+ )
+
+ _advance(m, d, d.qacc_integration)
+ else:
+ _advance(m, d, d.qacc)
+
+
+def _rk_perturb_state(m: Model, d: Data, scale: float):
+ # position
+ wp.launch(
+ _next_position,
+ dim=(d.nworld, m.njnt),
+ inputs=[m.opt.timestep, m.jnt_type, m.jnt_qposadr, m.jnt_dofadr, d.qpos_t0, d.qvel, scale],
+ outputs=[d.qpos],
+ )
+
+ # velocity
+ wp.launch(
+ _next_velocity,
+ dim=(d.nworld, m.nv),
+ inputs=[m.opt.timestep, d.qvel_t0, d.qacc, scale],
+ outputs=[d.qvel],
+ )
+
+ # activation
+ if m.na:
+ wp.launch(
+ _next_activation,
+ dim=(d.nworld, m.na),
+ inputs=[m.opt.timestep, d.act_t0, d.act_dot, scale, False],
+ outputs=[d.act],
+ )
+
+
+@wp.kernel
+def _rk_accumulate_velocity_acceleration(
+ # Data in:
+ qvel_in: wp.array2d(dtype=float),
+ qacc_in: wp.array2d(dtype=float),
+ # In:
+ scale: float,
+ # Data out:
+ qvel_out: wp.array2d(dtype=float),
+ qacc_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ qvel_out[worldid, dofid] += scale * qvel_in[worldid, dofid]
+ qacc_out[worldid, dofid] += scale * qacc_in[worldid, dofid]
+
+
+@wp.kernel
+def _rk_accumulate_activation_velocity(
+ # Data in:
+ act_dot_in: wp.array2d(dtype=float),
+ # In:
+ scale: float,
+ # Data out:
+ act_dot_out: wp.array2d(dtype=float),
+):
+ worldid, actid = wp.tid()
+ act_dot_out[worldid, actid] += scale * act_dot_in[worldid, actid]
+
+
+def _rk_accumulate(m: Model, d: Data, scale: float):
+ """Computes one term of 1/6 k_1 + 1/3 k_2 + 1/3 k_3 + 1/6 k_4"""
+
+ wp.launch(
+ _rk_accumulate_velocity_acceleration,
+ dim=(d.nworld, m.nv),
+ inputs=[d.qvel, d.qacc, scale],
+ outputs=[d.qvel_rk, d.qacc_rk],
+ )
+
+ if m.na:
+ wp.launch(
+ _rk_accumulate_activation_velocity,
+ dim=(d.nworld, m.na),
+ inputs=[d.act_dot, scale],
+ outputs=[d.act_dot_rk],
+ )
+
+
+@event_scope
+def rungekutta4(m: Model, d: Data):
+ """Runge-Kutta explicit order 4 integrator."""
+
+ wp.copy(d.qpos_t0, d.qpos)
+ wp.copy(d.qvel_t0, d.qvel)
+
+ d.qvel_rk.zero_()
+ d.qacc_rk.zero_()
+ d.act_dot_rk.zero_()
+
+ if m.na:
+ wp.copy(d.act_t0, d.act)
+
+ A, B = _RK4_A, _RK4_B
+
+ _rk_accumulate(m, d, B[0])
+ for i in range(3):
+ a, b = float(A[i][i]), B[i + 1]
+ _rk_perturb_state(m, d, a)
+ forward(m, d)
+ _rk_accumulate(m, d, b)
+
+ wp.copy(d.qpos, d.qpos_t0)
+ wp.copy(d.qvel, d.qvel_t0)
+ if m.na:
+ wp.copy(d.act, d.act_t0)
+ wp.copy(d.act_dot, d.act_dot_rk)
+ _advance(m, d, d.qacc_rk, d.qvel_rk)
+
+
+@event_scope
+def implicit(m: Model, d: Data):
+ """Integrates fully implicit in velocity."""
+
+ # compile-time constants
+ passive_enabled = not m.opt.disableflags & DisableBit.PASSIVE.value
+ actuation_enabled = (not m.opt.disableflags & DisableBit.ACTUATION.value) and m.actuator_affine_bias_gain
+
+ if passive_enabled or actuation_enabled:
+ derivative.deriv_smooth_vel(m, d)
+ smooth.factor_solve_i(
+ m, d, d.qM_integration, d.qLD_integration, d.qLDiagInv_integration, d.qacc_integration, d.qfrc_integration
+ )
+ _advance(m, d, d.qacc_integration)
+ else:
+ _advance(m, d, d.qacc)
+
+
+@event_scope
+def fwd_position(m: Model, d: Data, factorize: bool = True):
+ """Position-dependent computations."""
+
+ smooth.kinematics(m, d)
+ smooth.com_pos(m, d)
+ smooth.camlight(m, d)
+ smooth.tendon(m, d)
+ smooth.crb(m, d)
+ smooth.tendon_armature(m, d)
+ if factorize:
+ smooth.factor_m(m, d)
+ if m.opt.run_collision_detection:
+ collision_driver.collision(m, d)
+ constraint.make_constraint(m, d)
+ smooth.transmission(m, d)
+
+
+# TODO(team): sparse actuator_moment version
+def _actuator_velocity(m: Model, d: Data):
+ NV = m.nv
+
+ @kernel
+ def actuator_velocity(
+ # Data in:
+ qvel_in: wp.array2d(dtype=float),
+ actuator_moment_in: wp.array3d(dtype=float),
+ # Data out:
+ actuator_velocity_out: wp.array2d(dtype=float),
+ ):
+ worldid, actid = wp.tid()
+ moment_tile = wp.tile_load(actuator_moment_in[worldid, actid], shape=NV)
+ qvel_tile = wp.tile_load(qvel_in[worldid], shape=NV)
+ moment_qvel_tile = wp.tile_map(wp.mul, moment_tile, qvel_tile)
+ actuator_velocity_tile = wp.tile_reduce(wp.add, moment_qvel_tile)
+ actuator_velocity_out[worldid, actid] = actuator_velocity_tile[0]
+
+ wp.launch_tiled(
+ actuator_velocity,
+ dim=(d.nworld, m.nu),
+ inputs=[
+ d.qvel,
+ d.actuator_moment,
+ ],
+ outputs=[
+ d.actuator_velocity,
+ ],
+ block_dim=m.block_dim.actuator_velocity,
+ )
+
+
+def _tendon_velocity(m: Model, d: Data):
+ NV = m.nv
+
+ @kernel
+ def tendon_velocity(
+ # Data in:
+ qvel_in: wp.array2d(dtype=float),
+ ten_J_in: wp.array3d(dtype=float),
+ # Data out:
+ ten_velocity_out: wp.array2d(dtype=float),
+ ):
+ worldid, tenid = wp.tid()
+ ten_J_tile = wp.tile_load(ten_J_in[worldid, tenid], shape=NV)
+ qvel_tile = wp.tile_load(qvel_in[worldid], shape=NV)
+ ten_J_qvel_tile = wp.tile_map(wp.mul, ten_J_tile, qvel_tile)
+ ten_velocity_tile = wp.tile_reduce(wp.add, ten_J_qvel_tile)
+ ten_velocity_out[worldid, tenid] = ten_velocity_tile[0]
+
+ wp.launch_tiled(
+ tendon_velocity,
+ dim=(d.nworld, m.ntendon),
+ inputs=[
+ d.qvel,
+ d.ten_J,
+ ],
+ outputs=[
+ d.ten_velocity,
+ ],
+ block_dim=m.block_dim.tendon_velocity,
+ )
+
+
+@event_scope
+def fwd_velocity(m: Model, d: Data):
+ """Velocity-dependent computations."""
+
+ _actuator_velocity(m, d)
+
+ if m.ntendon > 0:
+ # TODO(team): sparse version
+ _tendon_velocity(m, d)
+
+ smooth.com_vel(m, d)
+ passive.passive(m, d)
+ smooth.rne(m, d)
+ smooth.tendon_bias(m, d, d.qfrc_bias)
+
+
+@wp.kernel
+def _actuator_force(
+ # Model:
+ na: int,
+ opt_timestep: wp.array(dtype=float),
+ actuator_dyntype: wp.array(dtype=int),
+ actuator_gaintype: wp.array(dtype=int),
+ actuator_biastype: wp.array(dtype=int),
+ actuator_actadr: wp.array(dtype=int),
+ actuator_actnum: wp.array(dtype=int),
+ actuator_ctrllimited: wp.array(dtype=bool),
+ actuator_forcelimited: wp.array(dtype=bool),
+ actuator_actlimited: wp.array(dtype=bool),
+ actuator_dynprm: wp.array2d(dtype=vec10f),
+ actuator_gainprm: wp.array2d(dtype=vec10f),
+ actuator_biasprm: wp.array2d(dtype=vec10f),
+ actuator_actearly: wp.array(dtype=bool),
+ actuator_ctrlrange: wp.array2d(dtype=wp.vec2),
+ actuator_forcerange: wp.array2d(dtype=wp.vec2),
+ actuator_actrange: wp.array2d(dtype=wp.vec2),
+ actuator_acc0: wp.array(dtype=float),
+ actuator_lengthrange: wp.array(dtype=wp.vec2),
+ # Data in:
+ act_in: wp.array2d(dtype=float),
+ ctrl_in: wp.array2d(dtype=float),
+ actuator_length_in: wp.array2d(dtype=float),
+ actuator_velocity_in: wp.array2d(dtype=float),
+ # In:
+ dsbl_clampctrl: int,
+ # Data out:
+ act_dot_out: wp.array2d(dtype=float),
+ actuator_force_out: wp.array2d(dtype=float),
+):
+ worldid, uid = wp.tid()
+
+ ctrl = ctrl_in[worldid, uid]
+
+ if actuator_ctrllimited[uid] and not dsbl_clampctrl:
+ ctrlrange = actuator_ctrlrange[worldid, uid]
+ ctrl = wp.clamp(ctrl, ctrlrange[0], ctrlrange[1])
+ ctrl_act = ctrl
+
+ act_first = actuator_actadr[uid]
+ if na and act_first >= 0:
+ act_last = act_first + actuator_actnum[uid] - 1
+ dyntype = actuator_dyntype[uid]
+
+ if dyntype == int(DynType.INTEGRATOR.value):
+ act_dot = ctrl
+ elif dyntype == int(DynType.FILTER.value) or dyntype == int(DynType.FILTEREXACT.value):
+ dynprm = actuator_dynprm[worldid, uid]
+ act = act_in[worldid, act_last]
+ act_dot = (ctrl - act) / wp.max(dynprm[0], MJ_MINVAL)
+ elif dyntype == int(DynType.MUSCLE.value):
+ dynprm = actuator_dynprm[worldid, uid]
+ act = act_in[worldid, act_last]
+ act_dot = util_misc.muscle_dynamics(ctrl, act, dynprm)
+ else: # DynType.NONE
+ act_dot = 0.0
+
+ act_dot_out[worldid, act_last] = act_dot
+
+ if actuator_actearly[uid]:
+ if dyntype == int(DynType.INTEGRATOR.value) or dyntype == int(DynType.NONE.value):
+ dynprm = actuator_dynprm[worldid, uid]
+ act = act_in[worldid, act_last]
+
+ ctrl_act = _next_act(
+ opt_timestep[worldid],
+ dyntype,
+ dynprm,
+ actuator_actrange[worldid, uid],
+ act,
+ act_dot,
+ 1.0,
+ actuator_actlimited[uid],
+ )
+ else:
+ ctrl_act = act_in[worldid, act_last]
+
+ length = actuator_length_in[worldid, uid]
+ velocity = actuator_velocity_in[worldid, uid]
+
+ # gain
+ gaintype = actuator_gaintype[uid]
+ gainprm = actuator_gainprm[worldid, uid]
+
+ gain = 0.0
+ if gaintype == int(GainType.FIXED.value):
+ gain = gainprm[0]
+ elif gaintype == int(GainType.AFFINE.value):
+ gain = gainprm[0] + gainprm[1] * length + gainprm[2] * velocity
+ elif gaintype == int(GainType.MUSCLE.value):
+ acc0 = actuator_acc0[uid]
+ lengthrange = actuator_lengthrange[uid]
+ gain = util_misc.muscle_gain(length, velocity, lengthrange, acc0, gainprm)
+
+ # bias
+ biastype = actuator_biastype[uid]
+ biasprm = actuator_biasprm[worldid, uid]
+
+ bias = 0.0 # BiasType.NONE
+ if biastype == int(BiasType.AFFINE.value):
+ bias = biasprm[0] + biasprm[1] * length + biasprm[2] * velocity
+ elif biastype == int(BiasType.MUSCLE.value):
+ acc0 = actuator_acc0[uid]
+ lengthrange = actuator_lengthrange[uid]
+ bias = util_misc.muscle_bias(length, lengthrange, acc0, biasprm)
+
+ force = gain * ctrl_act + bias
+
+ # TODO(team): tendon total force clamping
+
+ if actuator_forcelimited[uid]:
+ forcerange = actuator_forcerange[worldid, uid]
+ force = wp.clamp(force, forcerange[0], forcerange[1])
+
+ actuator_force_out[worldid, uid] = force
+
+
+@wp.kernel
+def _tendon_actuator_force(
+ # Model:
+ actuator_trntype: wp.array(dtype=int),
+ actuator_trnid: wp.array(dtype=wp.vec2i),
+ # Data in:
+ actuator_force_in: wp.array2d(dtype=float),
+ # Data out:
+ ten_actfrc_out: wp.array2d(dtype=float),
+):
+ worldid, actid = wp.tid()
+
+ if actuator_trntype[actid] == int(TrnType.TENDON.value):
+ tenid = actuator_trnid[actid][0]
+ # TODO(team): only compute for tendons with force limits?
+ wp.atomic_add(ten_actfrc_out[worldid], tenid, actuator_force_in[worldid, actid])
+
+
+@wp.kernel
+def _tendon_actuator_force_clamp(
+ # Model:
+ actuator_trntype: wp.array(dtype=int),
+ actuator_trnid: wp.array(dtype=wp.vec2i),
+ tendon_actfrclimited: wp.array(dtype=bool),
+ tendon_actfrcrange: wp.array2d(dtype=wp.vec2),
+ # Data in:
+ ten_actfrc_in: wp.array2d(dtype=float),
+ # Data out:
+ actuator_force_out: wp.array2d(dtype=float),
+):
+ worldid, actid = wp.tid()
+
+ if actuator_trntype[actid] == int(TrnType.TENDON.value):
+ tenid = actuator_trnid[actid][0]
+ if tendon_actfrclimited[tenid]:
+ ten_actfrc = ten_actfrc_in[worldid, tenid]
+ actfrcrange = tendon_actfrcrange[worldid, tenid]
+
+ if ten_actfrc < actfrcrange[0]:
+ actuator_force_out[worldid, actid] *= actfrcrange[0] / ten_actfrc
+ elif ten_actfrc > actfrcrange[1]:
+ actuator_force_out[worldid, actid] *= actfrcrange[1] / ten_actfrc
+
+
+def _qfrc_actuator(m: Model, d: Data):
+ NU = m.nu
+
+ @wp.kernel
+ def qfrc_actuator(
+ # Model:
+ ngravcomp: int,
+ jnt_actfrclimited: wp.array(dtype=bool),
+ jnt_actfrcrange: wp.array2d(dtype=wp.vec2),
+ jnt_actgravcomp: wp.array(dtype=int),
+ dof_jntid: wp.array(dtype=int),
+ # Data in:
+ actuator_moment_in: wp.array3d(dtype=float),
+ qfrc_gravcomp_in: wp.array2d(dtype=float),
+ actuator_force_in: wp.array2d(dtype=float),
+ # Data out:
+ qfrc_actuator_out: wp.array2d(dtype=float),
+ ):
+ worldid, dofid = wp.tid()
+
+ actuator_moment_tile = wp.tile_load(actuator_moment_in[worldid], shape=(NU, 1), offset=(0, dofid))
+ actuator_moment_tile = wp.tile_squeeze(actuator_moment_tile, axis=(1,))
+ actuator_force_tile = wp.tile_load(actuator_force_in[worldid], shape=NU)
+ actuator_moment_force_tile = wp.tile_map(wp.mul, actuator_moment_tile, actuator_force_tile)
+ qfrc_tile = wp.tile_reduce(wp.add, actuator_moment_force_tile)
+ qfrc = qfrc_tile[0]
+
+ jntid = dof_jntid[dofid]
+
+ # actuator-level gravity compensation, skip if added as passive force
+ if ngravcomp and jnt_actgravcomp[jntid]:
+ qfrc += qfrc_gravcomp_in[worldid, dofid]
+
+ if jnt_actfrclimited[jntid]:
+ frcrange = jnt_actfrcrange[worldid, jntid]
+ qfrc = wp.clamp(qfrc, frcrange[0], frcrange[1])
+
+ qfrc_actuator_out[worldid, dofid] = qfrc
+
+ wp.launch_tiled(
+ qfrc_actuator,
+ dim=(d.nworld, m.nv),
+ inputs=[
+ m.ngravcomp,
+ m.jnt_actfrclimited,
+ m.jnt_actfrcrange,
+ m.jnt_actgravcomp,
+ m.dof_jntid,
+ d.actuator_moment,
+ d.qfrc_gravcomp,
+ d.actuator_force,
+ ],
+ outputs=[d.qfrc_actuator],
+ block_dim=m.block_dim.qfrc_actuator,
+ )
+
+
+@event_scope
+def fwd_actuation(m: Model, d: Data):
+ """Actuation-dependent computations."""
+ if not m.nu or (m.opt.disableflags & DisableBit.ACTUATION):
+ d.act_dot.zero_()
+ d.qfrc_actuator.zero_()
+ return
+
+ wp.launch(
+ _actuator_force,
+ dim=(d.nworld, m.nu),
+ inputs=[
+ m.na,
+ m.opt.timestep,
+ m.actuator_dyntype,
+ m.actuator_gaintype,
+ m.actuator_biastype,
+ m.actuator_actadr,
+ m.actuator_actnum,
+ m.actuator_ctrllimited,
+ m.actuator_forcelimited,
+ m.actuator_actlimited,
+ m.actuator_dynprm,
+ m.actuator_gainprm,
+ m.actuator_biasprm,
+ m.actuator_actearly,
+ m.actuator_ctrlrange,
+ m.actuator_forcerange,
+ m.actuator_actrange,
+ m.actuator_acc0,
+ m.actuator_lengthrange,
+ d.act,
+ d.ctrl,
+ d.actuator_length,
+ d.actuator_velocity,
+ m.opt.disableflags & DisableBit.CLAMPCTRL,
+ ],
+ outputs=[d.act_dot, d.actuator_force],
+ )
+
+ if m.ntendon:
+ d.ten_actfrc.zero_()
+
+ wp.launch(
+ _tendon_actuator_force,
+ dim=(d.nworld, m.nu),
+ inputs=[
+ m.actuator_trntype,
+ m.actuator_trnid,
+ d.actuator_force,
+ ],
+ outputs=[d.ten_actfrc],
+ )
+
+ wp.launch(
+ _tendon_actuator_force_clamp,
+ dim=(d.nworld, m.nu),
+ inputs=[
+ m.actuator_trntype,
+ m.actuator_trnid,
+ m.tendon_actfrclimited,
+ m.tendon_actfrcrange,
+ d.ten_actfrc,
+ ],
+ outputs=[d.actuator_force],
+ )
+
+ _qfrc_actuator(m, d)
+
+
+@wp.kernel
+def _qfrc_smooth(
+ # Data in:
+ qfrc_applied_in: wp.array2d(dtype=float),
+ qfrc_bias_in: wp.array2d(dtype=float),
+ qfrc_passive_in: wp.array2d(dtype=float),
+ qfrc_actuator_in: wp.array2d(dtype=float),
+ # Data out:
+ qfrc_smooth_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ qfrc_smooth_out[worldid, dofid] = (
+ qfrc_passive_in[worldid, dofid]
+ - qfrc_bias_in[worldid, dofid]
+ + qfrc_actuator_in[worldid, dofid]
+ + qfrc_applied_in[worldid, dofid]
+ )
+
+
+@event_scope
+def fwd_acceleration(m: Model, d: Data, factorize: bool = False):
+ """Add up all non-constraint forces, compute qacc_smooth."""
+
+ wp.launch(
+ _qfrc_smooth,
+ dim=(d.nworld, m.nv),
+ inputs=[
+ d.qfrc_applied,
+ d.qfrc_bias,
+ d.qfrc_passive,
+ d.qfrc_actuator,
+ ],
+ outputs=[
+ d.qfrc_smooth,
+ ],
+ )
+ xfrc_accumulate(m, d, d.qfrc_smooth)
+
+ if factorize:
+ smooth.factor_solve_i(m, d, d.qM, d.qLD, d.qLDiagInv, d.qacc_smooth, d.qfrc_smooth)
+ else:
+ smooth.solve_m(m, d, d.qacc_smooth, d.qfrc_smooth)
+
+
+@wp.kernel
+def _zero_energy(
+ # Data out:
+ energy_out: wp.array(dtype=wp.vec2),
+):
+ tid = wp.tid()
+ energy_out[tid] = wp.vec2(0.0, 0.0)
+
+
+@event_scope
+def forward(m: Model, d: Data):
+ """Forward dynamics."""
+ energy = m.opt.enableflags & EnableBit.ENERGY
+
+ fwd_position(m, d, factorize=False)
+ sensor.sensor_pos(m, d)
+
+ if energy:
+ if m.sensor_e_potential == 0: # not computed by sensor
+ sensor.energy_pos(m, d)
+ else:
+ wp.launch(
+ _zero_energy,
+ dim=d.nworld,
+ inputs=[d.energy],
+ )
+
+ fwd_velocity(m, d)
+ sensor.sensor_vel(m, d)
+
+ if energy:
+ if m.sensor_e_kinetic == 0: # not computed by sensor
+ sensor.energy_vel(m, d)
+
+ fwd_actuation(m, d)
+ fwd_acceleration(m, d, factorize=True)
+ sensor.sensor_acc(m, d)
+
+ solver.solve(m, d)
+
+
+@event_scope
+def step(m: Model, d: Data):
+ """Advance simulation."""
+ forward(m, d)
+
+ if m.opt.integrator == IntegratorType.EULER:
+ euler(m, d)
+ elif m.opt.integrator == IntegratorType.RK4:
+ rungekutta4(m, d)
+ elif m.opt.integrator == IntegratorType.IMPLICITFAST:
+ implicit(m, d)
+ else:
+ raise NotImplementedError(f"integrator {m.opt.integrator} not implemented.")
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py
new file mode 100644
index 00000000..f18d4150
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py
@@ -0,0 +1,285 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 forward dynamics functions."""
+
+import mujoco
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+from absl.testing import parameterized
+
+import mujoco_warp as mjwarp
+
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType
+from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import GainType
+from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType
+
+# tolerance for difference between MuJoCo and mjwarp smooth calculations - mostly
+# due to float precision
+_TOLERANCE = 5e-5
+
+
+def _assert_eq(a, b, name):
+ tol = _TOLERANCE * 10 # avoid test noise
+ err_msg = f"mismatch: {name}"
+ np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+
+
+class ForwardTest(parameterized.TestCase):
+ # TODO(team): test sparse when actuator_moment and/or ten_J have sparse representation
+ @parameterized.product(xml=["humanoid/humanoid.xml", "pendula.xml"])
+ def test_fwd_velocity(self, xml):
+ _, mjd, m, d = test_util.fixture(xml, kick=True)
+
+ for arr in (d.actuator_velocity, d.qfrc_bias):
+ arr.zero_()
+
+ mjwarp.fwd_velocity(m, d)
+
+ _assert_eq(d.actuator_velocity.numpy()[0], mjd.actuator_velocity, "actuator_velocity")
+ _assert_eq(d.qfrc_bias.numpy()[0], mjd.qfrc_bias, "qfrc_bias")
+
+ def test_fwd_velocity_tendon(self):
+ _, mjd, m, d = test_util.fixture("tendon/fixed.xml", sparse=False)
+
+ d.ten_velocity.zero_()
+ mjwarp.fwd_velocity(m, d)
+
+ _assert_eq(d.ten_velocity.numpy()[0], mjd.ten_velocity, "ten_velocity")
+
+ @parameterized.parameters(
+ ("actuation/actuation.xml", True),
+ ("actuation/actuation.xml", False),
+ ("actuation/actuators.xml", True),
+ ("actuation/actuators.xml", False),
+ ("actuation/muscle.xml", True),
+ ("actuation/muscle.xml", False),
+ )
+ def test_actuation(self, xml, actuation):
+ mjm, mjd, m, d = test_util.fixture(xml, actuation=actuation, keyframe=0)
+
+ for arr in (d.qfrc_actuator, d.actuator_force, d.act_dot):
+ arr.zero_()
+
+ mjwarp.fwd_actuation(m, d)
+
+ _assert_eq(d.qfrc_actuator.numpy()[0], mjd.qfrc_actuator, "qfrc_actuator")
+ _assert_eq(d.actuator_force.numpy()[0], mjd.actuator_force, "actuator_force")
+
+ if mjm.na:
+ _assert_eq(d.act_dot.numpy()[0], mjd.act_dot, "act_dot")
+
+ # next activations
+ mujoco.mj_step(mjm, mjd)
+ mjwarp.step(m, d)
+
+ _assert_eq(d.act.numpy()[0], mjd.act, "act")
+
+ # TODO(team): test actearly
+
+ @parameterized.parameters(True, False)
+ def test_clampctrl(self, clampctrl):
+ _, mjd, _, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ clampctrl=clampctrl,
+ keyframe=0,
+ )
+
+ _assert_eq(d.ctrl.numpy()[0], mjd.ctrl, "ctrl")
+
+ def test_fwd_acceleration(self):
+ _, mjd, m, d = test_util.fixture("humanoid/humanoid.xml", kick=True)
+
+ for arr in (d.qfrc_smooth, d.qacc_smooth):
+ arr.zero_()
+
+ mjwarp.fwd_acceleration(m, d)
+
+ _assert_eq(d.qfrc_smooth.numpy()[0], mjd.qfrc_smooth, "qfrc_smooth")
+ _assert_eq(d.qacc_smooth.numpy()[0], mjd.qacc_smooth, "qacc_smooth")
+
+ @parameterized.parameters((True, True), (True, False), (False, True), (False, False))
+ def test_euler(self, eulerdamp, sparse):
+ mjm, mjd, _, _ = test_util.fixture("pendula.xml", kick=True, eulerdamp=eulerdamp, sparse=sparse)
+ self.assertTrue((mjm.dof_damping > 0).any())
+
+ mjd.qvel[:] = 1.0
+ mjd.qacc[:] = 1.0
+ mujoco.mj_forward(mjm, mjd)
+
+ m = mjwarp.put_model(mjm)
+ d = mjwarp.put_data(mjm, mjd)
+
+ mujoco.mj_Euler(mjm, mjd)
+ mjwarp.euler(m, d)
+
+ _assert_eq(d.qpos.numpy()[0], mjd.qpos, "qpos")
+ _assert_eq(d.qvel.numpy()[0], mjd.qvel, "qvel")
+ _assert_eq(d.act.numpy()[0], mjd.act, "act")
+
+ def test_rungekutta4(self):
+ mjm, mjd, m, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ keyframe=0,
+ )
+
+ mjwarp.rungekutta4(m, d)
+ mujoco.mj_RungeKutta(mjm, mjd, 4)
+
+ _assert_eq(d.qpos.numpy()[0], mjd.qpos, "qpos")
+ _assert_eq(d.qvel.numpy()[0], mjd.qvel, "qvel")
+ _assert_eq(d.time.numpy()[0], mjd.time, "time")
+ _assert_eq(d.xpos.numpy()[0], mjd.xpos, "xpos")
+
+ # test rungekutta determinism
+ def rk_step() -> wp.array(dtype=wp.float32, ndim=2):
+ d.qpos = wp.ones_like(d.qpos)
+ d.qvel = wp.ones_like(d.qvel)
+ d.act = wp.ones_like(d.act)
+ mjwarp.rungekutta4(m, d)
+ return d.qpos
+
+ _assert_eq(rk_step().numpy()[0], rk_step().numpy()[0], "qpos")
+
+ @parameterized.product(actuation=[True, False], passive=[True, False], sparse=[True, False])
+ def test_implicit(self, actuation, passive, sparse):
+ mjm, mjd, _, _ = test_util.fixture(
+ "pendula.xml",
+ integrator=IntegratorType.IMPLICITFAST,
+ actuation=actuation,
+ passive=passive,
+ sparse=sparse,
+ )
+
+ mjm.actuator_gainprm[:, 2] = np.random.uniform(low=0.01, high=10.0, size=mjm.actuator_gainprm[:, 2].shape)
+
+ # change actuators to velocity/damper to cover all codepaths
+ mjm.actuator_gaintype[3] = GainType.AFFINE
+ mjm.actuator_gaintype[6] = GainType.AFFINE
+ mjm.actuator_biastype[0:3] = BiasType.AFFINE
+ mjm.actuator_biastype[4:6] = BiasType.AFFINE
+ mjm.actuator_biasprm[0:3, 2] = -1.0
+ mjm.actuator_biasprm[4:6, 2] = -1.0
+ mjm.actuator_ctrlrange[3:7] = 10.0
+ mjm.actuator_gear[:] = 1.0
+
+ mjd.qvel = np.random.uniform(low=-0.01, high=0.01, size=mjd.qvel.shape)
+ mjd.ctrl = np.random.uniform(low=-0.1, high=0.1, size=mjd.ctrl.shape)
+ mjd.act = np.random.uniform(low=-0.1, high=0.1, size=mjd.act.shape)
+ mujoco.mj_forward(mjm, mjd)
+
+ m = mjwarp.put_model(mjm)
+ d = mjwarp.put_data(mjm, mjd)
+
+ mjwarp.implicit(m, d)
+ mujoco.mj_implicit(mjm, mjd)
+
+ _assert_eq(d.qpos.numpy()[0], mjd.qpos, "qpos")
+ _assert_eq(d.act.numpy()[0], mjd.act, "act")
+
+ def test_implicit_position(self):
+ mjm, mjd, m, d = test_util.fixture("actuation/position.xml", keyframe=0, integrator=IntegratorType.IMPLICITFAST, kick=True)
+
+ mujoco.mj_implicit(mjm, mjd)
+ mjwarp.implicit(m, d)
+
+ _assert_eq(d.qpos.numpy()[0], mjd.qpos, "qpos")
+ _assert_eq(d.qvel.numpy()[0], mjd.qvel, "qvel")
+
+ def test_implicit_tendon_damping(self):
+ mjm, mjd, m, d = test_util.fixture("tendon/damping.xml", keyframe=0, integrator=IntegratorType.IMPLICITFAST, kick=True)
+
+ mujoco.mj_implicit(mjm, mjd)
+ mjwarp.implicit(m, d)
+
+ _assert_eq(d.qpos.numpy()[0], mjd.qpos, "qpos")
+ _assert_eq(d.qvel.numpy()[0], mjd.qvel, "qvel")
+
+ @parameterized.product(
+ xml=("humanoid/humanoid.xml", "pendula.xml", "constraints.xml", "collision.xml"), graph_conditional=(True, False)
+ )
+ def test_graph_capture(self, xml, graph_conditional):
+ # TODO(team): test more environments
+ if wp.get_device().is_cuda and wp.config.verify_cuda == False:
+ _, _, m, d = test_util.fixture(xml)
+ m.opt.graph_conditional = graph_conditional
+
+ with wp.ScopedCapture() as capture:
+ mjwarp.step(m, d)
+
+ # step a few times to ensure no errors at the step boundary
+ wp.capture_launch(capture.graph)
+ wp.capture_launch(capture.graph)
+ wp.capture_launch(capture.graph)
+
+ self.assertTrue(d.time.numpy()[0] > 0.0)
+
+ def test_forward_energy(self):
+ _, mjd, _, d = test_util.fixture("humanoid/humanoid.xml", kick=True, energy=True)
+
+ _assert_eq(d.energy.numpy()[0][0], mjd.energy[0], "potential energy")
+ _assert_eq(d.energy.numpy()[0][1], mjd.energy[1], "kinetic energy")
+
+ def test_tendon_actuator_force_limits(self):
+ for keyframe in range(7):
+ _, mjd, m, d = test_util.fixture("actuation/tendon_force_limit.xml", keyframe=keyframe)
+
+ d.actuator_force.zero_()
+
+ mjwarp.forward(m, d)
+
+ _assert_eq(d.actuator_force.numpy()[0], mjd.actuator_force, "actuator_force")
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py
new file mode 100644
index 00000000..3856bfc1
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py
@@ -0,0 +1,147 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src import derivative
+from mujoco.mjx.third_party.mujoco_warp._src import forward
+from mujoco.mjx.third_party.mujoco_warp._src import sensor
+from mujoco.mjx.third_party.mujoco_warp._src import smooth
+from mujoco.mjx.third_party.mujoco_warp._src import solver
+from mujoco.mjx.third_party.mujoco_warp._src import support
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+
+
+@wp.kernel
+def _qfrc_eulerdamp(
+ # Model:
+ opt_timestep: wp.array(dtype=float),
+ dof_damping: wp.array2d(dtype=float),
+ # Data in:
+ qacc_in: wp.array2d(dtype=float),
+ # Out:
+ qfrc_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ timestep = opt_timestep[worldid]
+ qfrc_out[worldid, dofid] += timestep * dof_damping[worldid, dofid] * qacc_in[worldid, dofid]
+
+
+@wp.kernel
+def _qfrc_inverse(
+ # Data in:
+ qfrc_bias_in: wp.array2d(dtype=float),
+ qfrc_passive_in: wp.array2d(dtype=float),
+ qfrc_constraint_in: wp.array2d(dtype=float),
+ # In:
+ Ma: wp.array2d(dtype=float),
+ # Data out:
+ qfrc_inverse_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+
+ qfrc_inverse = qfrc_bias_in[worldid, dofid]
+ qfrc_inverse += Ma[worldid, dofid]
+ qfrc_inverse -= qfrc_passive_in[worldid, dofid]
+ qfrc_inverse -= qfrc_constraint_in[worldid, dofid]
+
+ qfrc_inverse_out[worldid, dofid] = qfrc_inverse
+
+
+def discrete_acc(m: Model, d: Data, qacc: wp.array2d(dtype=float), qfrc: wp.array2d(dtype=float)):
+ """Convert discrete-time qacc to continuous-time qacc."""
+
+ if m.opt.integrator == IntegratorType.RK4:
+ raise NotImplementedError("discrete inverse dynamics is not supported by RK4 integrator")
+ elif m.opt.integrator == IntegratorType.EULER:
+ if m.opt.disableflags & DisableBit.EULERDAMP:
+ wp.copy(qacc, d.qacc)
+ return
+
+ # TODO(team): qacc = d.qacc if (m.dof_damping == 0.0).all()
+
+ # set qfrc = (d.qM + m.opt.timestep * diag(m.dof_damping)) * d.qacc
+
+ # d.qM @ d.qacc
+ support.mul_m(m, d, qfrc, d.qacc, d.inverse_mul_m_skip)
+
+ # qfrc += m.opt.timestep * m.dof_damping * d.qacc
+ wp.launch(
+ _qfrc_eulerdamp,
+ dim=(d.nworld, m.nv),
+ inputs=[m.opt.timestep, m.dof_damping, d.qacc],
+ outputs=[qfrc],
+ )
+ elif m.opt.integrator == IntegratorType.IMPLICITFAST:
+ derivative.deriv_smooth_vel(m, d, flg_forward=False)
+ smooth.factor_solve_i(m, d, d.qM, d.qLD, d.qLDiagInv, qacc, qfrc)
+ else:
+ raise NotImplementedError(f"integrator {m.opt.integrator} not implemented.")
+
+ # solve for qacc: qfrc = d.qM @ d.qacc
+ smooth.solve_m(m, d, qacc, qfrc)
+
+
+def inv_constraint(m: Model, d: Data):
+ """Inverse constraint solver."""
+
+ # no constraints
+ if d.njmax == 0:
+ d.qfrc_constraint.zero_()
+ return
+
+ # update
+ solver.create_context(m, d, grad=False)
+
+
+def inverse(m: Model, d: Data):
+ """Inverse dynamics."""
+ forward.fwd_position(m, d)
+ sensor.sensor_pos(m, d)
+ forward.fwd_velocity(m, d)
+ sensor.sensor_vel(m, d)
+
+ invdiscrete = m.opt.enableflags & EnableBit.INVDISCRETE
+ if invdiscrete:
+ # save discrete-time qacc and compute continuous-time qacc
+ wp.copy(d.qacc_discrete, d.qacc)
+ discrete_acc(m, d, d.qacc, d.qfrc_integration)
+
+ inv_constraint(m, d)
+ smooth.rne(m, d)
+ smooth.tendon_bias(m, d, d.qfrc_bias)
+ sensor.sensor_acc(m, d)
+
+ support.mul_m(m, d, d.qfrc_inverse, d.qacc, d.inverse_mul_m_skip)
+
+ wp.launch(
+ _qfrc_inverse,
+ dim=(d.nworld, m.nv),
+ inputs=[
+ d.qfrc_bias,
+ d.qfrc_passive,
+ d.qfrc_constraint,
+ d.qfrc_inverse,
+ ],
+ outputs=[d.qfrc_inverse],
+ )
+
+ if invdiscrete:
+ # restore discrete-time qacc
+ wp.copy(d.qacc, d.qacc_discrete)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse_test.py
new file mode 100644
index 00000000..e6a7a065
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse_test.py
@@ -0,0 +1,169 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 inverse dynamics."""
+
+import mujoco
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+from absl.testing import parameterized
+
+import mujoco_warp as mjwarp
+
+from mujoco.mjx.third_party.mujoco_warp._src import inverse
+from mujoco.mjx.third_party.mujoco_warp._src import support
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType
+
+
+def _assert_eq(a, b, name):
+ tol = 5e-3 # avoid test noise
+ err_msg = f"mismatch: {name}"
+ np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+
+
+_XML = """
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+"""
+
+
+class InverseTest(parameterized.TestCase):
+ @parameterized.product(
+ integrator=[IntegratorType.EULER, IntegratorType.IMPLICITFAST],
+ invdiscrete=[True, False],
+ sparse=[True, False],
+ )
+ def test_inverse(self, integrator, invdiscrete, sparse):
+ """Tests inverse dynamics."""
+ mjm, mjd, m, d = test_util.fixture(
+ xml=_XML,
+ contact=False,
+ integrator=integrator,
+ kick=True,
+ applied=True,
+ nstep=10,
+ sparse=sparse,
+ )
+
+ # discrete qacc
+ if invdiscrete:
+ mjm.opt.enableflags |= mujoco.mjtEnableBit.mjENBL_INVDISCRETE
+
+ # save state
+ qpos = mjd.qpos.copy()
+ qvel = mjd.qvel.copy()
+
+ # call step, save new qvel
+ mujoco.mj_step(mjm, mjd)
+ qvel_next = mjd.qvel.copy()
+
+ # reset the state, compute discrete-time (finite-differenced) qacc
+ mjd.qpos = qpos
+ mjd.qvel = qvel
+ qacc_fd = (qvel_next - qvel) / mjm.opt.timestep
+
+ # call forward, overwrite qacc with qacc_fd
+ mujoco.mj_forward(mjm, mjd)
+ mjd.qacc = qacc_fd
+
+ m = mjwarp.put_model(mjm)
+ d = mjwarp.put_data(mjm, mjd)
+
+ qacc = d.qacc.numpy()[0].copy()
+ qfrc_constraint = d.qfrc_constraint.numpy()[0].copy()
+
+ # qfrc_inverse = qfrc_applied + J.T @ xfrc_applied + qfrc_actuator
+ qfrc_xfrc_applied = wp.zeros((d.nworld, m.nv), dtype=float)
+ support.xfrc_accumulate(m, d, qfrc_xfrc_applied)
+ qfrc_inverse = d.qfrc_applied.numpy()[0] + d.qfrc_actuator.numpy()[0] + qfrc_xfrc_applied.numpy()[0]
+
+ for arr in (d.qfrc_constraint, d.qfrc_inverse):
+ arr.zero_()
+
+ mjwarp.inverse(m, d)
+
+ _assert_eq(d.qfrc_constraint.numpy()[0], qfrc_constraint, "qfrc_constraint")
+ _assert_eq(d.qfrc_inverse.numpy()[0], qfrc_inverse, "qfrc_inverse")
+ _assert_eq(d.qacc.numpy()[0], qacc, "qacc")
+
+ def test_discrete_acc_eulerdamp(self):
+ _, _, m, d = test_util.fixture(
+ xml=_XML, integrator=IntegratorType.EULER, eulerdamp=False, kick=True, applied=True, nstep=10
+ )
+ qacc = wp.zeros((1, m.nv), dtype=float)
+ qfrc = wp.zeros((1, m.nv), dtype=float)
+ inverse.discrete_acc(m, d, qacc, qfrc)
+ _assert_eq(qacc.numpy()[0], d.qacc.numpy()[0], "qacc")
+
+ def test_discrete_acc_rk4(self):
+ _, _, m, d = test_util.fixture(xml=_XML, integrator=IntegratorType.RK4)
+ qacc = wp.zeros((1, m.nv), dtype=float)
+ qfrc = wp.zeros((1, m.nv), dtype=float)
+
+ with self.assertRaises(NotImplementedError):
+ inverse.discrete_acc(m, d, qacc, qfrc)
+
+ def test_inverse_tendon_armature(self):
+ """Tests inverse dynamics with tendon armature."""
+ _, _, m, d = test_util.fixture(
+ "tendon/armature.xml",
+ constraint=False,
+ gravity=False,
+ kick=True,
+ applied=True,
+ nstep=10,
+ keyframe=0,
+ )
+
+ qacc = d.qacc.numpy()[0].copy()
+ qfrc_constraint = d.qfrc_constraint.numpy()[0].copy()
+
+ # qfrc_inverse = qfrc_applied + J.T @ xfrc_applied + qfrc_actuator
+ qfrc_xfrc_applied = wp.zeros((d.nworld, m.nv), dtype=float)
+ support.xfrc_accumulate(m, d, qfrc_xfrc_applied)
+ qfrc_inverse = d.qfrc_applied.numpy()[0] + d.qfrc_actuator.numpy()[0] + qfrc_xfrc_applied.numpy()[0]
+
+ for arr in (d.qfrc_constraint, d.qfrc_inverse):
+ arr.zero_()
+
+ mjwarp.inverse(m, d)
+
+ _assert_eq(d.qfrc_constraint.numpy()[0], qfrc_constraint, "qfrc_constraint")
+ _assert_eq(d.qfrc_inverse.numpy()[0], qfrc_inverse, "qfrc_inverse")
+ _assert_eq(d.qacc.numpy()[0], qacc, "qacc")
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py
new file mode 100644
index 00000000..f2f2faa5
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py
@@ -0,0 +1,1598 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+from typing import Optional, Tuple
+
+import mujoco
+import numpy as np
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import conditional_graph_supported
+
+from mujoco.mjx.third_party.mujoco_warp._src import math
+from mujoco.mjx.third_party.mujoco_warp._src import types
+
+# number of max iterations to run GJK/EPA
+MJ_CCD_ITERATIONS = 12
+
+
+def _hfield_geom_pair(mjm: mujoco.MjModel) -> Tuple[int, np.array]:
+ geom1, geom2 = np.triu_indices(mjm.ngeom, k=1)
+ geom_type_hf = mujoco.mjtGeom.mjGEOM_HFIELD
+ has_hfield = (mjm.geom_type[geom1] == geom_type_hf) | (mjm.geom_type[geom2] == geom_type_hf)
+ nhfieldgeompair = np.sum(has_hfield)
+ geompair2hfgeompair = -1 * np.ones(mjm.ngeom * (mjm.ngeom - 1) // 2, dtype=int)
+ geompair2hfgeompair[has_hfield] = np.arange(nhfieldgeompair)
+
+ return nhfieldgeompair, geompair2hfgeompair
+
+
+def put_model(mjm: mujoco.MjModel) -> types.Model:
+ """
+ Creates a model on device.
+
+ Args:
+ mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host).
+
+ Returns:
+ Model: The model containing kinematic and dynamic information (device).
+ """
+ # check supported features
+ for field, field_types, field_str in (
+ (mjm.actuator_trntype, types.TrnType, "Actuator transmission type"),
+ (mjm.actuator_dyntype, types.DynType, "Actuator dynamics type"),
+ (mjm.actuator_gaintype, types.GainType, "Gain type"),
+ (mjm.actuator_biastype, types.BiasType, "Bias type"),
+ (mjm.eq_type, types.EqType, "Equality constraint types"),
+ (mjm.geom_type, types.GeomType, "Geom type"),
+ (mjm.sensor_type, types.SensorType, "Sensor types"),
+ (mjm.wrap_type, types.WrapType, "Wrap types"),
+ ):
+ unsupported = ~np.isin(field, list(field_types))
+ if unsupported.any():
+ raise NotImplementedError(f"{field_str} {field[unsupported]} not supported.")
+
+ plugin_id = []
+ plugin_attr = []
+ geom_plugin_index = np.full_like(mjm.geom_type, -1)
+
+ if mjm.nplugin > 0:
+ for i in range(len(mjm.geom_plugin)):
+ if mjm.geom_plugin[i] != -1:
+ p = mjm.geom_plugin[i]
+ geom_plugin_index[i] = len(plugin_id)
+ plugin_id.append(mjm.plugin[p])
+ start = mjm.plugin_attradr[p]
+ end = mjm.plugin_attradr[p + 1] if p + 1 < mjm.nplugin else len(mjm.plugin_attr)
+ values = mjm.plugin_attr[start:end]
+ attr_values = []
+ current = []
+ for v in values:
+ if v == 0:
+ if current:
+ s = "".join(chr(int(x)) for x in current)
+ attr_values.append(float(s))
+ current = []
+ else:
+ current.append(v)
+ # Pad with zeros if less than 3
+ attr_values += [0.0] * (3 - len(attr_values))
+ plugin_attr.append(attr_values[:3])
+
+ plugin_id = np.array(plugin_id)
+ plugin_attr = np.array(plugin_attr)
+
+ if mjm.nflex > 1:
+ raise NotImplementedError("Only one flex is unsupported.")
+
+ if ((mjm.flex_contype != 0) | (mjm.flex_conaffinity != 0)).any():
+ raise NotImplementedError("Flex collisions are not implemented.")
+
+ if mjm.geom_fluid.any():
+ raise NotImplementedError("Ellipsoid fluid model not implemented.")
+
+ # check options
+ for opt, opt_types, msg in (
+ (mjm.opt.integrator, types.IntegratorType, "Integrator"),
+ (mjm.opt.cone, types.ConeType, "Cone"),
+ (mjm.opt.solver, types.SolverType, "Solver"),
+ ):
+ if opt not in set(opt_types):
+ raise NotImplementedError(f"{msg} {opt} is unsupported.")
+
+ if mjm.opt.noslip_iterations > 0:
+ raise NotImplementedError(f"noslip solver not implemented.")
+
+ # TODO(team): remove after _update_gradient for Newton uses tile operations for islands
+ nv_max = 60
+ if mjm.nv > nv_max and mjm.opt.jacobian == mujoco.mjtJacobian.mjJAC_DENSE:
+ raise ValueError(f"Dense is unsupported for nv > {nv_max} (nv = {mjm.nv}).")
+
+ is_sparse = mujoco.mj_isSparse(mjm)
+
+ # calculate some fields that cannot be easily computed inline
+ nlsp = mjm.opt.ls_iterations # TODO(team): how to set nlsp?
+
+ # unfortunately we must create Data in order to get some model fields like M_rownnz
+ mjd = mujoco.MjData(mjm)
+
+ # dof lower triangle row and column indices (used in solver)
+ dof_tri_row, dof_tri_col = np.tril_indices(mjm.nv)
+
+ # indices for sparse qM_fullm (used in solver)
+ qM_fullm_i, qM_fullm_j = [], []
+ for i in range(mjm.nv):
+ j = i
+ while j > -1:
+ qM_fullm_i.append(i)
+ qM_fullm_j.append(j)
+ j = mjm.dof_parentid[j]
+
+ # indices for sparse qM mul_m (used in support)
+ qM_mulm_i, qM_mulm_j, qM_madr_ij = [], [], []
+ for i in range(mjm.nv):
+ madr_ij, j = mjm.dof_Madr[i], i
+
+ while True:
+ madr_ij, j = madr_ij + 1, mjm.dof_parentid[j]
+ if j == -1:
+ break
+ qM_mulm_i.append(i)
+ qM_mulm_j.append(j)
+ qM_madr_ij.append(madr_ij)
+
+ # body_tree is a list of body ids grouped by tree level
+ bodies, body_depth = {}, np.zeros(mjm.nbody, dtype=int) - 1
+ for i in range(mjm.nbody):
+ body_depth[i] = body_depth[mjm.body_parentid[i]] + 1
+ bodies.setdefault(body_depth[i], []).append(i)
+ body_tree = tuple(wp.array(bodies[i], dtype=int) for i in sorted(bodies))
+
+ # qLD_updates has dof tree ordering of qLD updates for sparse factor m
+ qLD_updates, dof_depth = {}, np.zeros(mjm.nv, dtype=int) - 1
+
+ for k in range(mjm.nv):
+ # skip diagonal rows
+ if mjd.M_rownnz[k] == 1:
+ continue
+ dof_depth[k] = dof_depth[mjm.dof_parentid[k]] + 1
+ i = mjm.dof_parentid[k]
+ diag_k = mjd.M_rowadr[k] + mjd.M_rownnz[k] - 1
+ Madr_ki = diag_k - 1
+ while i > -1:
+ qLD_updates.setdefault(dof_depth[i], []).append((i, k, Madr_ki))
+ i = mjm.dof_parentid[i]
+ Madr_ki -= 1
+
+ qLD_updates = tuple(wp.array(qLD_updates[i], dtype=wp.vec3i) for i in sorted(qLD_updates))
+
+ # qM_tiles records the block diagonal structure of qM
+ tile_corners = [i for i in range(mjm.nv) if mjm.dof_parentid[i] == -1]
+ tiles = {}
+ for i in range(len(tile_corners)):
+ tile_beg = tile_corners[i]
+ tile_end = mjm.nv if i == len(tile_corners) - 1 else tile_corners[i + 1]
+ tiles.setdefault(tile_end - tile_beg, []).append(tile_beg)
+
+ qM_tiles = tuple(types.TileSet(adr=wp.array(tiles[sz], dtype=int), size=sz) for sz in sorted(tiles.keys()))
+
+ # subtree_mass is a precalculated array used in smooth
+ subtree_mass = np.copy(mjm.body_mass)
+ # TODO(team): should this be [mjm.nbody - 1, 0) ?
+ for i in range(mjm.nbody - 1, -1, -1):
+ subtree_mass[mjm.body_parentid[i]] += subtree_mass[i]
+
+ # actuator_moment tiles are grouped by dof size and number of actuators
+ tree_id = np.arange(len(tile_corners), dtype=np.int32)
+ num_trees = int(np.max(tree_id)) if len(tree_id) > 0 else 0
+ bodyid = []
+ for i in range(mjm.nu):
+ trntype = mjm.actuator_trntype[i]
+ if trntype == mujoco.mjtTrn.mjTRN_JOINT or trntype == mujoco.mjtTrn.mjTRN_JOINTINPARENT:
+ jntid = mjm.actuator_trnid[i, 0]
+ bodyid.append(mjm.jnt_bodyid[jntid])
+ elif trntype == mujoco.mjtTrn.mjTRN_TENDON:
+ tenid = mjm.actuator_trnid[i, 0]
+ adr = mjm.tendon_adr[tenid]
+ if mjm.wrap_type[adr] == mujoco.mjtWrap.mjWRAP_JOINT:
+ ten_num = mjm.tendon_num[tenid]
+ for i in range(ten_num):
+ bodyid.append(mjm.jnt_bodyid[mjm.wrap_objid[adr + i]])
+ else:
+ for i in range(mjm.nv):
+ bodyid.append(mjm.dof_bodyid[i])
+ elif trntype == mujoco.mjtTrn.mjTRN_BODY:
+ pass
+ elif trntype == mujoco.mjtTrn.mjTRN_SITE:
+ siteid = mjm.actuator_trnid[i, 0]
+ bid = mjm.site_bodyid[siteid]
+ while bid > 0:
+ bodyid.append(bid)
+ bid = mjm.body_parentid[bid]
+ elif trntype == mujoco.mjtTrn.mjTRN_SLIDERCRANK:
+ for i in range(mjm.nv):
+ bodyid.append(mjm.dof_bodyid[i])
+ else:
+ raise NotImplementedError(f"Transmission type {trntype} not implemented.")
+ tree = mjm.body_treeid[np.array(bodyid, dtype=int)]
+ counts, ids = np.histogram(tree, bins=np.arange(0, num_trees + 2))
+ acts_per_tree = dict(zip(ids, counts))
+
+ tiles = {}
+ act_beg = 0
+ for i in range(len(tile_corners)):
+ tile_beg = tile_corners[i]
+ tile_end = mjm.nv if i == len(tile_corners) - 1 else tile_corners[i + 1]
+ tree = int(tree_id[i])
+ act_num = acts_per_tree[tree]
+ tiles.setdefault((tile_end - tile_beg, act_num), []).append((tile_beg, act_beg))
+ act_beg += act_num
+
+ actuator_moment_tiles_nv, actuator_moment_tiles_nu = tuple(), tuple()
+
+ for (nv, nu), adr in sorted(tiles.items()):
+ adr_nv = wp.array([nv for nv, _ in adr], dtype=int)
+ adr_nu = wp.array([nu for _, nu in adr], dtype=int)
+ actuator_moment_tiles_nv += (types.TileSet(adr=adr_nv, size=nv),)
+ actuator_moment_tiles_nu += (types.TileSet(adr=adr_nu, size=nu),)
+
+ # fixed tendon
+ tendon_jnt_adr = []
+ wrap_jnt_adr = []
+ for i in range(mjm.ntendon):
+ adr = mjm.tendon_adr[i]
+ if mjm.wrap_type[adr] == mujoco.mjtWrap.mjWRAP_JOINT:
+ tendon_num = mjm.tendon_num[i]
+ for j in range(tendon_num):
+ tendon_jnt_adr.append(i)
+ wrap_jnt_adr.append(adr + j)
+
+ # spatial tendon
+ tendon_site_pair_adr = []
+ tendon_geom_adr = []
+
+ ten_wrapadr_site = [0]
+ ten_wrapnum_site = []
+ for i, tendon_num in enumerate(mjm.tendon_num):
+ adr = mjm.tendon_adr[i]
+ # sites
+ if (mjm.wrap_type[adr : adr + tendon_num] == mujoco.mjtWrap.mjWRAP_SITE).all():
+ if i < mjm.ntendon:
+ ten_wrapadr_site.append(ten_wrapadr_site[-1] + tendon_num)
+ ten_wrapnum_site.append(tendon_num)
+ else:
+ if i < mjm.ntendon:
+ ten_wrapadr_site.append(ten_wrapadr_site[-1])
+ ten_wrapnum_site.append(0)
+
+ # geoms
+ for j in range(tendon_num):
+ wrap_type = mjm.wrap_type[adr + j]
+ if j < tendon_num - 1:
+ next_wrap_type = mjm.wrap_type[adr + j + 1]
+ if wrap_type == mujoco.mjtWrap.mjWRAP_SITE and next_wrap_type == mujoco.mjtWrap.mjWRAP_SITE:
+ tendon_site_pair_adr.append(i)
+ if wrap_type == mujoco.mjtWrap.mjWRAP_SPHERE or wrap_type == mujoco.mjtWrap.mjWRAP_CYLINDER:
+ tendon_geom_adr.append(i)
+
+ wrap_site_adr = np.nonzero(mjm.wrap_type == mujoco.mjtWrap.mjWRAP_SITE)[0]
+ wrap_site_pair_adr = np.setdiff1d(wrap_site_adr[np.nonzero(np.diff(wrap_site_adr) == 1)[0]], mjm.tendon_adr[1:] - 1)
+ wrap_geom_adr = np.nonzero(np.isin(mjm.wrap_type, [mujoco.mjtWrap.mjWRAP_SPHERE, mujoco.mjtWrap.mjWRAP_CYLINDER]))[0]
+
+ # pulley scaling
+ wrap_pulley_scale = np.ones(mjm.nwrap, dtype=float)
+ pulley_adr = np.nonzero(mjm.wrap_type == mujoco.mjtWrap.mjWRAP_PULLEY)[0]
+ for tadr, tnum in zip(mjm.tendon_adr, mjm.tendon_num):
+ for padr in pulley_adr:
+ if tadr <= padr < tadr + tnum:
+ wrap_pulley_scale[padr : tadr + tnum] = 1.0 / mjm.wrap_prm[padr]
+
+ # mocap
+ mocap_bodyid = np.arange(mjm.nbody)[mjm.body_mocapid >= 0]
+ mocap_bodyid = mocap_bodyid[mjm.body_mocapid[mjm.body_mocapid >= 0].argsort()]
+
+ # precalculated geom pairs
+ filterparent = not (mjm.opt.disableflags & types.DisableBit.FILTERPARENT.value)
+
+ geom1, geom2 = np.triu_indices(mjm.ngeom, k=1)
+ nxn_geom_pair = np.stack((geom1, geom2), axis=1)
+
+ bodyid1 = mjm.geom_bodyid[geom1]
+ bodyid2 = mjm.geom_bodyid[geom2]
+ contype1 = mjm.geom_contype[geom1]
+ contype2 = mjm.geom_contype[geom2]
+ conaffinity1 = mjm.geom_conaffinity[geom1]
+ conaffinity2 = mjm.geom_conaffinity[geom2]
+ weldid1 = mjm.body_weldid[bodyid1]
+ weldid2 = mjm.body_weldid[bodyid2]
+ weld_parentid1 = mjm.body_weldid[mjm.body_parentid[weldid1]]
+ weld_parentid2 = mjm.body_weldid[mjm.body_parentid[weldid2]]
+
+ self_collision = weldid1 == weldid2
+ parent_child_collision = (
+ filterparent & (weldid1 != 0) & (weldid2 != 0) & ((weldid1 == weld_parentid2) | (weldid2 == weld_parentid1))
+ )
+ mask = np.array((contype1 & conaffinity2) | (contype2 & conaffinity1), dtype=bool)
+ exclude = np.isin((bodyid1 << 16) + bodyid2, mjm.exclude_signature)
+
+ nxn_pairid = -1 * np.ones(len(geom1), dtype=int)
+ nxn_pairid[~(mask & ~self_collision & ~parent_child_collision & ~exclude)] = -2
+
+ # contact pairs
+ for i in range(mjm.npair):
+ pair_geom1 = mjm.pair_geom1[i]
+ pair_geom2 = mjm.pair_geom2[i]
+
+ if pair_geom2 < pair_geom1:
+ pairid = np.int32(math.upper_tri_index(mjm.ngeom, int(pair_geom2), int(pair_geom1)))
+ else:
+ pairid = np.int32(math.upper_tri_index(mjm.ngeom, int(pair_geom1), int(pair_geom2)))
+
+ nxn_pairid[pairid] = i
+
+ include = nxn_pairid > -2
+ nxn_pairid_filtered = nxn_pairid[include]
+ nxn_geom_pair_filtered = nxn_geom_pair[include]
+
+ # count contact pair types
+ geom_type_pair_count = np.bincount(
+ [
+ math.upper_trid_index(len(types.GeomType), int(mjm.geom_type[geom1[i]]), int(mjm.geom_type[geom2[i]]))
+ for i in np.arange(len(geom1))
+ if nxn_pairid[i] > -2
+ ],
+ minlength=len(types.GeomType) * (len(types.GeomType) + 1) // 2,
+ )
+
+ # Disable collisions if there are no potentially colliding pairs
+ if np.sum(geom_type_pair_count) == 0:
+ mjm.opt.disableflags |= types.DisableBit.CONTACT.value
+
+ def create_nmodel_batched_array(mjm_array, dtype, expand_dim=True):
+ array = wp.array(mjm_array, dtype=dtype)
+ # add private attribute for JAX to determine which fields are batched
+ array._is_batched = True
+ if not expand_dim:
+ array.strides = (0,) + array.strides[1:]
+ return array
+ array.strides = (0,) + array.strides
+ array.ndim += 1
+ array.shape = (1,) + array.shape
+ return array
+
+ # rangefinder
+ is_rangefinder = mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER
+ sensor_rangefinder_adr = np.nonzero(is_rangefinder)[0]
+ rangefinder_sensor_adr = np.full(mjm.nsensor, -1)
+ rangefinder_sensor_adr[sensor_rangefinder_adr] = np.arange(len(sensor_rangefinder_adr))
+
+ # TODO(team): improve heuristic for selecting broadphase routine
+ if mjm.ngeom > 1000:
+ broadphase = types.BroadphaseType.SAP_SEGMENTED
+ elif mjm.ngeom > 100:
+ broadphase = types.BroadphaseType.SAP_TILE
+ else:
+ broadphase = types.BroadphaseType.NXN
+
+ condim = np.concatenate((mjm.geom_condim, mjm.pair_dim))
+ condim_max = np.max(condim) if len(condim) > 0 else 0
+
+ m = types.Model(
+ nq=mjm.nq,
+ nv=mjm.nv,
+ nu=mjm.nu,
+ na=mjm.na,
+ nbody=mjm.nbody,
+ njnt=mjm.njnt,
+ ngeom=mjm.ngeom,
+ nsite=mjm.nsite,
+ ncam=mjm.ncam,
+ nlight=mjm.nlight,
+ nflex=mjm.nflex,
+ nflexvert=mjm.nflexvert,
+ nflexedge=mjm.nflexedge,
+ nflexelem=mjm.nflexelem,
+ nflexelemdata=mjm.nflexelemdata,
+ nexclude=mjm.nexclude,
+ neq=mjm.neq,
+ nmocap=mjm.nmocap,
+ ngravcomp=mjm.ngravcomp,
+ nM=mjm.nM,
+ nC=mjm.nC,
+ ntendon=mjm.ntendon,
+ nwrap=mjm.nwrap,
+ nsensor=mjm.nsensor,
+ nsensordata=mjm.nsensordata,
+ nmeshvert=mjm.nmeshvert,
+ nmeshface=mjm.nmeshface,
+ nmeshgraph=mjm.nmeshgraph,
+ nmeshpoly=mjm.nmeshpoly,
+ nmeshpolyvert=mjm.nmeshpolyvert,
+ nmeshpolymap=mjm.nmeshpolymap,
+ nlsp=nlsp,
+ npair=mjm.npair,
+ opt=types.Option(
+ timestep=create_nmodel_batched_array(np.array(mjm.opt.timestep), dtype=float, expand_dim=False),
+ tolerance=create_nmodel_batched_array(np.array(mjm.opt.tolerance), dtype=float, expand_dim=False),
+ ls_tolerance=create_nmodel_batched_array(np.array(mjm.opt.ls_tolerance), dtype=float, expand_dim=False),
+ gravity=create_nmodel_batched_array(mjm.opt.gravity, dtype=wp.vec3, expand_dim=False),
+ magnetic=create_nmodel_batched_array(mjm.opt.magnetic, dtype=wp.vec3, expand_dim=False),
+ wind=create_nmodel_batched_array(mjm.opt.wind, dtype=wp.vec3, expand_dim=False),
+ has_fluid=bool(mjm.opt.wind.any() or mjm.opt.density or mjm.opt.viscosity),
+ density=create_nmodel_batched_array(np.array(mjm.opt.density), dtype=float, expand_dim=False),
+ viscosity=create_nmodel_batched_array(np.array(mjm.opt.viscosity), dtype=float, expand_dim=False),
+ cone=mjm.opt.cone,
+ solver=mjm.opt.solver,
+ iterations=mjm.opt.iterations,
+ ls_iterations=mjm.opt.ls_iterations,
+ integrator=mjm.opt.integrator,
+ disableflags=mjm.opt.disableflags,
+ enableflags=mjm.opt.enableflags,
+ impratio=create_nmodel_batched_array(np.array(mjm.opt.impratio), dtype=float, expand_dim=False),
+ is_sparse=bool(is_sparse),
+ ls_parallel=False,
+ gjk_iterations=MJ_CCD_ITERATIONS,
+ epa_iterations=MJ_CCD_ITERATIONS,
+ broadphase=int(broadphase),
+ broadphase_filter=int(
+ types.BroadphaseFilter.PLANE.value | types.BroadphaseFilter.SPHERE.value | types.BroadphaseFilter.OBB.value
+ ),
+ graph_conditional=True and conditional_graph_supported(),
+ sdf_initpoints=mjm.opt.sdf_initpoints,
+ sdf_iterations=mjm.opt.sdf_iterations,
+ run_collision_detection=True,
+ ),
+ stat=types.Statistic(
+ meaninertia=mjm.stat.meaninertia,
+ ),
+ qpos0=create_nmodel_batched_array(mjm.qpos0, dtype=float),
+ qpos_spring=create_nmodel_batched_array(mjm.qpos_spring, dtype=float),
+ qM_fullm_i=wp.array(qM_fullm_i, dtype=int),
+ qM_fullm_j=wp.array(qM_fullm_j, dtype=int),
+ qM_mulm_i=wp.array(qM_mulm_i, dtype=int),
+ qM_mulm_j=wp.array(qM_mulm_j, dtype=int),
+ qM_madr_ij=wp.array(qM_madr_ij, dtype=int),
+ qLD_updates=qLD_updates,
+ M_rownnz=wp.array(mjd.M_rownnz, dtype=int),
+ M_rowadr=wp.array(mjd.M_rowadr, dtype=int),
+ M_colind=wp.array(mjd.M_colind, dtype=int),
+ mapM2M=wp.array(mjd.mapM2M, dtype=int),
+ qM_tiles=qM_tiles,
+ body_tree=body_tree,
+ body_parentid=wp.array(mjm.body_parentid, dtype=int),
+ body_rootid=wp.array(mjm.body_rootid, dtype=int),
+ body_weldid=wp.array(mjm.body_weldid, dtype=int),
+ body_mocapid=wp.array(mjm.body_mocapid, dtype=int),
+ mocap_bodyid=wp.array(mocap_bodyid, dtype=int),
+ body_jntnum=wp.array(mjm.body_jntnum, dtype=int),
+ body_jntadr=wp.array(mjm.body_jntadr, dtype=int),
+ body_dofnum=wp.array(mjm.body_dofnum, dtype=int),
+ body_dofadr=wp.array(mjm.body_dofadr, dtype=int),
+ body_geomnum=wp.array(mjm.body_geomnum, dtype=int),
+ body_geomadr=wp.array(mjm.body_geomadr, dtype=int),
+ body_pos=create_nmodel_batched_array(mjm.body_pos, dtype=wp.vec3),
+ body_quat=create_nmodel_batched_array(mjm.body_quat, dtype=wp.quat),
+ body_ipos=create_nmodel_batched_array(mjm.body_ipos, dtype=wp.vec3),
+ body_iquat=create_nmodel_batched_array(mjm.body_iquat, dtype=wp.quat),
+ body_mass=create_nmodel_batched_array(mjm.body_mass, dtype=float),
+ body_subtreemass=create_nmodel_batched_array(mjm.body_subtreemass, dtype=float),
+ subtree_mass=create_nmodel_batched_array(subtree_mass, dtype=float),
+ body_inertia=create_nmodel_batched_array(mjm.body_inertia, dtype=wp.vec3),
+ body_invweight0=create_nmodel_batched_array(mjm.body_invweight0, dtype=wp.vec2),
+ body_contype=wp.array(mjm.body_contype, dtype=int),
+ body_conaffinity=wp.array(mjm.body_conaffinity, dtype=int),
+ body_gravcomp=create_nmodel_batched_array(mjm.body_gravcomp, dtype=float),
+ jnt_type=wp.array(mjm.jnt_type, dtype=int),
+ jnt_qposadr=wp.array(mjm.jnt_qposadr, dtype=int),
+ jnt_dofadr=wp.array(mjm.jnt_dofadr, dtype=int),
+ jnt_bodyid=wp.array(mjm.jnt_bodyid, dtype=int),
+ jnt_limited=wp.array(mjm.jnt_limited, dtype=int),
+ jnt_actfrclimited=wp.array(mjm.jnt_actfrclimited, dtype=bool),
+ jnt_solref=create_nmodel_batched_array(mjm.jnt_solref, dtype=wp.vec2),
+ jnt_solimp=create_nmodel_batched_array(mjm.jnt_solimp, dtype=types.vec5),
+ jnt_pos=create_nmodel_batched_array(mjm.jnt_pos, dtype=wp.vec3),
+ jnt_axis=create_nmodel_batched_array(mjm.jnt_axis, dtype=wp.vec3),
+ jnt_stiffness=create_nmodel_batched_array(mjm.jnt_stiffness, dtype=float),
+ jnt_range=create_nmodel_batched_array(mjm.jnt_range, dtype=wp.vec2),
+ jnt_actfrcrange=create_nmodel_batched_array(mjm.jnt_actfrcrange, dtype=wp.vec2),
+ jnt_margin=create_nmodel_batched_array(mjm.jnt_margin, dtype=float),
+ # these jnt_limited adrs are used in constraint.py
+ jnt_limited_slide_hinge_adr=wp.array(
+ np.nonzero(
+ mjm.jnt_limited & ((mjm.jnt_type == mujoco.mjtJoint.mjJNT_SLIDE) | (mjm.jnt_type == mujoco.mjtJoint.mjJNT_HINGE))
+ )[0],
+ dtype=int,
+ ),
+ jnt_limited_ball_adr=wp.array(
+ np.nonzero(mjm.jnt_limited & (mjm.jnt_type == mujoco.mjtJoint.mjJNT_BALL))[0],
+ dtype=int,
+ ),
+ jnt_actgravcomp=wp.array(mjm.jnt_actgravcomp, dtype=int),
+ dof_bodyid=wp.array(mjm.dof_bodyid, dtype=int),
+ dof_jntid=wp.array(mjm.dof_jntid, dtype=int),
+ dof_parentid=wp.array(mjm.dof_parentid, dtype=int),
+ dof_Madr=wp.array(mjm.dof_Madr, dtype=int),
+ dof_armature=create_nmodel_batched_array(mjm.dof_armature, dtype=float),
+ dof_damping=create_nmodel_batched_array(mjm.dof_damping, dtype=float),
+ dof_invweight0=create_nmodel_batched_array(mjm.dof_invweight0, dtype=float),
+ dof_frictionloss=create_nmodel_batched_array(mjm.dof_frictionloss, dtype=float),
+ dof_solimp=create_nmodel_batched_array(mjm.dof_solimp, dtype=types.vec5),
+ dof_solref=create_nmodel_batched_array(mjm.dof_solref, dtype=wp.vec2),
+ dof_tri_row=wp.array(dof_tri_row, dtype=int),
+ dof_tri_col=wp.array(dof_tri_col, dtype=int),
+ geom_type=wp.array(mjm.geom_type, dtype=int),
+ geom_contype=wp.array(mjm.geom_contype, dtype=int),
+ geom_conaffinity=wp.array(mjm.geom_conaffinity, dtype=int),
+ geom_condim=wp.array(mjm.geom_condim, dtype=int),
+ geom_bodyid=wp.array(mjm.geom_bodyid, dtype=int),
+ geom_dataid=wp.array(mjm.geom_dataid, dtype=int),
+ geom_group=wp.array(mjm.geom_group, dtype=int),
+ geom_matid=create_nmodel_batched_array(mjm.geom_matid, dtype=int),
+ geom_priority=wp.array(mjm.geom_priority, dtype=int),
+ geom_solmix=create_nmodel_batched_array(mjm.geom_solmix, dtype=float),
+ geom_solref=create_nmodel_batched_array(mjm.geom_solref, dtype=wp.vec2),
+ geom_solimp=create_nmodel_batched_array(mjm.geom_solimp, dtype=types.vec5),
+ geom_size=create_nmodel_batched_array(mjm.geom_size, dtype=wp.vec3),
+ geom_aabb=wp.array2d(mjm.geom_aabb, dtype=wp.vec3),
+ geom_rbound=create_nmodel_batched_array(mjm.geom_rbound, dtype=float),
+ geom_pos=create_nmodel_batched_array(mjm.geom_pos, dtype=wp.vec3),
+ geom_quat=create_nmodel_batched_array(mjm.geom_quat, dtype=wp.quat),
+ geom_friction=create_nmodel_batched_array(mjm.geom_friction, dtype=wp.vec3),
+ geom_margin=create_nmodel_batched_array(mjm.geom_margin, dtype=float),
+ geom_gap=create_nmodel_batched_array(mjm.geom_gap, dtype=float),
+ geom_rgba=create_nmodel_batched_array(mjm.geom_rgba, dtype=wp.vec4),
+ site_type=wp.array(mjm.site_type, dtype=int),
+ site_bodyid=wp.array(mjm.site_bodyid, dtype=int),
+ site_size=wp.array(mjm.site_size, dtype=wp.vec3),
+ site_pos=create_nmodel_batched_array(mjm.site_pos, dtype=wp.vec3),
+ site_quat=create_nmodel_batched_array(mjm.site_quat, dtype=wp.quat),
+ cam_mode=wp.array(mjm.cam_mode, dtype=int),
+ cam_bodyid=wp.array(mjm.cam_bodyid, dtype=int),
+ cam_targetbodyid=wp.array(mjm.cam_targetbodyid, dtype=int),
+ cam_pos=create_nmodel_batched_array(mjm.cam_pos, dtype=wp.vec3),
+ cam_quat=create_nmodel_batched_array(mjm.cam_quat, dtype=wp.quat),
+ cam_poscom0=create_nmodel_batched_array(mjm.cam_poscom0, dtype=wp.vec3),
+ cam_pos0=create_nmodel_batched_array(mjm.cam_pos0, dtype=wp.vec3),
+ cam_mat0=create_nmodel_batched_array(mjm.cam_mat0, dtype=wp.mat33),
+ cam_fovy=wp.array(mjm.cam_fovy, dtype=float),
+ cam_resolution=wp.array(mjm.cam_resolution, dtype=wp.vec2i),
+ cam_sensorsize=wp.array(mjm.cam_sensorsize, dtype=wp.vec2),
+ cam_intrinsic=wp.array(mjm.cam_intrinsic, dtype=wp.vec4),
+ light_mode=wp.array(mjm.light_mode, dtype=int),
+ light_bodyid=wp.array(mjm.light_bodyid, dtype=int),
+ light_targetbodyid=wp.array(mjm.light_targetbodyid, dtype=int),
+ light_pos=create_nmodel_batched_array(mjm.light_pos, dtype=wp.vec3),
+ light_dir=create_nmodel_batched_array(mjm.light_dir, dtype=wp.vec3),
+ light_poscom0=create_nmodel_batched_array(mjm.light_poscom0, dtype=wp.vec3),
+ light_pos0=create_nmodel_batched_array(mjm.light_pos0, dtype=wp.vec3),
+ light_dir0=create_nmodel_batched_array(mjm.light_dir0, dtype=wp.vec3),
+ flex_dim=wp.array(mjm.flex_dim, dtype=int),
+ flex_vertadr=wp.array(mjm.flex_vertadr, dtype=int),
+ flex_vertnum=wp.array(mjm.flex_vertnum, dtype=int),
+ flex_edgeadr=wp.array(mjm.flex_edgeadr, dtype=int),
+ flex_elemedgeadr=wp.array(mjm.flex_elemedgeadr, dtype=int),
+ flex_vertbodyid=wp.array(mjm.flex_vertbodyid, dtype=int),
+ flex_edge=wp.array(mjm.flex_edge, dtype=wp.vec2i),
+ flex_edgeflap=wp.array(mjm.flex_edgeflap, dtype=wp.vec2i),
+ flex_elem=wp.array(mjm.flex_elem, dtype=int),
+ flex_elemedge=wp.array(mjm.flex_elemedge, dtype=int),
+ flexedge_length0=wp.array(mjm.flexedge_length0, dtype=float),
+ flex_stiffness=wp.array(mjm.flex_stiffness.flatten(), dtype=float),
+ flex_bending=wp.array(mjm.flex_bending, dtype=wp.mat44f),
+ flex_damping=wp.array(mjm.flex_damping, dtype=float),
+ mesh_vertadr=wp.array(mjm.mesh_vertadr, dtype=int),
+ mesh_vertnum=wp.array(mjm.mesh_vertnum, dtype=int),
+ mesh_vert=wp.array(mjm.mesh_vert, dtype=wp.vec3),
+ mesh_faceadr=wp.array(mjm.mesh_faceadr, dtype=int),
+ mesh_face=wp.array(mjm.mesh_face, dtype=wp.vec3i),
+ mesh_graphadr=wp.array(mjm.mesh_graphadr, dtype=int),
+ mesh_graph=wp.array(mjm.mesh_graph, dtype=int),
+ mesh_polynum=wp.array(mjm.mesh_polynum, dtype=int),
+ mesh_polyadr=wp.array(mjm.mesh_polyadr, dtype=int),
+ mesh_polynormal=wp.array(mjm.mesh_polynormal, dtype=wp.vec3),
+ mesh_polyvertadr=wp.array(mjm.mesh_polyvertadr, dtype=int),
+ mesh_polyvertnum=wp.array(mjm.mesh_polyvertnum, dtype=int),
+ mesh_polyvert=wp.array(mjm.mesh_polyvert, dtype=int),
+ mesh_polymapadr=wp.array(mjm.mesh_polymapadr, dtype=int),
+ mesh_polymapnum=wp.array(mjm.mesh_polymapnum, dtype=int),
+ mesh_polymap=wp.array(mjm.mesh_polymap, dtype=int),
+ nhfield=mjm.nhfield,
+ nhfielddata=mjm.nhfielddata,
+ hfield_adr=wp.array(mjm.hfield_adr, dtype=int),
+ hfield_nrow=wp.array(mjm.hfield_nrow, dtype=int),
+ hfield_ncol=wp.array(mjm.hfield_ncol, dtype=int),
+ hfield_size=wp.array(mjm.hfield_size, dtype=wp.vec4),
+ hfield_data=wp.array(mjm.hfield_data, dtype=float),
+ eq_type=wp.array(mjm.eq_type, dtype=int),
+ eq_obj1id=wp.array(mjm.eq_obj1id, dtype=int),
+ eq_obj2id=wp.array(mjm.eq_obj2id, dtype=int),
+ eq_objtype=wp.array(mjm.eq_objtype, dtype=int),
+ eq_active0=wp.array(mjm.eq_active0, dtype=bool),
+ eq_solref=create_nmodel_batched_array(mjm.eq_solref, dtype=wp.vec2),
+ eq_solimp=create_nmodel_batched_array(mjm.eq_solimp, dtype=types.vec5),
+ eq_data=create_nmodel_batched_array(mjm.eq_data, dtype=types.vec11),
+ # pre-compute indices of equality constraints
+ eq_connect_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.CONNECT.value)[0], dtype=int),
+ eq_wld_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.WELD.value)[0], dtype=int),
+ eq_jnt_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.JOINT.value)[0], dtype=int),
+ eq_ten_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.TENDON.value)[0], dtype=int),
+ actuator_moment_tiles_nv=actuator_moment_tiles_nv,
+ actuator_moment_tiles_nu=actuator_moment_tiles_nu,
+ actuator_trntype=wp.array(mjm.actuator_trntype, dtype=int),
+ actuator_dyntype=wp.array(mjm.actuator_dyntype, dtype=int),
+ actuator_gaintype=wp.array(mjm.actuator_gaintype, dtype=int),
+ actuator_biastype=wp.array(mjm.actuator_biastype, dtype=int),
+ actuator_trnid=wp.array(mjm.actuator_trnid, dtype=wp.vec2i),
+ actuator_actadr=wp.array(mjm.actuator_actadr, dtype=int),
+ actuator_actnum=wp.array(mjm.actuator_actnum, dtype=int),
+ actuator_ctrllimited=wp.array(mjm.actuator_ctrllimited, dtype=bool),
+ actuator_forcelimited=wp.array(mjm.actuator_forcelimited, dtype=bool),
+ actuator_actlimited=wp.array(mjm.actuator_actlimited, dtype=bool),
+ actuator_dynprm=create_nmodel_batched_array(mjm.actuator_dynprm, dtype=types.vec10f),
+ actuator_gainprm=create_nmodel_batched_array(mjm.actuator_gainprm, dtype=types.vec10f),
+ actuator_biasprm=create_nmodel_batched_array(mjm.actuator_biasprm, dtype=types.vec10f),
+ actuator_actearly=wp.array(mjm.actuator_actearly, dtype=bool),
+ actuator_ctrlrange=create_nmodel_batched_array(mjm.actuator_ctrlrange, dtype=wp.vec2),
+ actuator_forcerange=create_nmodel_batched_array(mjm.actuator_forcerange, dtype=wp.vec2),
+ actuator_actrange=create_nmodel_batched_array(mjm.actuator_actrange, dtype=wp.vec2),
+ actuator_gear=create_nmodel_batched_array(mjm.actuator_gear, dtype=wp.spatial_vector),
+ actuator_cranklength=wp.array(mjm.actuator_cranklength, dtype=float),
+ actuator_acc0=wp.array(mjm.actuator_acc0, dtype=float),
+ actuator_lengthrange=wp.array(mjm.actuator_lengthrange, dtype=wp.vec2),
+ exclude_signature=wp.array(mjm.exclude_signature, dtype=int),
+ # short-circuiting here allows us to skip a lot of code in implicit integration
+ actuator_affine_bias_gain=bool(
+ np.any(mjm.actuator_biastype == types.BiasType.AFFINE.value)
+ or np.any(mjm.actuator_gaintype == types.GainType.AFFINE.value)
+ ),
+ nxn_geom_pair=wp.array(nxn_geom_pair, dtype=wp.vec2i),
+ nxn_geom_pair_filtered=wp.array(nxn_geom_pair_filtered, dtype=wp.vec2i),
+ nxn_pairid=wp.array(nxn_pairid, dtype=int),
+ nxn_pairid_filtered=wp.array(nxn_pairid_filtered, dtype=int),
+ pair_dim=wp.array(mjm.pair_dim, dtype=int),
+ pair_geom1=wp.array(mjm.pair_geom1, dtype=int),
+ pair_geom2=wp.array(mjm.pair_geom2, dtype=int),
+ pair_solref=create_nmodel_batched_array(mjm.pair_solref, dtype=wp.vec2),
+ pair_solreffriction=create_nmodel_batched_array(mjm.pair_solreffriction, dtype=wp.vec2),
+ pair_solimp=create_nmodel_batched_array(mjm.pair_solimp, dtype=types.vec5),
+ pair_margin=create_nmodel_batched_array(mjm.pair_margin, dtype=float),
+ pair_gap=create_nmodel_batched_array(mjm.pair_gap, dtype=float),
+ pair_friction=create_nmodel_batched_array(mjm.pair_friction, dtype=types.vec5),
+ condim_max=condim_max, # TODO(team): get max after filtering,
+ tendon_adr=wp.array(mjm.tendon_adr, dtype=int),
+ tendon_num=wp.array(mjm.tendon_num, dtype=int),
+ tendon_limited=wp.array(mjm.tendon_limited, dtype=int),
+ tendon_limited_adr=wp.array(np.nonzero(mjm.tendon_limited)[0], dtype=int),
+ tendon_actfrclimited=wp.array(mjm.tendon_actfrclimited, dtype=bool),
+ tendon_solref_lim=create_nmodel_batched_array(mjm.tendon_solref_lim, dtype=wp.vec2f),
+ tendon_solimp_lim=create_nmodel_batched_array(mjm.tendon_solimp_lim, dtype=types.vec5),
+ tendon_solref_fri=create_nmodel_batched_array(mjm.tendon_solref_fri, dtype=wp.vec2f),
+ tendon_solimp_fri=create_nmodel_batched_array(mjm.tendon_solimp_fri, dtype=types.vec5),
+ tendon_range=create_nmodel_batched_array(mjm.tendon_range, dtype=wp.vec2f),
+ tendon_actfrcrange=create_nmodel_batched_array(mjm.tendon_actfrcrange, dtype=wp.vec2),
+ tendon_margin=create_nmodel_batched_array(mjm.tendon_margin, dtype=float),
+ tendon_stiffness=create_nmodel_batched_array(mjm.tendon_stiffness, dtype=float),
+ tendon_damping=create_nmodel_batched_array(mjm.tendon_damping, dtype=float),
+ tendon_armature=create_nmodel_batched_array(mjm.tendon_armature, dtype=float),
+ tendon_frictionloss=create_nmodel_batched_array(mjm.tendon_frictionloss, dtype=float),
+ tendon_lengthspring=create_nmodel_batched_array(mjm.tendon_lengthspring, dtype=wp.vec2),
+ tendon_length0=create_nmodel_batched_array(mjm.tendon_length0, dtype=float),
+ tendon_invweight0=create_nmodel_batched_array(mjm.tendon_invweight0, dtype=float),
+ wrap_objid=wp.array(mjm.wrap_objid, dtype=int),
+ wrap_prm=wp.array(mjm.wrap_prm, dtype=float),
+ wrap_type=wp.array(mjm.wrap_type, dtype=int),
+ tendon_jnt_adr=wp.array(tendon_jnt_adr, dtype=int),
+ tendon_site_pair_adr=wp.array(tendon_site_pair_adr, dtype=int),
+ tendon_geom_adr=wp.array(tendon_geom_adr, dtype=int),
+ ten_wrapadr_site=wp.array(ten_wrapadr_site, dtype=int),
+ ten_wrapnum_site=wp.array(ten_wrapnum_site, dtype=int),
+ wrap_jnt_adr=wp.array(wrap_jnt_adr, dtype=int),
+ wrap_site_adr=wp.array(wrap_site_adr, dtype=int),
+ wrap_site_pair_adr=wp.array(wrap_site_pair_adr, dtype=int),
+ wrap_geom_adr=wp.array(wrap_geom_adr, dtype=int),
+ wrap_pulley_scale=wp.array(wrap_pulley_scale, dtype=float),
+ sensor_type=wp.array(mjm.sensor_type, dtype=int),
+ sensor_datatype=wp.array(mjm.sensor_datatype, dtype=int),
+ sensor_objtype=wp.array(mjm.sensor_objtype, dtype=int),
+ sensor_objid=wp.array(mjm.sensor_objid, dtype=int),
+ sensor_reftype=wp.array(mjm.sensor_reftype, dtype=int),
+ sensor_refid=wp.array(mjm.sensor_refid, dtype=int),
+ sensor_dim=wp.array(mjm.sensor_dim, dtype=int),
+ sensor_adr=wp.array(mjm.sensor_adr, dtype=int),
+ sensor_cutoff=wp.array(mjm.sensor_cutoff, dtype=float),
+ sensor_pos_adr=wp.array(
+ np.nonzero(
+ (mjm.sensor_needstage == mujoco.mjtStage.mjSTAGE_POS)
+ & (mjm.sensor_type != mujoco.mjtSensor.mjSENS_JOINTLIMITPOS)
+ & (mjm.sensor_type != mujoco.mjtSensor.mjSENS_TENDONLIMITPOS)
+ )[0],
+ dtype=int,
+ ),
+ sensor_limitpos_adr=wp.array(
+ np.nonzero(
+ (mjm.sensor_type == mujoco.mjtSensor.mjSENS_JOINTLIMITPOS) | (mjm.sensor_type == mujoco.mjtSensor.mjSENS_TENDONLIMITPOS)
+ )[0],
+ dtype=int,
+ ),
+ sensor_vel_adr=wp.array(
+ np.nonzero(
+ (mjm.sensor_needstage == mujoco.mjtStage.mjSTAGE_VEL)
+ & (
+ (mjm.sensor_type != mujoco.mjtSensor.mjSENS_JOINTLIMITVEL)
+ | (mjm.sensor_type != mujoco.mjtSensor.mjSENS_TENDONLIMITVEL)
+ )
+ )[0],
+ dtype=int,
+ ),
+ sensor_limitvel_adr=wp.array(
+ np.nonzero(
+ (mjm.sensor_type == mujoco.mjtSensor.mjSENS_JOINTLIMITVEL) | (mjm.sensor_type == mujoco.mjtSensor.mjSENS_TENDONLIMITVEL)
+ )[0],
+ dtype=int,
+ ),
+ sensor_acc_adr=wp.array(
+ np.nonzero(
+ (mjm.sensor_needstage == mujoco.mjtStage.mjSTAGE_ACC)
+ & (
+ (mjm.sensor_type != mujoco.mjtSensor.mjSENS_TOUCH)
+ | (mjm.sensor_type != mujoco.mjtSensor.mjSENS_JOINTLIMITFRC)
+ | (mjm.sensor_type != mujoco.mjtSensor.mjSENS_TENDONLIMITFRC)
+ | (mjm.sensor_type != mujoco.mjtSensor.mjSENS_TENDONACTFRC)
+ )
+ )[0],
+ dtype=int,
+ ),
+ sensor_rangefinder_adr=wp.array(sensor_rangefinder_adr, dtype=int),
+ rangefinder_sensor_adr=wp.array(rangefinder_sensor_adr, dtype=int),
+ sensor_touch_adr=wp.array(
+ np.nonzero(mjm.sensor_type == mujoco.mjtSensor.mjSENS_TOUCH)[0],
+ dtype=int,
+ ),
+ sensor_limitfrc_adr=wp.array(
+ np.nonzero(
+ (mjm.sensor_type == mujoco.mjtSensor.mjSENS_JOINTLIMITFRC) | (mjm.sensor_type == mujoco.mjtSensor.mjSENS_TENDONLIMITFRC)
+ )[0],
+ dtype=int,
+ ),
+ sensor_e_potential=(mjm.sensor_type == mujoco.mjtSensor.mjSENS_E_POTENTIAL).any(),
+ sensor_e_kinetic=(mjm.sensor_type == mujoco.mjtSensor.mjSENS_E_KINETIC).any(),
+ sensor_tendonactfrc_adr=wp.array(
+ np.nonzero(mjm.sensor_type == mujoco.mjtSensor.mjSENS_TENDONACTFRC)[0],
+ dtype=int,
+ ),
+ sensor_subtree_vel=np.isin(
+ mjm.sensor_type,
+ [mujoco.mjtSensor.mjSENS_SUBTREELINVEL, mujoco.mjtSensor.mjSENS_SUBTREEANGMOM],
+ ).any(),
+ sensor_rne_postconstraint=np.isin(
+ mjm.sensor_type,
+ [
+ mujoco.mjtSensor.mjSENS_ACCELEROMETER,
+ mujoco.mjtSensor.mjSENS_FORCE,
+ mujoco.mjtSensor.mjSENS_TORQUE,
+ mujoco.mjtSensor.mjSENS_FRAMELINACC,
+ mujoco.mjtSensor.mjSENS_FRAMEANGACC,
+ ],
+ ).any(),
+ sensor_rangefinder_bodyid=wp.array(
+ mjm.site_bodyid[mjm.sensor_objid[mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER]], dtype=int
+ ),
+ plugin=wp.array(plugin_id, dtype=int),
+ plugin_attr=wp.array(plugin_attr, dtype=wp.vec3f),
+ geom_plugin_index=wp.array(geom_plugin_index, dtype=int),
+ mat_rgba=create_nmodel_batched_array(mjm.mat_rgba, dtype=wp.vec4),
+ actuator_trntype_body_adr=wp.array(np.nonzero(mjm.actuator_trntype == mujoco.mjtTrn.mjTRN_BODY)[0], dtype=int),
+ geompair2hfgeompair=wp.array(_hfield_geom_pair(mjm)[1], dtype=int),
+ block_dim=types.BlockDim(),
+ geom_pair_type_count=tuple(geom_type_pair_count),
+ has_sdf_geom=bool(np.any(mjm.geom_type == mujoco.mjtGeom.mjGEOM_SDF)),
+ )
+
+ return m
+
+
+def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: int = -1) -> types.Data:
+ """
+ Creates a data object on device.
+
+ Args:
+ mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host).
+ nworld (int, optional): Number of worlds. Defaults to 1.
+ nconmax (int, optional): Maximum number of contacts for all worlds. Defaults to -1.
+ njmax (int, optional): Maximum number of constraints for all worlds. Defaults to -1.
+
+ Returns:
+ Data: The data object containing the current state and output arrays (device).
+ """
+ # TODO(team): move to Model?
+ if nconmax == -1:
+ # TODO(team): heuristic for nconmax
+ nconmax = nworld * 20
+ if njmax == -1:
+ # TODO(team): heuristic for njmax
+ njmax = 20 * 6
+ condim = np.concatenate((mjm.geom_condim, mjm.pair_dim))
+ condim_max = np.max(condim) if len(condim) > 0 else 0
+
+ if mujoco.mj_isSparse(mjm):
+ qM = wp.zeros((nworld, 1, mjm.nM), dtype=float)
+ qLD = wp.zeros((nworld, 1, mjm.nM), dtype=float)
+ else:
+ qM = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float)
+ qLD = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float)
+
+ nrangefinder = sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER)
+
+ return types.Data(
+ nworld=nworld,
+ nconmax=nconmax,
+ njmax=njmax,
+ solver_niter=wp.zeros(nworld, dtype=int),
+ ncon=wp.zeros(1, dtype=int),
+ ncon_hfield=wp.zeros((nworld, _hfield_geom_pair(mjm)[0]), dtype=int), # warp only
+ ne=wp.zeros(nworld, dtype=int),
+ ne_connect=wp.zeros(nworld, dtype=int), # warp only
+ ne_weld=wp.zeros(nworld, dtype=int), # warp only
+ ne_jnt=wp.zeros(nworld, dtype=int), # warp only
+ ne_ten=wp.zeros(nworld, dtype=int), # warp only
+ nf=wp.zeros(nworld, dtype=int),
+ nl=wp.zeros(nworld, dtype=int),
+ nefc=wp.zeros(nworld, dtype=int),
+ nsolving=wp.zeros(1, dtype=int), # warp only
+ time=wp.zeros(nworld, dtype=float),
+ energy=wp.zeros(nworld, dtype=wp.vec2),
+ qpos=wp.zeros((nworld, mjm.nq), dtype=float),
+ qvel=wp.zeros((nworld, mjm.nv), dtype=float),
+ act=wp.zeros((nworld, mjm.na), dtype=float),
+ qacc_warmstart=wp.zeros((nworld, mjm.nv), dtype=float),
+ qacc_discrete=wp.zeros((nworld, mjm.nv), dtype=float),
+ ctrl=wp.zeros((nworld, mjm.nu), dtype=float),
+ qfrc_applied=wp.zeros((nworld, mjm.nv), dtype=float),
+ xfrc_applied=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector),
+ fluid_applied=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector),
+ eq_active=wp.array(np.tile(mjm.eq_active0, (nworld, 1)), dtype=bool),
+ mocap_pos=wp.zeros((nworld, mjm.nmocap), dtype=wp.vec3),
+ mocap_quat=wp.zeros((nworld, mjm.nmocap), dtype=wp.quat),
+ qacc=wp.zeros((nworld, mjm.nv), dtype=float),
+ act_dot=wp.zeros((nworld, mjm.na), dtype=float),
+ xpos=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3),
+ xquat=wp.zeros((nworld, mjm.nbody), dtype=wp.quat),
+ xmat=wp.zeros((nworld, mjm.nbody), dtype=wp.mat33),
+ xipos=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3),
+ ximat=wp.zeros((nworld, mjm.nbody), dtype=wp.mat33),
+ xanchor=wp.zeros((nworld, mjm.njnt), dtype=wp.vec3),
+ xaxis=wp.zeros((nworld, mjm.njnt), dtype=wp.vec3),
+ geom_skip=wp.zeros(mjm.ngeom, dtype=bool), # warp only
+ geom_xpos=wp.zeros((nworld, mjm.ngeom), dtype=wp.vec3),
+ geom_xmat=wp.zeros((nworld, mjm.ngeom), dtype=wp.mat33),
+ site_xpos=wp.zeros((nworld, mjm.nsite), dtype=wp.vec3),
+ site_xmat=wp.zeros((nworld, mjm.nsite), dtype=wp.mat33),
+ cam_xpos=wp.zeros((nworld, mjm.ncam), dtype=wp.vec3),
+ cam_xmat=wp.zeros((nworld, mjm.ncam), dtype=wp.mat33),
+ light_xpos=wp.zeros((nworld, mjm.nlight), dtype=wp.vec3),
+ light_xdir=wp.zeros((nworld, mjm.nlight), dtype=wp.vec3),
+ subtree_com=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3),
+ cdof=wp.zeros((nworld, mjm.nv), dtype=wp.spatial_vector),
+ cinert=wp.zeros((nworld, mjm.nbody), dtype=types.vec10),
+ flexvert_xpos=wp.zeros((nworld, mjm.nflexvert), dtype=wp.vec3),
+ flexedge_length=wp.zeros((nworld, mjm.nflexedge), dtype=wp.float32),
+ flexedge_velocity=wp.zeros((nworld, mjm.nflexedge), dtype=wp.float32),
+ actuator_length=wp.zeros((nworld, mjm.nu), dtype=float),
+ actuator_moment=wp.zeros((nworld, mjm.nu, mjm.nv), dtype=float),
+ crb=wp.zeros((nworld, mjm.nbody), dtype=types.vec10),
+ qM=qM,
+ qLD=qLD,
+ qLDiagInv=wp.zeros((nworld, mjm.nv), dtype=float),
+ ten_velocity=wp.zeros((nworld, mjm.ntendon), dtype=float),
+ actuator_velocity=wp.zeros((nworld, mjm.nu), dtype=float),
+ cvel=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector),
+ cdof_dot=wp.zeros((nworld, mjm.nv), dtype=wp.spatial_vector),
+ qfrc_bias=wp.zeros((nworld, mjm.nv), dtype=float),
+ qfrc_spring=wp.zeros((nworld, mjm.nv), dtype=float),
+ qfrc_damper=wp.zeros((nworld, mjm.nv), dtype=float),
+ qfrc_gravcomp=wp.zeros((nworld, mjm.nv), dtype=float),
+ qfrc_fluid=wp.zeros((nworld, mjm.nv), dtype=float),
+ qfrc_passive=wp.zeros((nworld, mjm.nv), dtype=float),
+ subtree_linvel=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3),
+ subtree_angmom=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3),
+ subtree_bodyvel=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), # warp only
+ actuator_force=wp.zeros((nworld, mjm.nu), dtype=float),
+ qfrc_actuator=wp.zeros((nworld, mjm.nv), dtype=float),
+ qfrc_smooth=wp.zeros((nworld, mjm.nv), dtype=float),
+ qacc_smooth=wp.zeros((nworld, mjm.nv), dtype=float),
+ qfrc_constraint=wp.zeros((nworld, mjm.nv), dtype=float),
+ qfrc_inverse=wp.zeros((nworld, mjm.nv), dtype=float),
+ contact=types.Contact(
+ dist=wp.zeros((nconmax,), dtype=float),
+ pos=wp.zeros((nconmax,), dtype=wp.vec3f),
+ frame=wp.zeros((nconmax,), dtype=wp.mat33f),
+ includemargin=wp.zeros((nconmax,), dtype=float),
+ friction=wp.zeros((nconmax,), dtype=types.vec5),
+ solref=wp.zeros((nconmax,), dtype=wp.vec2f),
+ solreffriction=wp.zeros((nconmax,), dtype=wp.vec2f),
+ solimp=wp.zeros((nconmax,), dtype=types.vec5),
+ dim=wp.zeros((nconmax,), dtype=int),
+ geom=wp.zeros((nconmax,), dtype=wp.vec2i),
+ efc_address=wp.zeros(
+ (nconmax, np.maximum(1, 2 * (condim_max - 1))),
+ dtype=int,
+ ),
+ worldid=wp.zeros((nconmax,), dtype=int),
+ ),
+ efc=types.Constraint(
+ type=wp.zeros((nworld, njmax), dtype=int),
+ id=wp.zeros((nworld, njmax), dtype=int),
+ J=wp.zeros((nworld, njmax, mjm.nv), dtype=float),
+ pos=wp.zeros((nworld, njmax), dtype=float),
+ margin=wp.zeros((nworld, njmax), dtype=float),
+ D=wp.zeros((nworld, njmax), dtype=float),
+ vel=wp.zeros((nworld, njmax), dtype=float),
+ aref=wp.zeros((nworld, njmax), dtype=float),
+ frictionloss=wp.zeros((nworld, njmax), dtype=float),
+ force=wp.zeros((nworld, njmax), dtype=float),
+ Jaref=wp.zeros((nworld, njmax), dtype=float),
+ Ma=wp.zeros((nworld, mjm.nv), dtype=float),
+ grad=wp.zeros((nworld, mjm.nv), dtype=float),
+ cholesky_L_tmp=wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float),
+ cholesky_y_tmp=wp.zeros((nworld, mjm.nv), dtype=float),
+ grad_dot=wp.zeros((nworld,), dtype=float),
+ Mgrad=wp.zeros((nworld, mjm.nv), dtype=float),
+ search=wp.zeros((nworld, mjm.nv), dtype=float),
+ search_dot=wp.zeros((nworld,), dtype=float),
+ gauss=wp.zeros((nworld,), dtype=float),
+ cost=wp.zeros((nworld,), dtype=float),
+ prev_cost=wp.zeros((nworld,), dtype=float),
+ active=wp.zeros((nworld, njmax), dtype=bool),
+ gtol=wp.zeros((nworld,), dtype=float),
+ mv=wp.zeros((nworld, mjm.nv), dtype=float),
+ jv=wp.zeros((nworld, njmax), dtype=float),
+ quad=wp.zeros((nworld, njmax), dtype=wp.vec3f),
+ quad_gauss=wp.zeros((nworld,), dtype=wp.vec3f),
+ h=wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float),
+ alpha=wp.zeros((nworld,), dtype=float),
+ prev_grad=wp.zeros((nworld, mjm.nv), dtype=float),
+ prev_Mgrad=wp.zeros((nworld, mjm.nv), dtype=float),
+ beta=wp.zeros((nworld,), dtype=float),
+ beta_num=wp.zeros((nworld,), dtype=float),
+ beta_den=wp.zeros((nworld,), dtype=float),
+ done=wp.zeros((nworld,), dtype=bool),
+ # linesearch
+ ls_done=wp.zeros((nworld,), dtype=bool),
+ p0=wp.zeros((nworld,), dtype=wp.vec3),
+ lo=wp.zeros((nworld,), dtype=wp.vec3),
+ lo_alpha=wp.zeros((nworld,), dtype=float),
+ hi=wp.zeros((nworld,), dtype=wp.vec3),
+ hi_alpha=wp.zeros((nworld,), dtype=float),
+ lo_next=wp.zeros((nworld,), dtype=wp.vec3),
+ lo_next_alpha=wp.zeros((nworld,), dtype=float),
+ hi_next=wp.zeros((nworld,), dtype=wp.vec3),
+ hi_next_alpha=wp.zeros((nworld,), dtype=float),
+ mid=wp.zeros((nworld,), dtype=wp.vec3),
+ mid_alpha=wp.zeros((nworld,), dtype=float),
+ cost_candidate=wp.zeros((nworld, mjm.opt.ls_iterations), dtype=float),
+ # elliptic cone
+ u=wp.zeros((nconmax,), dtype=types.vec6),
+ uu=wp.zeros((nconmax,), dtype=float),
+ uv=wp.zeros((nconmax,), dtype=float),
+ vv=wp.zeros((nconmax,), dtype=float),
+ condim=wp.zeros((nworld, njmax), dtype=int),
+ ),
+ # RK4
+ qpos_t0=wp.zeros((nworld, mjm.nq), dtype=float),
+ qvel_t0=wp.zeros((nworld, mjm.nv), dtype=float),
+ act_t0=wp.zeros((nworld, mjm.na), dtype=float),
+ qvel_rk=wp.zeros((nworld, mjm.nv), dtype=float),
+ qacc_rk=wp.zeros((nworld, mjm.nv), dtype=float),
+ act_dot_rk=wp.zeros((nworld, mjm.na), dtype=float),
+ # euler + implicit integration
+ qfrc_integration=wp.zeros((nworld, mjm.nv), dtype=float),
+ qacc_integration=wp.zeros((nworld, mjm.nv), dtype=float),
+ act_vel_integration=wp.zeros((nworld, mjm.nu), dtype=float),
+ qM_integration=wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float),
+ qLD_integration=wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float),
+ qLDiagInv_integration=wp.zeros((nworld, mjm.nv), dtype=float),
+ # sweep-and-prune broadphase
+ sap_projection_lower=wp.zeros((nworld, mjm.ngeom, 2), dtype=float),
+ sap_projection_upper=wp.zeros((nworld, mjm.ngeom), dtype=float),
+ sap_sort_index=wp.zeros((nworld, mjm.ngeom, 2), dtype=int),
+ sap_range=wp.zeros((nworld, mjm.ngeom), dtype=int),
+ sap_cumulative_sum=wp.zeros((nworld, mjm.ngeom), dtype=int),
+ sap_segment_index=wp.array(
+ np.array([i * mjm.ngeom if i < nworld + 1 else 0 for i in range(2 * nworld)]).reshape((nworld, 2)), dtype=int
+ ),
+ # collision driver
+ collision_pair=wp.zeros((nconmax,), dtype=wp.vec2i),
+ collision_hftri_index=wp.zeros((nconmax,), dtype=int),
+ collision_pairid=wp.zeros((nconmax,), dtype=int),
+ collision_worldid=wp.zeros((nconmax,), dtype=int),
+ ncollision=wp.zeros((1,), dtype=int),
+ # narrowphase (EPA polytope)
+ epa_vert=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=wp.vec3),
+ epa_vert1=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=wp.vec3),
+ epa_vert2=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=wp.vec3),
+ epa_vert_index1=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=int),
+ epa_vert_index2=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=int),
+ epa_face=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=wp.vec3i),
+ epa_pr=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=wp.vec3),
+ epa_norm2=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=float),
+ epa_index=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=int),
+ epa_map=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=int),
+ epa_horizon=wp.zeros(shape=(nconmax, 6 * MJ_CCD_ITERATIONS), dtype=int),
+ # rne_postconstraint
+ cacc=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector),
+ cfrc_int=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector),
+ cfrc_ext=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector),
+ # tendon
+ ten_length=wp.zeros((nworld, mjm.ntendon), dtype=float),
+ ten_J=wp.zeros((nworld, mjm.ntendon, mjm.nv), dtype=float),
+ ten_Jdot=wp.zeros((nworld, mjm.ntendon, mjm.nv), dtype=float),
+ ten_bias_coef=wp.zeros((nworld, mjm.ntendon), dtype=float),
+ ten_wrapadr=wp.zeros((nworld, mjm.ntendon), dtype=int),
+ ten_wrapnum=wp.zeros((nworld, mjm.ntendon), dtype=int),
+ ten_actfrc=wp.zeros((nworld, mjm.ntendon), dtype=float),
+ wrap_obj=wp.zeros((nworld, mjm.nwrap), dtype=wp.vec2i),
+ wrap_xpos=wp.zeros((nworld, mjm.nwrap), dtype=wp.spatial_vector),
+ wrap_geom_xpos=wp.zeros((nworld, mjm.nwrap), dtype=wp.spatial_vector),
+ # sensors
+ sensordata=wp.zeros((nworld, mjm.nsensordata), dtype=float),
+ sensor_rangefinder_pnt=wp.zeros((nworld, nrangefinder), dtype=wp.vec3),
+ sensor_rangefinder_vec=wp.zeros((nworld, nrangefinder), dtype=wp.vec3),
+ sensor_rangefinder_dist=wp.zeros((nworld, nrangefinder), dtype=float),
+ sensor_rangefinder_geomid=wp.zeros((nworld, nrangefinder), dtype=int),
+ # ray
+ ray_bodyexclude=wp.zeros(1, dtype=int),
+ ray_dist=wp.zeros((nworld, 1), dtype=float),
+ ray_geomid=wp.zeros((nworld, 1), dtype=int),
+ # mul_m
+ energy_vel_mul_m_skip=wp.zeros((nworld,), dtype=bool),
+ inverse_mul_m_skip=wp.zeros((nworld,), dtype=bool),
+ # actuator
+ actuator_trntype_body_ncon=wp.zeros((nworld, np.sum(mjm.actuator_trntype == mujoco.mjtTrn.mjTRN_BODY)), dtype=int),
+ )
+
+
+def put_data(
+ mjm: mujoco.MjModel,
+ mjd: mujoco.MjData,
+ nworld: Optional[int] = None,
+ nconmax: Optional[int] = None,
+ njmax: Optional[int] = None,
+) -> types.Data:
+ """
+ Moves data from host to a device.
+
+ Args:
+ mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host).
+ mjd (mujoco.MjData): The data object containing current state and output arrays (host).
+ nworld (int, optional): The number of worlds. Defaults to 1.
+ nconmax (int, optional): The maximum number of contacts for all worlds. Defaults to -1.
+ njmax (int, optional): The maximum number of constraints for all worlds. Defaults to -1.
+
+ Returns:
+ Data: The data object containing the current state and output arrays (device).
+ """
+ # TODO(team): move nconmax and njmax to Model?
+ # TODO(team): decide what to do about uninitialized warp-only fields created by put_data
+ # we need to ensure these are only workspace fields and don't carry state
+
+ nworld = nworld or 1
+ # TODO(team): better heuristic for nconmax
+ nconmax = nconmax or max(512, mjd.ncon * nworld)
+ # TODO(team): better heuristic for njmax
+ njmax = njmax or max(5, mjd.nefc)
+
+ if nworld < 1:
+ raise ValueError("nworld must be >= 1")
+
+ if nconmax < 1:
+ raise ValueError("nconmax must be >= 1")
+
+ if njmax < 1:
+ raise ValueError("njmax must be >= 1")
+
+ if nworld * mjd.ncon > nconmax:
+ raise ValueError(f"nconmax overflow (nconmax must be >= {nworld * mjd.ncon})")
+
+ if mjd.nefc > njmax:
+ raise ValueError(f"njmax overflow (njmax must be >= {mjd.nefc})")
+
+ # calculate some fields that cannot be easily computed inline:
+ if mujoco.mj_isSparse(mjm):
+ qM = np.expand_dims(mjd.qM, axis=0)
+ qLD = np.expand_dims(mjd.qLD, axis=0)
+ qM_integration = np.zeros((1, mjm.nM), dtype=float)
+ qLD_integration = np.zeros((1, mjm.nM), dtype=float)
+ efc_J = np.zeros((mjd.nefc, mjm.nv))
+ mujoco.mju_sparse2dense(efc_J, mjd.efc_J, mjd.efc_J_rownnz, mjd.efc_J_rowadr, mjd.efc_J_colind)
+ ten_J = np.zeros((mjm.ntendon, mjm.nv))
+ mujoco.mju_sparse2dense(
+ ten_J,
+ mjd.ten_J.reshape(-1),
+ mjd.ten_J_rownnz,
+ mjd.ten_J_rowadr,
+ mjd.ten_J_colind.reshape(-1),
+ )
+ else:
+ qM = np.zeros((mjm.nv, mjm.nv))
+ mujoco.mj_fullM(mjm, qM, mjd.qM)
+ if (mjd.qM == 0.0).all() or (mjd.qLD == 0.0).all():
+ qLD = np.zeros((mjm.nv, mjm.nv))
+ else:
+ qLD = np.linalg.cholesky(qM)
+ qM_integration = np.zeros((mjm.nv, mjm.nv), dtype=float)
+ qLD_integration = np.zeros((mjm.nv, mjm.nv), dtype=float)
+ efc_J = mjd.efc_J.reshape((mjd.nefc, mjm.nv))
+ ten_J = mjd.ten_J.reshape((mjm.ntendon, mjm.nv))
+
+ # TODO(taylorhowell): sparse actuator_moment
+ actuator_moment = np.zeros((mjm.nu, mjm.nv))
+ mujoco.mju_sparse2dense(
+ actuator_moment,
+ mjd.actuator_moment,
+ mjd.moment_rownnz,
+ mjd.moment_rowadr,
+ mjd.moment_colind,
+ )
+
+ condim = np.concatenate((mjm.geom_condim, mjm.pair_dim))
+ condim_max = np.max(condim) if len(condim) > 0 else 0
+ contact_efc_address = np.zeros((nconmax, np.maximum(1, 2 * (condim_max - 1))), dtype=int)
+ for i in range(nworld):
+ for j in range(mjd.ncon):
+ condim = mjd.contact.dim[j]
+ efc_address = mjd.contact.efc_address[j]
+ if efc_address == -1:
+ continue
+ if condim == 1:
+ nconvar = 1
+ else:
+ nconvar = condim if mjm.opt.cone == mujoco.mjtCone.mjCONE_ELLIPTIC else 2 * (condim - 1)
+ for k in range(nconvar):
+ contact_efc_address[i * mjd.ncon + j, k] = mjd.nefc * i + efc_address + k
+
+ contact_worldid = np.pad(np.repeat(np.arange(nworld), mjd.ncon), (0, nconmax - nworld * mjd.ncon))
+
+ ne_connect = int(3 * np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_CONNECT) & mjd.eq_active))
+ ne_weld = int(6 * np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_WELD) & mjd.eq_active))
+ ne_jnt = int(np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_JOINT) & mjd.eq_active))
+ ne_ten = int(np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_TENDON) & mjd.eq_active))
+
+ efc_type_fill = np.zeros((nworld, njmax))
+ efc_id_fill = np.zeros((nworld, njmax))
+ efc_J_fill = np.zeros((nworld, njmax, mjm.nv))
+ efc_D_fill = np.zeros((nworld, njmax))
+ efc_vel_fill = np.zeros((nworld, njmax))
+ efc_pos_fill = np.zeros((nworld, njmax))
+ efc_aref_fill = np.zeros((nworld, njmax))
+ efc_frictionloss_fill = np.zeros((nworld, njmax))
+ efc_force_fill = np.zeros((nworld, njmax))
+ efc_margin_fill = np.zeros((nworld, njmax))
+
+ nefc = mjd.nefc
+ efc_type_fill[:, :nefc] = np.tile(mjd.efc_type, (nworld, 1))
+ efc_id_fill[:, :nefc] = np.tile(mjd.efc_id, (nworld, 1))
+ efc_J_fill[:, :nefc, :] = np.tile(efc_J, (nworld, 1, 1))
+ efc_D_fill[:, :nefc] = np.tile(mjd.efc_D, (nworld, 1))
+ efc_vel_fill[:, :nefc] = np.tile(mjd.efc_vel, (nworld, 1))
+ efc_pos_fill[:, :nefc] = np.tile(mjd.efc_pos, (nworld, 1))
+ efc_aref_fill[:, :nefc] = np.tile(mjd.efc_aref, (nworld, 1))
+ efc_frictionloss_fill[:, :nefc] = np.tile(mjd.efc_frictionloss, (nworld, 1))
+ efc_force_fill[:, :nefc] = np.tile(mjd.efc_force, (nworld, 1))
+ efc_margin_fill[:, :nefc] = np.tile(mjd.efc_margin, (nworld, 1))
+
+ nrangefinder = sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER)
+
+ # some helper functions to simplify the data field definitions below
+
+ def arr(x, dtype=None):
+ if not isinstance(x, np.ndarray):
+ x = np.array(x)
+ if dtype is None:
+ if np.issubdtype(x.dtype, np.integer):
+ dtype = wp.int32
+ elif np.issubdtype(x.dtype, np.floating):
+ dtype = wp.float32
+ elif np.issubdtype(x.dtype, bool):
+ dtype = wp.bool
+ else:
+ raise ValueError(f"Unsupported dtype: {x.dtype}")
+ wp_array = {1: wp.array, 2: wp.array2d, 3: wp.array3d}[x.ndim]
+ return wp_array(x, dtype=dtype)
+
+ def tile(x, dtype=None):
+ return arr(np.tile(x, (nworld,) + (1,) * len(x.shape)), dtype)
+
+ def padtile(x, length, dtype=None):
+ x = np.repeat(x, nworld, axis=0)
+ width = ((0, length - x.shape[0]),) + ((0, 0),) * (x.ndim - 1)
+ return arr(np.pad(x, width), dtype)
+
+ return types.Data(
+ nworld=nworld,
+ nconmax=nconmax,
+ njmax=njmax,
+ solver_niter=tile(mjd.solver_niter[0]),
+ ncon=arr([mjd.ncon * nworld]),
+ ncon_hfield=wp.zeros((nworld, _hfield_geom_pair(mjm)[0]), dtype=int), # warp only
+ ne=wp.full(shape=(nworld), value=mjd.ne),
+ ne_connect=wp.full(shape=(nworld), value=ne_connect),
+ ne_weld=wp.full(shape=(nworld), value=ne_weld),
+ ne_jnt=wp.full(shape=(nworld), value=ne_jnt),
+ ne_ten=wp.full(shape=(nworld), value=ne_ten),
+ nf=wp.full(shape=(nworld), value=mjd.nf),
+ nl=wp.full(shape=(nworld), value=mjd.nl),
+ nefc=wp.full(shape=(nworld), value=mjd.nefc),
+ nsolving=arr([nworld]),
+ time=arr(mjd.time * np.ones(nworld)),
+ energy=tile(mjd.energy, dtype=wp.vec2),
+ qpos=tile(mjd.qpos),
+ qvel=tile(mjd.qvel),
+ act=tile(mjd.act),
+ qacc_warmstart=tile(mjd.qacc_warmstart),
+ qacc_discrete=wp.zeros((nworld, mjm.nv), dtype=float),
+ ctrl=tile(mjd.ctrl),
+ qfrc_applied=tile(mjd.qfrc_applied),
+ xfrc_applied=tile(mjd.xfrc_applied, dtype=wp.spatial_vector),
+ fluid_applied=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector),
+ eq_active=tile(mjd.eq_active.astype(bool)),
+ mocap_pos=tile(mjd.mocap_pos, dtype=wp.vec3),
+ mocap_quat=tile(mjd.mocap_quat, dtype=wp.quat),
+ qacc=tile(mjd.qacc),
+ act_dot=tile(mjd.act_dot),
+ xpos=tile(mjd.xpos, dtype=wp.vec3),
+ xquat=tile(mjd.xquat, dtype=wp.quat),
+ xmat=tile(mjd.xmat, dtype=wp.mat33),
+ xipos=tile(mjd.xipos, dtype=wp.vec3),
+ ximat=tile(mjd.ximat, dtype=wp.mat33),
+ xanchor=tile(mjd.xanchor, dtype=wp.vec3),
+ xaxis=tile(mjd.xaxis, dtype=wp.vec3),
+ geom_skip=wp.zeros(mjm.ngeom, dtype=bool), # warp only
+ geom_xpos=tile(mjd.geom_xpos, dtype=wp.vec3),
+ geom_xmat=tile(mjd.geom_xmat, dtype=wp.mat33),
+ site_xpos=tile(mjd.site_xpos, dtype=wp.vec3),
+ site_xmat=tile(mjd.site_xmat, dtype=wp.mat33),
+ cam_xpos=tile(mjd.cam_xpos, dtype=wp.vec3),
+ cam_xmat=tile(mjd.cam_xmat, dtype=wp.mat33),
+ light_xpos=tile(mjd.light_xpos, dtype=wp.vec3),
+ light_xdir=tile(mjd.light_xdir, dtype=wp.vec3),
+ subtree_com=tile(mjd.subtree_com, dtype=wp.vec3),
+ cdof=tile(mjd.cdof, dtype=wp.spatial_vector),
+ cinert=tile(mjd.cinert, dtype=types.vec10),
+ flexvert_xpos=tile(mjd.flexvert_xpos, dtype=wp.vec3),
+ flexedge_length=tile(mjd.flexedge_length),
+ flexedge_velocity=tile(mjd.flexedge_velocity),
+ actuator_length=tile(mjd.actuator_length),
+ actuator_moment=tile(actuator_moment),
+ crb=tile(mjd.crb, dtype=types.vec10),
+ qM=tile(qM),
+ qLD=tile(qLD),
+ qLDiagInv=tile(mjd.qLDiagInv),
+ ten_velocity=tile(mjd.ten_velocity),
+ actuator_velocity=tile(mjd.actuator_velocity),
+ cvel=tile(mjd.cvel, dtype=wp.spatial_vector),
+ cdof_dot=tile(mjd.cdof_dot, dtype=wp.spatial_vector),
+ qfrc_bias=tile(mjd.qfrc_bias),
+ qfrc_spring=tile(mjd.qfrc_spring),
+ qfrc_damper=tile(mjd.qfrc_damper),
+ qfrc_gravcomp=tile(mjd.qfrc_gravcomp),
+ qfrc_fluid=tile(mjd.qfrc_fluid),
+ qfrc_passive=tile(mjd.qfrc_passive),
+ subtree_linvel=tile(mjd.subtree_linvel, dtype=wp.vec3),
+ subtree_angmom=tile(mjd.subtree_angmom, dtype=wp.vec3),
+ subtree_bodyvel=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector),
+ actuator_force=tile(mjd.actuator_force),
+ qfrc_actuator=tile(mjd.qfrc_actuator),
+ qfrc_smooth=tile(mjd.qfrc_smooth),
+ qacc_smooth=tile(mjd.qacc_smooth),
+ qfrc_constraint=tile(mjd.qfrc_constraint),
+ qfrc_inverse=tile(mjd.qfrc_inverse),
+ contact=types.Contact(
+ dist=padtile(mjd.contact.dist, nconmax),
+ pos=padtile(mjd.contact.pos, nconmax, dtype=wp.vec3),
+ frame=padtile(mjd.contact.frame, nconmax, dtype=wp.mat33),
+ includemargin=padtile(mjd.contact.includemargin, nconmax),
+ friction=padtile(mjd.contact.friction, nconmax, dtype=types.vec5),
+ solref=padtile(mjd.contact.solref, nconmax, dtype=wp.vec2f),
+ solreffriction=padtile(mjd.contact.solreffriction, nconmax, dtype=wp.vec2f),
+ solimp=padtile(mjd.contact.solimp, nconmax, dtype=types.vec5),
+ dim=padtile(mjd.contact.dim, nconmax),
+ geom=padtile(mjd.contact.geom, nconmax, dtype=wp.vec2i),
+ efc_address=arr(contact_efc_address),
+ worldid=arr(contact_worldid),
+ ),
+ efc=types.Constraint(
+ type=wp.array2d(efc_type_fill, dtype=int),
+ id=wp.array2d(efc_id_fill, dtype=int),
+ J=wp.array3d(efc_J_fill, dtype=float),
+ pos=wp.array2d(efc_pos_fill, dtype=float),
+ margin=wp.array2d(efc_margin_fill, dtype=float),
+ D=wp.array2d(efc_D_fill, dtype=float),
+ vel=wp.array2d(efc_vel_fill, dtype=float),
+ aref=wp.array2d(efc_aref_fill, dtype=float),
+ frictionloss=wp.array2d(efc_frictionloss_fill, dtype=float),
+ force=wp.array2d(efc_force_fill, dtype=float),
+ Jaref=wp.empty(shape=(nworld, njmax), dtype=float),
+ Ma=wp.empty(shape=(nworld, mjm.nv), dtype=float),
+ grad=wp.empty(shape=(nworld, mjm.nv), dtype=float),
+ cholesky_L_tmp=wp.empty(shape=(nworld, mjm.nv, mjm.nv), dtype=float),
+ cholesky_y_tmp=wp.empty(shape=(nworld, mjm.nv), dtype=float),
+ grad_dot=wp.empty(shape=(nworld,), dtype=float),
+ Mgrad=wp.empty(shape=(nworld, mjm.nv), dtype=float),
+ search=wp.empty(shape=(nworld, mjm.nv), dtype=float),
+ search_dot=wp.empty(shape=(nworld,), dtype=float),
+ gauss=wp.empty(shape=(nworld,), dtype=float),
+ cost=wp.empty(shape=(nworld,), dtype=float),
+ prev_cost=wp.empty(shape=(nworld,), dtype=float),
+ active=wp.empty(shape=(nworld, njmax), dtype=bool),
+ gtol=wp.empty(shape=(nworld,), dtype=float),
+ mv=wp.empty(shape=(nworld, mjm.nv), dtype=float),
+ jv=wp.empty(shape=(nworld, njmax), dtype=float),
+ quad=wp.empty(shape=(nworld, njmax), dtype=wp.vec3f),
+ quad_gauss=wp.empty(shape=(nworld,), dtype=wp.vec3f),
+ h=wp.empty(shape=(nworld, mjm.nv, mjm.nv), dtype=float),
+ alpha=wp.empty(shape=(nworld,), dtype=float),
+ prev_grad=wp.empty(shape=(nworld, mjm.nv), dtype=float),
+ prev_Mgrad=wp.empty(shape=(nworld, mjm.nv), dtype=float),
+ beta=wp.empty(shape=(nworld,), dtype=float),
+ beta_num=wp.empty(shape=(nworld,), dtype=float),
+ beta_den=wp.empty(shape=(nworld,), dtype=float),
+ done=wp.empty(shape=(nworld,), dtype=bool),
+ ls_done=wp.zeros(shape=(nworld,), dtype=bool),
+ p0=wp.empty(shape=(nworld,), dtype=wp.vec3),
+ lo=wp.empty(shape=(nworld,), dtype=wp.vec3),
+ lo_alpha=wp.empty(shape=(nworld,), dtype=float),
+ hi=wp.empty(shape=(nworld,), dtype=wp.vec3),
+ hi_alpha=wp.empty(shape=(nworld,), dtype=float),
+ lo_next=wp.empty(shape=(nworld,), dtype=wp.vec3),
+ lo_next_alpha=wp.empty(shape=(nworld,), dtype=float),
+ hi_next=wp.empty(shape=(nworld,), dtype=wp.vec3),
+ hi_next_alpha=wp.empty(shape=(nworld,), dtype=float),
+ mid=wp.empty(shape=(nworld,), dtype=wp.vec3),
+ mid_alpha=wp.empty(shape=(nworld,), dtype=float),
+ cost_candidate=wp.empty(shape=(nworld, mjm.opt.ls_iterations), dtype=float),
+ # TODO(team): skip allocation if not elliptic
+ u=wp.empty((nconmax,), dtype=types.vec6),
+ uu=wp.empty((nconmax,), dtype=float),
+ uv=wp.empty((nconmax,), dtype=float),
+ vv=wp.empty((nconmax,), dtype=float),
+ condim=wp.empty((nworld, njmax), dtype=int),
+ ),
+ # TODO(team): skip allocation if integrator != RK4
+ qpos_t0=wp.empty((nworld, mjm.nq), dtype=float),
+ qvel_t0=wp.empty((nworld, mjm.nv), dtype=float),
+ act_t0=wp.empty((nworld, mjm.na), dtype=float),
+ qvel_rk=wp.empty((nworld, mjm.nv), dtype=float),
+ qacc_rk=wp.empty((nworld, mjm.nv), dtype=float),
+ act_dot_rk=wp.empty((nworld, mjm.na), dtype=float),
+ # TODO(team): skip allocation if integrator != euler | implicit
+ qfrc_integration=wp.zeros((nworld, mjm.nv), dtype=float),
+ qacc_integration=wp.zeros((nworld, mjm.nv), dtype=float),
+ act_vel_integration=wp.zeros((nworld, mjm.nu), dtype=float),
+ qM_integration=tile(qM_integration),
+ qLD_integration=tile(qLD_integration),
+ qLDiagInv_integration=wp.zeros((nworld, mjm.nv), dtype=float),
+ # TODO(team): skip allocation if broadphase != sap
+ sap_projection_lower=wp.zeros((nworld, mjm.ngeom, 2), dtype=float),
+ sap_projection_upper=wp.zeros((nworld, mjm.ngeom), dtype=float),
+ sap_sort_index=wp.zeros((nworld, mjm.ngeom, 2), dtype=int),
+ sap_range=wp.zeros((nworld, mjm.ngeom), dtype=int),
+ sap_cumulative_sum=wp.zeros((nworld, mjm.ngeom), dtype=int),
+ sap_segment_index=arr(np.array([i * mjm.ngeom if i < nworld + 1 else 0 for i in range(2 * nworld)]).reshape((nworld, 2))),
+ # collision driver
+ collision_pair=wp.empty(nconmax, dtype=wp.vec2i),
+ collision_hftri_index=wp.empty(nconmax, dtype=int),
+ collision_pairid=wp.empty(nconmax, dtype=int),
+ collision_worldid=wp.empty(nconmax, dtype=int),
+ ncollision=wp.zeros(1, dtype=int),
+ # narrowphase (EPA polytope)
+ epa_vert=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=wp.vec3),
+ epa_vert1=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=wp.vec3),
+ epa_vert2=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=wp.vec3),
+ epa_vert_index1=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=int),
+ epa_vert_index2=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=int),
+ epa_face=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=wp.vec3i),
+ epa_pr=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=wp.vec3),
+ epa_norm2=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=float),
+ epa_index=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=int),
+ epa_map=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=int),
+ epa_horizon=wp.zeros(shape=(nconmax, 6 * MJ_CCD_ITERATIONS), dtype=int),
+ # rne_postconstraint but also smooth
+ cacc=tile(mjd.cacc, dtype=wp.spatial_vector),
+ cfrc_int=tile(mjd.cfrc_int, dtype=wp.spatial_vector),
+ cfrc_ext=tile(mjd.cfrc_ext, dtype=wp.spatial_vector),
+ # tendon
+ ten_length=tile(mjd.ten_length),
+ ten_J=tile(ten_J),
+ ten_Jdot=wp.zeros((nworld, mjm.ntendon, mjm.nv), dtype=float),
+ ten_bias_coef=wp.zeros((nworld, mjm.ntendon), dtype=float),
+ ten_wrapadr=tile(mjd.ten_wrapadr),
+ ten_wrapnum=tile(mjd.ten_wrapnum),
+ ten_actfrc=wp.zeros((nworld, mjm.ntendon), dtype=float),
+ wrap_obj=tile(mjd.wrap_obj, dtype=wp.vec2i),
+ wrap_xpos=tile(mjd.wrap_xpos, dtype=wp.spatial_vector),
+ wrap_geom_xpos=wp.zeros((nworld, mjm.nwrap), dtype=wp.spatial_vector),
+ # sensors
+ sensordata=tile(mjd.sensordata),
+ sensor_rangefinder_pnt=wp.zeros((nworld, nrangefinder), dtype=wp.vec3),
+ sensor_rangefinder_vec=wp.zeros((nworld, nrangefinder), dtype=wp.vec3),
+ sensor_rangefinder_dist=wp.zeros((nworld, nrangefinder), dtype=float),
+ sensor_rangefinder_geomid=wp.zeros((nworld, nrangefinder), dtype=int),
+ # ray
+ ray_bodyexclude=wp.zeros(1, dtype=int),
+ ray_dist=wp.zeros((nworld, 1), dtype=float),
+ ray_geomid=wp.zeros((nworld, 1), dtype=int),
+ # mul_m
+ energy_vel_mul_m_skip=wp.zeros((nworld,), dtype=bool),
+ inverse_mul_m_skip=wp.zeros((nworld,), dtype=bool),
+ # actuator
+ actuator_trntype_body_ncon=wp.zeros((nworld, np.sum(mjm.actuator_trntype == mujoco.mjtTrn.mjTRN_BODY)), dtype=int),
+ )
+
+
+def get_data_into(
+ result: mujoco.MjData,
+ mjm: mujoco.MjModel,
+ d: types.Data,
+):
+ """Gets data from a device into an existing mujoco.MjData.
+
+ Args:
+ result (mujoco.MjData): The data object containing the current state and output arrays
+ (host).
+ mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host).
+ d (Data): The data object containing the current state and output arrays (device).
+ """
+ if d.nworld > 1:
+ raise NotImplementedError("only nworld == 1 supported for now")
+
+ result.solver_niter[0] = d.solver_niter.numpy()[0]
+
+ ncon = d.ncon.numpy()[0]
+ nefc = d.nefc.numpy()[0]
+
+ if ncon != result.ncon or nefc != result.nefc:
+ mujoco._functions._realloc_con_efc(result, ncon=ncon, nefc=nefc)
+
+ result.time = d.time.numpy()[0]
+ result.energy = d.energy.numpy()[0]
+ result.ne = d.ne.numpy()[0]
+ result.qpos[:] = d.qpos.numpy()[0]
+ result.qvel[:] = d.qvel.numpy()[0]
+ result.qacc_warmstart = d.qacc_warmstart.numpy()[0]
+ result.qfrc_applied = d.qfrc_applied.numpy()[0]
+ result.mocap_pos = d.mocap_pos.numpy()[0]
+ result.mocap_quat = d.mocap_quat.numpy()[0]
+ result.qacc = d.qacc.numpy()[0]
+ result.xanchor = d.xanchor.numpy()[0]
+ result.xaxis = d.xaxis.numpy()[0]
+ result.xmat = d.xmat.numpy().reshape((-1, 9))
+ result.xpos = d.xpos.numpy()[0]
+ result.xquat = d.xquat.numpy()[0]
+ result.xipos = d.xipos.numpy()[0]
+ result.ximat = d.ximat.numpy().reshape((-1, 9))
+ result.subtree_com = d.subtree_com.numpy()[0]
+ result.geom_xpos = d.geom_xpos.numpy()[0]
+ result.geom_xmat = d.geom_xmat.numpy().reshape((-1, 9))
+ result.site_xpos = d.site_xpos.numpy()[0]
+ result.site_xmat = d.site_xmat.numpy().reshape((-1, 9))
+ result.cam_xpos = d.cam_xpos.numpy()[0]
+ result.cam_xmat = d.cam_xmat.numpy().reshape((-1, 9))
+ result.light_xpos = d.light_xpos.numpy()[0]
+ result.light_xdir = d.light_xdir.numpy()[0]
+ result.cinert = d.cinert.numpy()[0]
+ result.flexvert_xpos = d.flexvert_xpos.numpy()[0]
+ result.flexedge_length = d.flexedge_length.numpy()[0]
+ result.flexedge_velocity = d.flexedge_velocity.numpy()[0]
+ result.cdof = d.cdof.numpy()[0]
+ result.crb = d.crb.numpy()[0]
+ result.qLDiagInv = d.qLDiagInv.numpy()[0]
+ result.ctrl = d.ctrl.numpy()[0]
+ result.ten_velocity = d.ten_velocity.numpy()[0]
+ result.actuator_velocity = d.actuator_velocity.numpy()[0]
+ result.actuator_force = d.actuator_force.numpy()[0]
+ result.actuator_length = d.actuator_length.numpy()[0]
+ mujoco.mju_dense2sparse(
+ result.actuator_moment,
+ d.actuator_moment.numpy()[0],
+ result.moment_rownnz,
+ result.moment_rowadr,
+ result.moment_colind,
+ )
+ result.cvel = d.cvel.numpy()[0]
+ result.cdof_dot = d.cdof_dot.numpy()[0]
+ result.qfrc_bias = d.qfrc_bias.numpy()[0]
+ result.qfrc_fluid = d.qfrc_fluid.numpy()[0]
+ result.qfrc_passive = d.qfrc_passive.numpy()[0]
+ result.subtree_linvel = d.subtree_linvel.numpy()[0]
+ result.subtree_angmom = d.subtree_angmom.numpy()[0]
+ result.qfrc_spring = d.qfrc_spring.numpy()[0]
+ result.qfrc_damper = d.qfrc_damper.numpy()[0]
+ result.qfrc_gravcomp = d.qfrc_gravcomp.numpy()[0]
+ result.qfrc_fluid = d.qfrc_fluid.numpy()[0]
+ result.qfrc_actuator = d.qfrc_actuator.numpy()[0]
+ result.qfrc_smooth = d.qfrc_smooth.numpy()[0]
+ result.qfrc_constraint = d.qfrc_constraint.numpy()[0]
+ result.qfrc_inverse = d.qfrc_inverse.numpy()[0]
+ result.qacc_smooth = d.qacc_smooth.numpy()[0]
+ result.act = d.act.numpy()[0]
+ result.act_dot = d.act_dot.numpy()[0]
+
+ result.contact.dist[:] = d.contact.dist.numpy()[:ncon]
+ result.contact.pos[:] = d.contact.pos.numpy()[:ncon]
+ result.contact.frame[:] = d.contact.frame.numpy()[:ncon].reshape((-1, 9))
+ result.contact.includemargin[:] = d.contact.includemargin.numpy()[:ncon]
+ result.contact.friction[:] = d.contact.friction.numpy()[:ncon]
+ result.contact.solref[:] = d.contact.solref.numpy()[:ncon]
+ result.contact.solreffriction[:] = d.contact.solreffriction.numpy()[:ncon]
+ result.contact.solimp[:] = d.contact.solimp.numpy()[:ncon]
+ result.contact.dim[:] = d.contact.dim.numpy()[:ncon]
+ result.contact.efc_address[:] = d.contact.efc_address.numpy()[:ncon, 0]
+
+ if mujoco.mj_isSparse(mjm):
+ result.qM[:] = d.qM.numpy()[0, 0]
+ result.qLD[:] = d.qLD.numpy()[0, 0]
+ # TODO(team): set efc_J after fix to _realloc_con_efc lands
+ # efc_J = d.efc_J.numpy()[0, :nefc]
+ # mujoco.mju_dense2sparse(
+ # result.efc_J, efc_J, result.efc_J_rownnz, result.efc_J_rowadr, result.efc_J_colind
+ # )
+ else:
+ qM = d.qM.numpy()
+ adr = 0
+ for i in range(mjm.nv):
+ j = i
+ while j >= 0:
+ result.qM[adr] = qM[0, i, j]
+ j = mjm.dof_parentid[j]
+ adr += 1
+ mujoco.mj_factorM(mjm, result)
+ # TODO(team): set efc_J after fix to _realloc_con_efc lands
+ # if nefc > 0:
+ # result.efc_J[:nefc * mjm.nv] = d.efc_J.numpy()[:nefc].flatten()
+ result.xfrc_applied[:] = d.xfrc_applied.numpy()[0]
+ result.eq_active[:] = d.eq_active.numpy()[0]
+
+ # TODO(team): set these efc_* fields after fix to _realloc_con_efc
+ # Safely copy only up to the minimum of the destination and source sizes
+ # n = min(result.efc_D.shape[0], d.efc.D.numpy()[:nefc].shape[0])
+ # result.efc_D[:n] = d.efc.D.numpy()[:nefc][:n]
+ # n_pos = min(result.efc_pos.shape[0], d.efc.pos.numpy()[:nefc].shape[0])
+ # result.efc_pos[:n_pos] = d.efc.pos.numpy()[:nefc][:n_pos]
+
+ # n_aref = min(result.efc_aref.shape[0], d.efc.aref.numpy()[:nefc].shape[0])
+ # result.efc_aref[:n_aref] = d.efc.aref.numpy()[:nefc][:n_aref]
+
+ # n_force = min(result.efc_force.shape[0], d.efc.force.numpy()[:nefc].shape[0])
+ # result.efc_force[:n_force] = d.efc.force.numpy()[:nefc][:n_force]
+
+ # n_margin = min(result.efc_margin.shape[0], d.efc.margin.numpy()[:nefc].shape[0])
+ # result.efc_margin[:n_margin] = d.efc.margin.numpy()[:nefc][:n_margin]
+
+ result.cacc[:] = d.cacc.numpy()[0]
+ result.cfrc_int[:] = d.cfrc_int.numpy()[0]
+ result.cfrc_ext[:] = d.cfrc_ext.numpy()[0]
+
+ # TODO: other efc_ fields, anything else missing
+
+ # tendon
+ result.ten_length[:] = d.ten_length.numpy()[0]
+ result.ten_J[:] = d.ten_J.numpy()[0]
+ result.ten_wrapadr[:] = d.ten_wrapadr.numpy()[0]
+ result.ten_wrapnum[:] = d.ten_wrapnum.numpy()[0]
+ result.wrap_obj[:] = d.wrap_obj.numpy()[0]
+ result.wrap_xpos[:] = d.wrap_xpos.numpy()[0]
+
+ # sensors
+ result.sensordata[:] = d.sensordata.numpy()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py
new file mode 100644
index 00000000..33503268
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py
@@ -0,0 +1,354 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 io functions."""
+
+import dataclasses
+import typing
+from typing import Any, Dict, Optional, Union
+
+import mujoco
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+
+import mujoco_warp as mjwarp
+
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+
+
+def _dims_match(test_obj, d1: Any, d2: Any, prefix: str = ""):
+ """Checks that two dataclasses have fields with the same leading dims."""
+ fields1, fields2 = dataclasses.fields(d1), dataclasses.fields(d2)
+ for f1, f2 in zip(fields1, fields2):
+ full_name = prefix + f1.name
+ a1, a2 = getattr(d1, f1.name), getattr(d2, f2.name)
+ if dataclasses.is_dataclass(a1) or dataclasses.is_dataclass(a2):
+ _dims_match(test_obj, a1, a2, prefix + f1.name + ".")
+ continue
+
+ if isinstance(f1.type, wp.types.array) or isinstance(f2.type, wp.types.array):
+ s1, s2 = a1.shape, a2.shape
+ test_obj.assertEqual(len(s1), len(s2), full_name + f" dims mismatch. Got {s1} and {s2}.")
+ test_obj.assertSequenceAlmostEqual(s1, s2, full_name + f" dims mismatch. Got {s1} and {s2}.")
+
+
+def _get_np_scalar_type(val: Any) -> Optional[Union[bool, int, float]]:
+ """Returns the python type from a numpy scalar."""
+ is_np_scalar = list(isinstance(val, t) for t in (np.integer, np.floating, np.bool_))
+ if any(is_np_scalar):
+ return [int, float, bool][is_np_scalar.index(True)]
+
+
+def _check_type_matches_annotation(test_obj, obj: Any, prefix: str = ""):
+ """Checks that dataclass annotations match the runtime types."""
+ assert dataclasses.is_dataclass(obj), prefix + " must be dataclass."
+ msg = "Type of {val_type} does not match annotation {type_} for field {prefix}{field_name}"
+
+ for field in dataclasses.fields(obj):
+ field_name = field.name
+ val = getattr(obj, field_name)
+ val_type = type(val)
+ type_ = field.type
+
+ if dataclasses.is_dataclass(val):
+ test_obj.assertTrue(dataclasses.is_dataclass(type_), msg.format(**locals()))
+ _check_type_matches_annotation(test_obj, val, prefix + field_name + ".")
+ continue
+
+ np_scalar_type = _get_np_scalar_type(val)
+ if np_scalar_type:
+ test_obj.assertIsInstance(np_scalar_type(val), type_, msg.format(**locals()))
+ continue
+
+ if isinstance(type_, wp.types.array):
+ test_obj.assertIsInstance(val, wp.types.array, msg.format(**locals()))
+ continue
+
+ origin_type = typing.get_origin(type_)
+ if tuple in (val_type, origin_type):
+ test_obj.assertEqual(val_type, origin_type, msg.format(**locals()))
+ field_name += ".tuple[]"
+ type_ = typing.get_args(type_)[0]
+
+ items = val
+ for val in items:
+ val_type = type(val)
+ if dataclasses.is_dataclass(val):
+ _check_type_matches_annotation(test_obj, val, prefix + field_name)
+ continue
+
+ np_scalar_type = _get_np_scalar_type(val)
+ if np_scalar_type:
+ test_obj.assertIsInstance(np_scalar_type(val), type_, msg.format(**locals()))
+ continue
+
+ if isinstance(type_, wp.types.array):
+ test_obj.assertIsInstance(val, wp.types.array, msg.format(**locals()))
+ continue
+
+ test_obj.assertEqual(type(val), type_, msg.format(**locals()))
+ continue
+
+ test_obj.assertEqual(type(val), field.type, msg.format(**locals()))
+
+
+def _check_annotation_compat(
+ annotations: Dict[str, Any], prefix: str = "", in_cls: bool = False, in_tuple: bool = False
+) -> Dict[str, Any]:
+ """Checks that dataclass annotations match criteria for JAX API compat."""
+ for k, v in annotations.items():
+ full_key = f"{prefix}{k}"
+ info = f"Found {v} for annotation {full_key}."
+
+ if v in (int, bool, float):
+ continue
+
+ if isinstance(v, wp.types.array):
+ continue
+
+ if v in wp.types.vector_types:
+ raise AssertionError(f"Vector types are not allowed. {info}")
+
+ if typing.get_origin(v) == tuple and (in_cls or in_tuple):
+ raise AssertionError(f"Nested args in Model/Data must not be tuple. {info}")
+
+ if typing.get_origin(v) == tuple:
+ tuple_args = typing.get_args(v)
+ if len(tuple_args) != 2 and tuple_args[1] != ...:
+ raise AssertionError(f"Tuple args must be variadic. {info}")
+
+ _check_annotation_compat(
+ {"[]": tuple_args[0]},
+ prefix=f"{full_key}.tuple",
+ in_cls=in_cls,
+ in_tuple=True,
+ )
+ continue
+
+ if hasattr(v, "__class__") and in_cls:
+ raise AssertionError(f"Nested object args in Model/Data are not allowed. {info}")
+
+ if hasattr(v, "__class__") and not dataclasses.is_dataclass(v):
+ raise AssertionError(f"Args that are objects must be dataclass. {info}")
+
+ if hasattr(v, "__class__") and not v.__module__.startswith("mujoco_warp"):
+ raise AssertionError(f"dataclass args must be within the mujoco_warp module. {info}")
+
+ if hasattr(v, "__class__"):
+ _check_annotation_compat(v.__annotations__, prefix=f"{full_key}{v.__name__}.", in_cls=True, in_tuple=in_tuple)
+ continue
+
+ raise AssertionError(f"Model/Data annotation is not allowed. {info}")
+
+
+def _leading_dims_scale_w_nworld(test_obj, d1: Any, d2: Any, nworld1: int, nworld2: int, prefix: str = ""):
+ """Checks that dataclass fields that scale with nworld have leading dim nworld."""
+ msg = "Arrays that scale with nworld should have leading dim nworld."
+ fields1, fields2 = dataclasses.fields(d1), dataclasses.fields(d2)
+ for f1, f2 in zip(fields1, fields2):
+ full_name = prefix + f1.name
+ a1, a2 = getattr(d1, f1.name), getattr(d2, f2.name)
+ if dataclasses.is_dataclass(a1) or dataclasses.is_dataclass(a2):
+ _leading_dims_scale_w_nworld(test_obj, a1, a2, nworld1, nworld2, prefix + f1.name + ".")
+ continue
+
+ if isinstance(f1.type, wp.types.array) or isinstance(f2.type, wp.types.array):
+ s1, s2 = a1.shape[0], a2.shape[0]
+ if s1 == s2:
+ continue
+ test_obj.assertEqual(s2, nworld2, full_name + f" has leading dim {s2} with nworld={nworld2}. {msg}")
+ test_obj.assertEqual(s1, nworld1, full_name + f" has leading dim {s1} with nworld={nworld1}. {msg}")
+
+
+class IOTest(absltest.TestCase):
+ def test_make_put_data(self):
+ """Tests that make_data and put_data are producing the same shapes for all arrays."""
+ mjm, _, _, d = test_util.fixture("pendula.xml")
+ md = mjwarp.make_data(mjm, nconmax=512, njmax=512)
+
+ # same number of fields
+ self.assertEqual(len(d.__dict__), len(md.__dict__))
+
+ # test shapes for all arrays
+ for attr, val in md.__dict__.items():
+ if isinstance(val, wp.array):
+ self.assertEqual(val.shape, getattr(d, attr).shape, f"{attr} shape mismatch")
+
+ # TODO(team): sensors
+
+ def test_get_data_into_m(self):
+ mjm = mujoco.MjModel.from_xml_string("""
+
+
+
+
+
+
+
+
+
+
+
+
+ """)
+
+ mjd = mujoco.MjData(mjm)
+ mujoco.mj_forward(mjm, mjd)
+
+ mjd_ref = mujoco.MjData(mjm)
+ mujoco.mj_forward(mjm, mjd_ref)
+
+ m = mjwarp.put_model(mjm)
+ d = mjwarp.put_data(mjm, mjd)
+
+ mjd.qLD.fill(-123)
+ mjd.qM.fill(-123)
+
+ mjwarp.get_data_into(mjd, mjm, d)
+ np.testing.assert_allclose(mjd.qLD, mjd_ref.qLD)
+ np.testing.assert_allclose(mjd.qM, mjd_ref.qM)
+
+ def test_ellipsoid_fluid_model(self):
+ with self.assertRaises(NotImplementedError):
+ mjm = mujoco.MjModel.from_xml_string(
+ """
+
+
+
+
+
+
+
+
+
+ """
+ )
+ mjwarp.put_model(mjm)
+
+ def test_jacobian_auto(self):
+ mjm = mujoco.MjModel.from_xml_string("""
+
+
+
+
+
+
+
+
+
+
+
+ """)
+ mjwarp.put_model(mjm)
+
+ def test_put_data_qLD(self):
+ mjm = mujoco.MjModel.from_xml_string("""
+
+
+
+
+
+
+
+
+ """)
+ mjd = mujoco.MjData(mjm)
+ d = mjwarp.put_data(mjm, mjd)
+ self.assertTrue((d.qLD.numpy() == 0.0).all())
+
+ mujoco.mj_forward(mjm, mjd)
+ mjd.qM[:] = 0.0
+ d = mjwarp.put_data(mjm, mjd)
+ self.assertTrue((d.qLD.numpy() == 0.0).all())
+
+ mujoco.mj_forward(mjm, mjd)
+ mjd.qLD[:] = 0.0
+ d = mjwarp.put_data(mjm, mjd)
+ self.assertTrue((d.qLD.numpy() == 0.0).all())
+
+ def test_noslip_solver(self):
+ with self.assertRaises(NotImplementedError):
+ test_util.fixture(
+ xml="""
+
+
+
+ """
+ )
+
+ def test_put_model_nworld_array(self):
+ """Tests that put_model arrays with nworld leading dim have `_is_batched`."""
+ mjm, *_ = test_util.fixture("pendula.xml")
+ m1 = mjwarp.put_model(mjm)
+
+ self.assertTrue(hasattr(m1.geom_pos, "_is_batched"))
+ self.assertEqual(m1.geom_pos.shape[0], 1)
+ self.assertEqual(m1.geom_pos.strides[0], 0)
+ self.assertLen(m1.geom_pos.strides, m1.geom_pos.ndim)
+ self.assertTrue(hasattr(m1.opt.gravity, "_is_batched"))
+ self.assertEqual(m1.opt.gravity.shape[0], 1)
+ self.assertEqual(m1.opt.gravity.strides[0], 0)
+ self.assertLen(m1.opt.gravity.strides, m1.opt.gravity.ndim)
+ self.assertFalse(hasattr(m1.body_parentid, "_is_batched"))
+ self.assertGreater(m1.body_parentid.shape[0], 0)
+ self.assertGreater(m1.body_parentid.strides[0], 0)
+ self.assertLen(m1.body_parentid.strides, m1.body_parentid.ndim)
+
+ def test_put_data_nworld_array(self):
+ """Tests that put_data arrays that scale with nworld have leading dim nworld."""
+ mjm, mjd, _, _ = test_util.fixture("pendula.xml")
+ d1 = mjwarp.put_data(mjm, mjd, nworld=1, nconmax=1_000, njmax=1_000)
+ dn = mjwarp.put_data(mjm, mjd, nworld=133, nconmax=1_000, njmax=1_000)
+ _leading_dims_scale_w_nworld(self, d1, dn, 1, 133)
+
+ def test_make_data_nworld_array(self):
+ """Tests that make_data arrays that scale with nworld have leading dim nworld."""
+ mjm, *_ = test_util.fixture("pendula.xml")
+ d1 = mjwarp.make_data(mjm, nworld=1, nconmax=1_000, njmax=1_000)
+ dn = mjwarp.make_data(mjm, nworld=133, nconmax=1_000, njmax=1_000)
+ _leading_dims_scale_w_nworld(self, d1, dn, 1, 133)
+
+ def test_public_api_jax_compat(self):
+ """Tests that annotations meet a set of criteria for JAX compat."""
+ _check_annotation_compat(mjwarp.Model.__annotations__, "Model.")
+ _check_annotation_compat(mjwarp.Data.__annotations__, "Data.")
+
+ def test_types_match_annotations(self):
+ """Tests that the types of dataclass fields match the annotations."""
+ mjm, _, m, d = test_util.fixture("pendula.xml")
+
+ _check_type_matches_annotation(self, m, "Model.")
+ _check_type_matches_annotation(self, d, "Data.")
+
+ d = mjwarp.make_data(mjm, nworld=2)
+ _check_type_matches_annotation(self, d, "Data.")
+
+ def test_make_put_data_dims_match(self):
+ """Tests that make_data and put_data have matching dimensions."""
+ mjm, mjd, _, _ = test_util.fixture("pendula.xml")
+ dm2 = mjwarp.make_data(mjm, nworld=2, nconmax=13, njmax=42)
+ dm3 = mjwarp.make_data(mjm, nworld=3, nconmax=13, njmax=42)
+
+ dp2 = mjwarp.put_data(mjm, mjd, nworld=2, nconmax=13, njmax=42)
+ dp3 = mjwarp.put_data(mjm, mjd, nworld=3, nconmax=13, njmax=42)
+
+ _dims_match(self, dm2, dp2)
+ _dims_match(self, dm3, dp3)
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/jax_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/jax_test.py
new file mode 100644
index 00000000..529a2288
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/jax_test.py
@@ -0,0 +1,101 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import os
+
+import warp as wp
+from absl.testing import absltest
+from absl.testing import parameterized
+
+import mujoco_warp as mjwarp
+from mujoco.mjx.third_party.mujoco_warp._src.test_util import fixture
+
+# TODO(team): JAX test is temporary, remove after we land MJX:Warp
+
+
+class JAXTest(parameterized.TestCase):
+ @parameterized.parameters("humanoid/humanoid.xml", "pendula.xml")
+ def test_jax(self, xml):
+ os.environ["XLA_FLAGS"] = "--xla_gpu_graph_min_graph_size=1"
+ # Force JAX to allocate memory on demand and deallocate when not needed (slow)
+ os.environ["XLA_PYTHON_CLIENT_ALLOCATOR"] = "platform"
+
+ try:
+ import jax
+ except ImportError:
+ self.skipTest("JAX not installed")
+
+ from jax import numpy as jp
+ from warp.jax_experimental.ffi import jax_callable
+
+ if jax.default_backend() != "gpu":
+ self.skipTest("JAX default backend is not GPU")
+
+ NWORLDS = 2
+ NCONTACTS = 16
+ UNROLL_LENGTH = 1
+
+ mjm, _, m, d = fixture(
+ xml,
+ nworld=NWORLDS,
+ nconmax=NWORLDS * NCONTACTS,
+ njmax=NWORLDS * NCONTACTS * 4,
+ iterations=1,
+ ls_iterations=4,
+ kick=True,
+ )
+
+ # Disable CUDA graph conditional
+ m.opt.graph_conditional = False
+
+ def warp_step(
+ qpos_in: wp.array(dtype=wp.float32, ndim=2),
+ qvel_in: wp.array(dtype=wp.float32, ndim=2),
+ qpos_out: wp.array(dtype=wp.float32, ndim=2),
+ qvel_out: wp.array(dtype=wp.float32, ndim=2),
+ ):
+ wp.copy(d.qpos, qpos_in)
+ wp.copy(d.qvel, qvel_in)
+ mjwarp.step(m, d)
+ wp.copy(qpos_out, d.qpos)
+ wp.copy(qvel_out, d.qvel)
+
+ def unroll(qpos, qvel):
+ def step(carry, _):
+ qpos, qvel = carry
+ qpos, qvel = warp_step_fn(qpos, qvel)
+ return (qpos, qvel), None
+
+ (qpos, qvel), _ = jax.lax.scan(step, (qpos, qvel), length=UNROLL_LENGTH)
+
+ return qpos, qvel
+
+ warp_step_fn = jax_callable(
+ warp_step,
+ num_outputs=2,
+ output_dims={"qpos_out": (NWORLDS, mjm.nq), "qvel_out": (NWORLDS, mjm.nv)},
+ graph_compatible=True,
+ )
+
+ jax_qpos = jp.tile(jp.array(m.qpos0.numpy()), (NWORLDS, 1))
+ jax_qvel = jp.zeros((NWORLDS, m.nv))
+
+ jax_unroll_fn = jax.jit(unroll).lower(jax_qpos, jax_qvel).compile()
+ jax_unroll_fn(jax_qpos, jax_qvel)
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py
new file mode 100644
index 00000000..99cc0b20
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py
@@ -0,0 +1,289 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+from typing import Any, Tuple
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src import types
+
+
+@wp.func
+def mul_quat(u: wp.quat, v: wp.quat) -> wp.quat:
+ return wp.quat(
+ u[0] * v[0] - u[1] * v[1] - u[2] * v[2] - u[3] * v[3],
+ u[0] * v[1] + u[1] * v[0] + u[2] * v[3] - u[3] * v[2],
+ u[0] * v[2] - u[1] * v[3] + u[2] * v[0] + u[3] * v[1],
+ u[0] * v[3] + u[1] * v[2] - u[2] * v[1] + u[3] * v[0],
+ )
+
+
+@wp.func
+def quat_mul_axis(q: wp.quat, axis: wp.vec3f) -> wp.quat:
+ """Multiplies a quaternion and an axis."""
+ return wp.quat(
+ -q[1] * axis[0] - q[2] * axis[1] - q[3] * axis[2],
+ q[0] * axis[0] + q[2] * axis[2] - q[3] * axis[1],
+ q[0] * axis[1] + q[3] * axis[0] - q[1] * axis[2],
+ q[0] * axis[2] + q[1] * axis[1] - q[2] * axis[0],
+ )
+
+
+@wp.func
+def rot_vec_quat(vec: wp.vec3, quat: wp.quat) -> wp.vec3:
+ s, u = quat[0], wp.vec3(quat[1], quat[2], quat[3])
+ r = 2.0 * (wp.dot(u, vec) * u) + (s * s - wp.dot(u, u)) * vec
+ r = r + 2.0 * s * wp.cross(u, vec)
+ return r
+
+
+@wp.func
+def axis_angle_to_quat(axis: wp.vec3, angle: float) -> wp.quat:
+ s, c = wp.sin(angle * 0.5), wp.cos(angle * 0.5)
+ axis = axis * s
+ return wp.quat(c, axis[0], axis[1], axis[2])
+
+
+@wp.func
+def quat_to_mat(quat: wp.quat) -> wp.mat33:
+ """Converts a quaternion into a 9-dimensional rotation matrix."""
+ vec = wp.vec4(quat[0], quat[1], quat[2], quat[3])
+ q = wp.outer(vec, vec)
+
+ return wp.mat33(
+ q[0, 0] + q[1, 1] - q[2, 2] - q[3, 3],
+ 2.0 * (q[1, 2] - q[0, 3]),
+ 2.0 * (q[1, 3] + q[0, 2]),
+ 2.0 * (q[1, 2] + q[0, 3]),
+ q[0, 0] - q[1, 1] + q[2, 2] - q[3, 3],
+ 2.0 * (q[2, 3] - q[0, 1]),
+ 2.0 * (q[1, 3] - q[0, 2]),
+ 2.0 * (q[2, 3] + q[0, 1]),
+ q[0, 0] - q[1, 1] - q[2, 2] + q[3, 3],
+ )
+
+
+@wp.func
+def quat_inv(quat: wp.quat) -> wp.quat:
+ return wp.quat(quat[0], -quat[1], -quat[2], -quat[3])
+
+
+@wp.func
+def inert_vec(i: types.vec10, v: wp.spatial_vector) -> wp.spatial_vector:
+ """mju_mulInertVec: multiply 6D vector (rotation, translation) by 6D inertia matrix."""
+ return wp.spatial_vector(
+ i[0] * v[0] + i[3] * v[1] + i[4] * v[2] - i[8] * v[4] + i[7] * v[5],
+ i[3] * v[0] + i[1] * v[1] + i[5] * v[2] + i[8] * v[3] - i[6] * v[5],
+ i[4] * v[0] + i[5] * v[1] + i[2] * v[2] - i[7] * v[3] + i[6] * v[4],
+ i[8] * v[1] - i[7] * v[2] + i[9] * v[3],
+ i[6] * v[2] - i[8] * v[0] + i[9] * v[4],
+ i[7] * v[0] - i[6] * v[1] + i[9] * v[5],
+ )
+
+
+@wp.func
+def motion_cross(u: wp.spatial_vector, v: wp.spatial_vector) -> wp.spatial_vector:
+ """Cross product of two motions."""
+
+ u0 = wp.vec3(u[0], u[1], u[2])
+ u1 = wp.vec3(u[3], u[4], u[5])
+ v0 = wp.vec3(v[0], v[1], v[2])
+ v1 = wp.vec3(v[3], v[4], v[5])
+
+ ang = wp.cross(u0, v0)
+ vel = wp.cross(u1, v0) + wp.cross(u0, v1)
+
+ return wp.spatial_vector(ang, vel)
+
+
+@wp.func
+def motion_cross_force(v: wp.spatial_vector, f: wp.spatial_vector) -> wp.spatial_vector:
+ """Cross product of a motion and a force."""
+
+ v0 = wp.vec3(v[0], v[1], v[2])
+ v1 = wp.vec3(v[3], v[4], v[5])
+ f0 = wp.vec3(f[0], f[1], f[2])
+ f1 = wp.vec3(f[3], f[4], f[5])
+
+ ang = wp.cross(v0, f0) + wp.cross(v1, f1)
+ vel = wp.cross(v0, f1)
+
+ return wp.spatial_vector(ang, vel)
+
+
+@wp.func
+def quat_to_vel(quat: wp.quat) -> wp.vec3:
+ axis = wp.vec3(quat[1], quat[2], quat[3])
+ sin_a_2 = wp.norm_l2(axis)
+
+ if sin_a_2 == 0.0:
+ return wp.vec3(0.0)
+
+ speed = 2.0 * wp.atan2(sin_a_2, quat[0])
+ # when axis-angle is larger than pi, rotation is in the opposite direction
+ if speed > wp.pi:
+ speed -= 2.0 * wp.pi
+
+ return axis * speed / sin_a_2
+
+
+@wp.func
+def quat_sub(qa: wp.quat, qb: wp.quat) -> wp.vec3:
+ """Subtract quaternions, express as 3D velocity: qb*quat(res) = qa."""
+ # qdif = neg(qb)*qa
+ qneg = wp.quat(qb[0], -qb[1], -qb[2], -qb[3])
+ qdif = mul_quat(qneg, qa)
+
+ # convert to 3D velocity
+ return quat_to_vel(qdif)
+
+
+@wp.func
+def quat_integrate(q: wp.quat, v: wp.vec3, dt: float) -> wp.quat:
+ """Integrates a quaternion given angular velocity and dt."""
+ norm_ = wp.length(v)
+ v = wp.normalize(v) # does that need proper zero gradient handling?
+ angle = dt * norm_
+
+ q_res = axis_angle_to_quat(v, angle)
+ q = wp.normalize(q)
+ q_res = mul_quat(q, q_res)
+
+ return wp.normalize(q_res)
+
+
+@wp.func
+def orthogonals(a: wp.vec3):
+ y = wp.vec3(0.0, 1.0, 0.0)
+ z = wp.vec3(0.0, 0.0, 1.0)
+ b = wp.where((-0.5 < a[1]) and (a[1] < 0.5), y, z)
+ b = b - a * wp.dot(a, b)
+ b = wp.normalize(b)
+ if wp.length(a) == 0.0:
+ b = wp.vec3(0.0, 0.0, 0.0)
+ c = wp.cross(a, b)
+
+ return b, c
+
+
+@wp.func
+def orthonormal(normal: wp.vec3) -> wp.vec3:
+ if wp.abs(normal[0]) < wp.abs(normal[1]) and wp.abs(normal[0]) < wp.abs(normal[2]):
+ dir = wp.vec3(1.0 - normal[0] * normal[0], -normal[0] * normal[1], -normal[0] * normal[2])
+ elif wp.abs(normal[1]) < wp.abs(normal[2]):
+ dir = wp.vec3(-normal[1] * normal[0], 1.0 - normal[1] * normal[1], -normal[1] * normal[2])
+ else:
+ dir = wp.vec3(-normal[2] * normal[0], -normal[2] * normal[1], 1.0 - normal[2] * normal[2])
+ dir, _ = gjk_normalize(dir)
+ return dir
+
+
+@wp.func
+def gjk_normalize(a: wp.vec3):
+ norm = wp.length(a)
+ if norm > 1e-8 and norm < 1e12:
+ return a / norm, True
+ return a, False
+
+
+@wp.func
+def make_frame(a: wp.vec3):
+ a = wp.normalize(a)
+ b, c = orthogonals(a)
+
+ # fmt: off
+ return wp.mat33(
+ a.x, a.y, a.z,
+ b.x, b.y, b.z,
+ c.x, c.y, c.z
+ )
+ # fmt: on
+
+
+@wp.func
+def normalize_with_norm(x: Any):
+ norm = wp.length(x)
+ if norm == 0.0:
+ return x, 0.0
+ return x / norm, norm
+
+
+@wp.func
+def closest_segment_point(a: wp.vec3, b: wp.vec3, pt: wp.vec3) -> wp.vec3:
+ """Returns the closest point on the a-b line segment to a point pt."""
+ ab = b - a
+ t = wp.dot(pt - a, ab) / (wp.dot(ab, ab) + 1e-6)
+ return a + wp.clamp(t, 0.0, 1.0) * ab
+
+
+@wp.func
+def closest_segment_point_and_dist(a: wp.vec3, b: wp.vec3, pt: wp.vec3) -> Tuple[wp.vec3, float]:
+ """Returns closest point on the line segment and the distance squared."""
+ closest = closest_segment_point(a, b, pt)
+ dist = wp.dot((pt - closest), (pt - closest))
+ return closest, dist
+
+
+@wp.func
+def closest_segment_to_segment_points(a0: wp.vec3, a1: wp.vec3, b0: wp.vec3, b1: wp.vec3) -> Tuple[wp.vec3, wp.vec3]:
+ """Returns closest points between two line segments."""
+
+ dir_a, len_a = normalize_with_norm(a1 - a0)
+ dir_b, len_b = normalize_with_norm(b1 - b0)
+
+ half_len_a = len_a * 0.5
+ half_len_b = len_b * 0.5
+ a_mid = a0 + dir_a * half_len_a
+ b_mid = b0 + dir_b * half_len_b
+
+ trans = a_mid - b_mid
+
+ dira_dot_dirb = wp.dot(dir_a, dir_b)
+ dira_dot_trans = wp.dot(dir_a, trans)
+ dirb_dot_trans = wp.dot(dir_b, trans)
+ denom = 1.0 - dira_dot_dirb * dira_dot_dirb
+
+ orig_t_a = (-dira_dot_trans + dira_dot_dirb * dirb_dot_trans) / (denom + 1e-6)
+ orig_t_b = dirb_dot_trans + orig_t_a * dira_dot_dirb
+ t_a = wp.clamp(orig_t_a, -half_len_a, half_len_a)
+ t_b = wp.clamp(orig_t_b, -half_len_b, half_len_b)
+
+ best_a = a_mid + dir_a * t_a
+ best_b = b_mid + dir_b * t_b
+
+ new_a, d1 = closest_segment_point_and_dist(a0, a1, best_b)
+ new_b, d2 = closest_segment_point_and_dist(b0, b1, best_a)
+ if d1 < d2:
+ return new_a, best_b
+ return best_a, new_b
+
+
+@wp.func
+def safe_div(x: float, y: float) -> float:
+ return x / wp.where(y != 0.0, y, types.MJ_MINVAL)
+
+
+@wp.func
+def upper_tri_index(n: int, i: int, j: int) -> int:
+ """Returns index of a_ij = a_ji in upper triangular matrix (excluding diagonal)."""
+ return (i * (2 * n - i - 3)) // 2 + j - 1
+
+
+@wp.func
+def upper_trid_index(n: int, i: int, j: int) -> int:
+ """Returns index of a_ij = a_ji in upper triangular matrix (including diagonal)."""
+ if j < i:
+ i, j = j, i
+ return (i * (2 * n - i - 1)) // 2 + j
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math_test.py
new file mode 100644
index 00000000..3f1de2c7
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math_test.py
@@ -0,0 +1,131 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import warp as wp
+from absl.testing import absltest
+
+from mujoco.mjx.third_party.mujoco_warp._src.math import closest_segment_to_segment_points
+from mujoco.mjx.third_party.mujoco_warp._src.math import upper_tri_index
+from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index
+
+
+class ClosestSegmentSegmentPointsTest(absltest.TestCase):
+ """Tests for closest segment-to-segment points."""
+
+ def test_closest_segments_points(self):
+ """Test closest points between two segments."""
+ a0 = wp.vec3([0.73432405, 0.12372768, 0.20272314])
+ a1 = wp.vec3([1.10600128, 0.88555209, 0.65209485])
+ b0 = wp.vec3([0.85599262, 0.61736299, 0.9843583])
+ b1 = wp.vec3([1.84270939, 0.92891793, 1.36343326])
+
+ best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1)
+ self.assertSequenceAlmostEqual(best_a, [1.09063, 0.85404, 0.63351], 5)
+ self.assertSequenceAlmostEqual(best_b, [0.99596, 0.66156, 1.03813], 5)
+
+ def test_intersecting_segments(self):
+ """Tests segments that intersect."""
+ a0, a1 = wp.vec3([0.0, 0.0, -1.0]), wp.vec3([0.0, 0.0, 1.0])
+ b0, b1 = wp.vec3([-1.0, 0.0, 0.0]), wp.vec3([1.0, 0.0, 0.0])
+
+ best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1)
+ self.assertSequenceAlmostEqual(best_a, [0.0, 0.0, 0.0], 5)
+ self.assertSequenceAlmostEqual(best_b, [0.0, 0.0, 0.0], 5)
+
+ def test_intersecting_lines(self):
+ """Tests that intersecting lines get clipped."""
+ a0, a1 = wp.vec3([0.2, 0.2, 0.0]), wp.vec3([1.0, 1.0, 0.0])
+ b0, b1 = wp.vec3([0.2, 0.4, 0.0]), wp.vec3([1.0, 2.0, 0.0])
+
+ best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1)
+ self.assertSequenceAlmostEqual(best_a, [0.3, 0.3, 0.0], 2)
+ self.assertSequenceAlmostEqual(best_b, [0.2, 0.4, 0.0], 2)
+
+ def test_parallel_segments(self):
+ """Tests that parallel segments have closest points at the midpoint."""
+ a0, a1 = wp.vec3([0.0, 0.0, -1.0]), wp.vec3([0.0, 0.0, 1.0])
+ b0, b1 = wp.vec3([1.0, 0.0, -1.0]), wp.vec3([1.0, 0.0, 1.0])
+
+ best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1)
+ self.assertSequenceAlmostEqual(best_a, [0.0, 0.0, 0.0], 5)
+ self.assertSequenceAlmostEqual(best_b, [1.0, 0.0, 0.0], 5)
+
+ def test_parallel_offset_segments(self):
+ """Tests that offset parallel segments are close at segment endpoints."""
+ a0, a1 = wp.vec3([0.0, 0.0, -1.0]), wp.vec3([0.0, 0.0, 1.0])
+ b0, b1 = wp.vec3([1.0, 0.0, 1.0]), wp.vec3([1.0, 0.0, 3.0])
+
+ best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1)
+ self.assertSequenceAlmostEqual(best_a, [0.0, 0.0, 1.0], 5)
+ self.assertSequenceAlmostEqual(best_b, [1.0, 0.0, 1.0], 5)
+
+ def test_zero_length_segments(self):
+ """Test that zero length segments don't return NaNs."""
+ a0, a1 = wp.vec3([0.0, 0.0, -1.0]), wp.vec3([0.0, 0.0, -1.0])
+ b0, b1 = wp.vec3([1.0, 0.0, 0.1]), wp.vec3([1.0, 0.0, 0.1])
+
+ best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1)
+ self.assertSequenceAlmostEqual(best_a, [0.0, 0.0, -1.0], 5)
+ self.assertSequenceAlmostEqual(best_b, [1.0, 0.0, 0.1], 5)
+
+ def test_overlapping_segments(self):
+ """Tests that perfectly overlapping segments intersect at the midpoints."""
+ a0, a1 = wp.vec3([0.0, 0.0, -1.0]), wp.vec3([0.0, 0.0, 1.0])
+ b0, b1 = wp.vec3([0.0, 0.0, -1.0]), wp.vec3([0.0, 0.0, 1.0])
+
+ best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1)
+ self.assertSequenceAlmostEqual(best_a, [0.0, 0.0, 0.0], 5)
+ self.assertSequenceAlmostEqual(best_b, [0.0, 0.0, 0.0], 5)
+
+ def test_upper_tri_index2(self):
+ """Tests upper_tri_index with size 2"""
+ arr = []
+ for i in range(2):
+ for j in range(i + 1, 2):
+ arr.append(upper_tri_index(2, i, j))
+ self.assertEqual(arr, list(range(0, 1)))
+
+ def test_upper_tri_index10(self):
+ """Tests upper_tri_index with size 10"""
+ arr = []
+ for i in range(10):
+ for j in range(i + 1, 10):
+ arr.append(upper_tri_index(10, i, j))
+ self.assertEqual(arr, list(range(0, 45)))
+
+ def test_upper_trid_index1(self):
+ """Tests upper_trid_index with size 1"""
+ arr = []
+ for i in range(1):
+ for j in range(i, 1):
+ arr.append(upper_trid_index(1, i, j))
+ self.assertEqual(arr, list(range(0, 1)))
+
+ def test_upper_trid_index10(self):
+ """Tests upper_trid_index with size 10"""
+ arr = []
+ for i in range(10):
+ for j in range(i, 10):
+ arr.append(upper_trid_index(10, i, j))
+ self.assertEqual(arr, list(range(0, 55)))
+
+ def test_upper_trid_index10(self):
+ """Tests upper_trid_index works with symmetric matrix"""
+ self.assertEqual(upper_trid_index(10, 1, 5), upper_trid_index(10, 5, 1))
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py
new file mode 100644
index 00000000..72b00d90
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py
@@ -0,0 +1,617 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src import math
+from mujoco.mjx.third_party.mujoco_warp._src import support
+from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import JointType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
+
+
+@wp.kernel
+def _spring_damper_dof_passive(
+ # Model:
+ qpos_spring: wp.array2d(dtype=float),
+ jnt_type: wp.array(dtype=int),
+ jnt_qposadr: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ jnt_stiffness: wp.array2d(dtype=float),
+ dof_damping: wp.array2d(dtype=float),
+ # Data in:
+ qpos_in: wp.array2d(dtype=float),
+ qvel_in: wp.array2d(dtype=float),
+ # Data out:
+ qfrc_spring_out: wp.array2d(dtype=float),
+ qfrc_damper_out: wp.array2d(dtype=float),
+):
+ worldid, jntid = wp.tid()
+ dofid = jnt_dofadr[jntid]
+ stiffness = jnt_stiffness[worldid, jntid]
+ damping = dof_damping[worldid, dofid]
+
+ has_stiffness = stiffness != 0.0
+ has_damping = damping != 0.0
+
+ if not has_stiffness:
+ qfrc_spring_out[worldid, dofid] = 0.0
+
+ if not has_damping:
+ qfrc_damper_out[worldid, dofid] = 0.0
+
+ if not (has_stiffness or has_damping):
+ return
+
+ jnttype = jnt_type[jntid]
+ qposid = jnt_qposadr[jntid]
+
+ if jnttype == wp.static(JointType.FREE.value):
+ # spring
+ if has_stiffness:
+ dif = wp.vec3(
+ qpos_in[worldid, qposid + 0] - qpos_spring[worldid, qposid + 0],
+ qpos_in[worldid, qposid + 1] - qpos_spring[worldid, qposid + 1],
+ qpos_in[worldid, qposid + 2] - qpos_spring[worldid, qposid + 2],
+ )
+ qfrc_spring_out[worldid, dofid + 0] = -stiffness * dif[0]
+ qfrc_spring_out[worldid, dofid + 1] = -stiffness * dif[1]
+ qfrc_spring_out[worldid, dofid + 2] = -stiffness * dif[2]
+ rot = wp.quat(
+ qpos_in[worldid, qposid + 3],
+ qpos_in[worldid, qposid + 4],
+ qpos_in[worldid, qposid + 5],
+ qpos_in[worldid, qposid + 6],
+ )
+ rot = wp.normalize(rot)
+ ref = wp.quat(
+ qpos_spring[worldid, qposid + 3],
+ qpos_spring[worldid, qposid + 4],
+ qpos_spring[worldid, qposid + 5],
+ qpos_spring[worldid, qposid + 6],
+ )
+ dif = math.quat_sub(rot, ref)
+ qfrc_spring_out[worldid, dofid + 3] = -stiffness * dif[0]
+ qfrc_spring_out[worldid, dofid + 4] = -stiffness * dif[1]
+ qfrc_spring_out[worldid, dofid + 5] = -stiffness * dif[2]
+
+ # damper
+ if has_damping:
+ qfrc_damper_out[worldid, dofid + 0] = -damping * qvel_in[worldid, dofid + 0]
+ qfrc_damper_out[worldid, dofid + 1] = -damping * qvel_in[worldid, dofid + 1]
+ qfrc_damper_out[worldid, dofid + 2] = -damping * qvel_in[worldid, dofid + 2]
+ qfrc_damper_out[worldid, dofid + 3] = -damping * qvel_in[worldid, dofid + 3]
+ qfrc_damper_out[worldid, dofid + 4] = -damping * qvel_in[worldid, dofid + 4]
+ qfrc_damper_out[worldid, dofid + 5] = -damping * qvel_in[worldid, dofid + 5]
+ elif jnttype == wp.static(JointType.BALL.value):
+ # spring
+ if has_stiffness:
+ rot = wp.quat(
+ qpos_in[worldid, qposid + 0],
+ qpos_in[worldid, qposid + 1],
+ qpos_in[worldid, qposid + 2],
+ qpos_in[worldid, qposid + 3],
+ )
+ rot = wp.normalize(rot)
+ ref = wp.quat(
+ qpos_spring[worldid, qposid + 0],
+ qpos_spring[worldid, qposid + 1],
+ qpos_spring[worldid, qposid + 2],
+ qpos_spring[worldid, qposid + 3],
+ )
+ dif = math.quat_sub(rot, ref)
+ qfrc_spring_out[worldid, dofid + 0] = -stiffness * dif[0]
+ qfrc_spring_out[worldid, dofid + 1] = -stiffness * dif[1]
+ qfrc_spring_out[worldid, dofid + 2] = -stiffness * dif[2]
+
+ # damper
+ if has_damping:
+ qfrc_damper_out[worldid, dofid + 0] = -damping * qvel_in[worldid, dofid + 0]
+ qfrc_damper_out[worldid, dofid + 1] = -damping * qvel_in[worldid, dofid + 1]
+ qfrc_damper_out[worldid, dofid + 2] = -damping * qvel_in[worldid, dofid + 2]
+ else: # mjJNT_SLIDE, mjJNT_HINGE
+ # spring
+ if has_stiffness:
+ fdif = qpos_in[worldid, qposid] - qpos_spring[worldid, qposid]
+ qfrc_spring_out[worldid, dofid] = -stiffness * fdif
+
+ # damper
+ if has_damping:
+ qfrc_damper_out[worldid, dofid] = -damping * qvel_in[worldid, dofid]
+
+
+@wp.kernel
+def _spring_damper_tendon_passive(
+ # Model:
+ tendon_stiffness: wp.array2d(dtype=float),
+ tendon_damping: wp.array2d(dtype=float),
+ tendon_lengthspring: wp.array2d(dtype=wp.vec2),
+ # Data in:
+ ten_velocity_in: wp.array2d(dtype=float),
+ ten_length_in: wp.array2d(dtype=float),
+ ten_J_in: wp.array3d(dtype=float),
+ # Data out:
+ qfrc_spring_out: wp.array2d(dtype=float),
+ qfrc_damper_out: wp.array2d(dtype=float),
+):
+ worldid, tenid, dofid = wp.tid()
+
+ stiffness = tendon_stiffness[worldid, tenid]
+ damping = tendon_damping[worldid, tenid]
+
+ if stiffness == 0.0 and damping == 0.0:
+ return
+
+ J = ten_J_in[worldid, tenid, dofid]
+
+ if stiffness:
+ # compute spring force along tendon
+ length = ten_length_in[worldid, tenid]
+ lengthspring = tendon_lengthspring[worldid, tenid]
+ lower = lengthspring[0]
+ upper = lengthspring[1]
+
+ if length > upper:
+ frc_spring = stiffness * (upper - length)
+ elif length < lower:
+ frc_spring = stiffness * (lower - length)
+ else:
+ frc_spring = 0.0
+
+ # transform to joint torque
+ wp.atomic_add(qfrc_spring_out[worldid], dofid, J * frc_spring)
+
+ if damping:
+ # compute damper linear force along tendon
+ frc_damper = -damping * ten_velocity_in[worldid, tenid]
+
+ # transform to joint torque
+ wp.atomic_add(qfrc_damper_out[worldid], dofid, J * frc_damper)
+
+
+@wp.kernel
+def _gravity_force(
+ # Model:
+ opt_gravity: wp.array(dtype=wp.vec3),
+ body_parentid: wp.array(dtype=int),
+ body_rootid: wp.array(dtype=int),
+ body_mass: wp.array2d(dtype=float),
+ body_gravcomp: wp.array2d(dtype=float),
+ dof_bodyid: wp.array(dtype=int),
+ # Data in:
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ # Data out:
+ qfrc_gravcomp_out: wp.array2d(dtype=float),
+):
+ worldid, bodyid, dofid = wp.tid()
+ bodyid += 1 # skip world body
+ gravcomp = body_gravcomp[worldid, bodyid]
+ gravity = opt_gravity[worldid]
+
+ if gravcomp:
+ force = -gravity * body_mass[worldid, bodyid] * gravcomp
+
+ pos = xipos_in[worldid, bodyid]
+ jac, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos, bodyid, dofid, worldid)
+
+ wp.atomic_add(qfrc_gravcomp_out[worldid], dofid, wp.dot(jac, force))
+
+
+@wp.kernel
+def _box_fluid(
+ # Model:
+ opt_wind: wp.array(dtype=wp.vec3),
+ opt_density: wp.array(dtype=float),
+ opt_viscosity: wp.array(dtype=float),
+ body_rootid: wp.array(dtype=int),
+ body_mass: wp.array2d(dtype=float),
+ body_inertia: wp.array2d(dtype=wp.vec3),
+ # Data in:
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ ximat_in: wp.array2d(dtype=wp.mat33),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ # Data out:
+ fluid_applied_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ """Fluid forces based on inertia-box approximation."""
+
+ worldid, bodyid = wp.tid()
+ wind = opt_wind[worldid]
+ density = opt_density[worldid]
+ viscosity = opt_viscosity[worldid]
+
+ # map from CoM-centered to local body-centered 6D velocity
+
+ # body-inertial
+ pos = xipos_in[worldid, bodyid]
+ rot = ximat_in[worldid, bodyid]
+ rotT = wp.transpose(rot)
+
+ # transform velocity
+ cvel = cvel_in[worldid, bodyid]
+ torque = wp.spatial_top(cvel)
+ force = wp.spatial_bottom(cvel)
+ subtree_com = subtree_com_in[worldid, body_rootid[bodyid]]
+ dif = pos - subtree_com
+ force -= wp.cross(dif, torque)
+
+ lvel_torque = rotT @ torque
+ lvel_force = rotT @ force
+
+ if wind[0] or wind[1] or wind[2]:
+ # subtract translational component from body velocity
+ lvel_force -= rotT @ wind
+
+ lfrc_torque = wp.vec3(0.0)
+ lfrc_force = wp.vec3(0.0)
+
+ has_viscosity = viscosity > 0.0
+ has_density = density > 0.0
+
+ if has_viscosity or has_density:
+ inertia = body_inertia[worldid, bodyid]
+ mass = body_mass[worldid, bodyid]
+ scl = 6.0 / mass
+ box0 = wp.sqrt(wp.max(MJ_MINVAL, inertia[1] + inertia[2] - inertia[0]) * scl)
+ box1 = wp.sqrt(wp.max(MJ_MINVAL, inertia[0] + inertia[2] - inertia[1]) * scl)
+ box2 = wp.sqrt(wp.max(MJ_MINVAL, inertia[0] + inertia[1] - inertia[2]) * scl)
+
+ if has_viscosity:
+ # diameter of sphere approximation
+ diam = (box0 + box1 + box2) / 3.0
+
+ # angular viscosity
+ lfrc_torque = -lvel_torque * wp.pow(diam, 3.0) * wp.pi * viscosity
+
+ # linear viscosity
+ lfrc_force = -3.0 * lvel_force * diam * wp.pi * viscosity
+
+ if has_density:
+ # force
+ lfrc_force -= wp.vec3(
+ 0.5 * density * box1 * box2 * wp.abs(lvel_force[0]) * lvel_force[0],
+ 0.5 * density * box0 * box2 * wp.abs(lvel_force[1]) * lvel_force[1],
+ 0.5 * density * box0 * box1 * wp.abs(lvel_force[2]) * lvel_force[2],
+ )
+
+ # torque
+ scl = density / 64.0
+ box0_pow4 = wp.pow(box0, 4.0)
+ box1_pow4 = wp.pow(box1, 4.0)
+ box2_pow4 = wp.pow(box2, 4.0)
+ lfrc_torque -= wp.vec3(
+ box0 * (box1_pow4 + box2_pow4) * wp.abs(lvel_torque[0]) * lvel_torque[0] * scl,
+ box1 * (box0_pow4 + box2_pow4) * wp.abs(lvel_torque[1]) * lvel_torque[1] * scl,
+ box2 * (box0_pow4 + box1_pow4) * wp.abs(lvel_torque[2]) * lvel_torque[2] * scl,
+ )
+
+ # rotate to global orientation: lfrc -> bfrc
+ torque = rot @ lfrc_torque
+ force = rot @ lfrc_force
+
+ fluid_applied_out[worldid, bodyid] = wp.spatial_vector(force, torque)
+
+
+def _fluid(m: Model, d: Data):
+ wp.launch(
+ _box_fluid,
+ dim=(d.nworld, m.nbody),
+ inputs=[
+ m.opt.wind,
+ m.opt.density,
+ m.opt.viscosity,
+ m.body_rootid,
+ m.body_mass,
+ m.body_inertia,
+ d.xipos,
+ d.ximat,
+ d.subtree_com,
+ d.cvel,
+ ],
+ outputs=[
+ d.fluid_applied,
+ ],
+ )
+
+ # TODO(team): ellipsoid fluid model
+
+ support.apply_ft(m, d, d.fluid_applied, d.qfrc_fluid, False)
+
+
+@wp.kernel
+def _qfrc_passive(
+ # Model:
+ opt_has_fluid: bool,
+ jnt_actgravcomp: wp.array(dtype=int),
+ dof_jntid: wp.array(dtype=int),
+ # Data in:
+ qfrc_spring_in: wp.array2d(dtype=float),
+ qfrc_damper_in: wp.array2d(dtype=float),
+ qfrc_gravcomp_in: wp.array2d(dtype=float),
+ qfrc_fluid_in: wp.array2d(dtype=float),
+ # In:
+ gravcomp: bool,
+ # Data out:
+ qfrc_passive_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ qfrc_passive = qfrc_spring_in[worldid, dofid]
+ qfrc_passive += qfrc_damper_in[worldid, dofid]
+
+ # add gravcomp unless added by actuators
+ if gravcomp and not jnt_actgravcomp[dof_jntid[dofid]]:
+ qfrc_passive += qfrc_gravcomp_in[worldid, dofid]
+
+ # add fluid force
+ if opt_has_fluid:
+ qfrc_passive += qfrc_fluid_in[worldid, dofid]
+
+ qfrc_passive_out[worldid, dofid] = qfrc_passive
+
+
+@wp.kernel
+def _flex_elasticity(
+ # Model:
+ opt_timestep: wp.array(dtype=float),
+ body_dofadr: wp.array(dtype=int),
+ flex_dim: wp.array(dtype=int),
+ flex_vertadr: wp.array(dtype=int),
+ flex_edgeadr: wp.array(dtype=int),
+ flex_elemedgeadr: wp.array(dtype=int),
+ flex_vertbodyid: wp.array(dtype=int),
+ flex_elem: wp.array(dtype=int),
+ flex_elemedge: wp.array(dtype=int),
+ flexedge_length0: wp.array(dtype=float),
+ flex_stiffness: wp.array(dtype=float),
+ flex_damping: wp.array(dtype=float),
+ # Data in:
+ flexvert_xpos_in: wp.array2d(dtype=wp.vec3),
+ flexedge_length_in: wp.array2d(dtype=float),
+ flexedge_velocity_in: wp.array2d(dtype=float),
+ # Data out:
+ qfrc_spring_out: wp.array2d(dtype=float),
+):
+ worldid, elemid = wp.tid()
+ timestep = opt_timestep[worldid]
+ f = 0 # TODO(quaglino): this should become a function of t
+
+ dim = flex_dim[f]
+ nvert = dim + 1
+ nedge = nvert * (nvert - 1) / 2
+ edges = wp.where(
+ dim == 3,
+ wp.mat(0, 1, 1, 2, 2, 0, 2, 3, 0, 3, 1, 3, shape=(6, 2), dtype=int),
+ wp.mat(1, 2, 2, 0, 0, 1, 0, 0, 0, 0, 0, 0, shape=(6, 2), dtype=int),
+ )
+ kD = flex_damping[f] / timestep
+
+ gradient = wp.mat(0.0, shape=(6, 6))
+ for e in range(nedge):
+ vert0 = flex_elem[(dim + 1) * elemid + edges[e, 0]]
+ vert1 = flex_elem[(dim + 1) * elemid + edges[e, 1]]
+ xpos0 = flexvert_xpos_in[worldid, vert0]
+ xpos1 = flexvert_xpos_in[worldid, vert1]
+ for i in range(3):
+ gradient[e, 0 + i] = xpos0[i] - xpos1[i]
+ gradient[e, 3 + i] = xpos1[i] - xpos0[i]
+
+ elongation = wp.spatial_vectorf(0.0)
+ for e in range(nedge):
+ idx = flex_elemedge[flex_elemedgeadr[f] + elemid * nedge + e]
+ vel = flexedge_velocity_in[worldid, flex_edgeadr[f] + idx]
+ deformed = flexedge_length_in[worldid, flex_edgeadr[f] + idx]
+ reference = flexedge_length0[flex_edgeadr[f] + idx]
+ previous = deformed - vel * timestep
+ elongation[e] = deformed * deformed - reference * reference + (deformed * deformed - previous * previous) * kD
+
+ metric = wp.mat(0.0, shape=(6, 6))
+ id = int(0)
+ for ed1 in range(nedge):
+ for ed2 in range(ed1, nedge):
+ metric[ed1, ed2] = flex_stiffness[21 * elemid + id]
+ metric[ed2, ed1] = flex_stiffness[21 * elemid + id]
+ id += 1
+
+ force = wp.mat(0.0, shape=(6, 3))
+ for ed1 in range(nedge):
+ for ed2 in range(nedge):
+ for i in range(2):
+ for x in range(3):
+ force[edges[ed2, i], x] -= elongation[ed1] * gradient[ed2, 3 * i + x] * metric[ed1, ed2]
+
+ for v in range(nvert):
+ vert = flex_elem[(dim + 1) * elemid + v]
+ bodyid = flex_vertbodyid[flex_vertadr[f] + vert]
+ for x in range(3):
+ wp.atomic_add(qfrc_spring_out, worldid, body_dofadr[bodyid] + x, force[v, x])
+
+
+@wp.kernel
+def _flex_bending(
+ # Model:
+ body_dofadr: wp.array(dtype=int),
+ flex_dim: wp.array(dtype=int),
+ flex_vertadr: wp.array(dtype=int),
+ flex_edgeadr: wp.array(dtype=int),
+ flex_vertbodyid: wp.array(dtype=int),
+ flex_edge: wp.array(dtype=wp.vec2i),
+ flex_edgeflap: wp.array(dtype=wp.vec2i),
+ flex_bending: wp.array(dtype=wp.mat44f),
+ # Data in:
+ flexvert_xpos_in: wp.array2d(dtype=wp.vec3),
+ # Data out:
+ qfrc_spring_out: wp.array2d(dtype=float),
+):
+ worldid, edgeid = wp.tid()
+ nvert = 4
+ f = 0 # TODO(quaglino): this should become a function of t
+
+ if flex_dim[f] != 2:
+ return
+
+ v = wp.vec4i(
+ flex_edge[edgeid + flex_edgeadr[f]][0],
+ flex_edge[edgeid + flex_edgeadr[f]][1],
+ flex_edgeflap[edgeid + flex_edgeadr[f]][0],
+ flex_edgeflap[edgeid + flex_edgeadr[f]][1],
+ )
+
+ if v[3] == -1:
+ return
+
+ force = wp.mat(0.0, shape=(nvert, 3))
+ for i in range(nvert):
+ for j in range(nvert):
+ for x in range(3):
+ force[i, x] -= flex_bending[edgeid][i, j] * flexvert_xpos_in[worldid, v[j]][x]
+
+ for i in range(nvert):
+ bodyid = flex_vertbodyid[flex_vertadr[f] + v[i]]
+ for x in range(3):
+ wp.atomic_add(qfrc_spring_out, worldid, body_dofadr[bodyid] + x, force[i, x])
+
+
+@event_scope
+def passive(m: Model, d: Data):
+ """Adds all passive forces."""
+
+ if m.opt.disableflags & DisableBit.PASSIVE:
+ d.qfrc_spring.zero_()
+ d.qfrc_damper.zero_()
+ d.qfrc_gravcomp.zero_()
+ d.qfrc_fluid.zero_()
+ d.qfrc_passive.zero_()
+ return
+
+ wp.launch(
+ _spring_damper_dof_passive,
+ dim=(d.nworld, m.njnt),
+ inputs=[
+ m.qpos_spring,
+ m.jnt_type,
+ m.jnt_qposadr,
+ m.jnt_dofadr,
+ m.jnt_stiffness,
+ m.dof_damping,
+ d.qpos,
+ d.qvel,
+ ],
+ outputs=[d.qfrc_spring, d.qfrc_damper],
+ )
+
+ if m.ntendon:
+ wp.launch(
+ _spring_damper_tendon_passive,
+ dim=(d.nworld, m.ntendon, m.nv),
+ inputs=[
+ m.tendon_stiffness,
+ m.tendon_damping,
+ m.tendon_lengthspring,
+ d.ten_velocity,
+ d.ten_length,
+ d.ten_J,
+ ],
+ outputs=[
+ d.qfrc_spring,
+ d.qfrc_damper,
+ ],
+ )
+
+ wp.launch(
+ _flex_elasticity,
+ dim=(d.nworld, m.nflexelem),
+ inputs=[
+ m.opt.timestep,
+ m.body_dofadr,
+ m.flex_dim,
+ m.flex_vertadr,
+ m.flex_edgeadr,
+ m.flex_elemedgeadr,
+ m.flex_vertbodyid,
+ m.flex_elem,
+ m.flex_elemedge,
+ m.flexedge_length0,
+ m.flex_stiffness,
+ m.flex_damping,
+ d.flexvert_xpos,
+ d.flexedge_length,
+ d.flexedge_velocity,
+ ],
+ outputs=[d.qfrc_spring],
+ )
+ wp.launch(
+ _flex_bending,
+ dim=(d.nworld, m.nflexedge),
+ inputs=[
+ m.body_dofadr,
+ m.flex_dim,
+ m.flex_vertadr,
+ m.flex_edgeadr,
+ m.flex_vertbodyid,
+ m.flex_edge,
+ m.flex_edgeflap,
+ m.flex_bending,
+ d.flexvert_xpos,
+ ],
+ outputs=[d.qfrc_spring],
+ )
+
+ gravcomp = m.ngravcomp and not (m.opt.disableflags & DisableBit.GRAVITY)
+
+ if gravcomp:
+ d.qfrc_gravcomp.zero_()
+ wp.launch(
+ _gravity_force,
+ dim=(d.nworld, m.nbody - 1, m.nv),
+ inputs=[
+ m.opt.gravity,
+ m.body_parentid,
+ m.body_rootid,
+ m.body_mass,
+ m.body_gravcomp,
+ m.dof_bodyid,
+ d.xipos,
+ d.subtree_com,
+ d.cdof,
+ ],
+ outputs=[d.qfrc_gravcomp],
+ )
+
+ if m.opt.has_fluid:
+ _fluid(m, d)
+
+ wp.launch(
+ _qfrc_passive,
+ dim=(d.nworld, m.nv),
+ inputs=[
+ m.opt.has_fluid,
+ m.jnt_actgravcomp,
+ m.dof_jntid,
+ d.qfrc_spring,
+ d.qfrc_damper,
+ d.qfrc_gravcomp,
+ d.qfrc_fluid,
+ gravcomp,
+ ],
+ outputs=[
+ d.qfrc_passive,
+ ],
+ )
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive_test.py
new file mode 100644
index 00000000..f914d268
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive_test.py
@@ -0,0 +1,144 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 passive force functions."""
+
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+from absl.testing import parameterized
+
+import mujoco_warp as mjwarp
+
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+
+# tolerance for difference between MuJoCo and MJWarp passive force calculations - mostly
+# due to float precision
+_TOLERANCE = 5e-5
+
+
+def _assert_eq(a, b, name):
+ tol = _TOLERANCE * 10 # avoid test noise
+ err_msg = f"mismatch: {name}"
+ np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+
+
+class PassiveTest(parameterized.TestCase):
+ @parameterized.product(passive=[True, False], gravity=[True, False])
+ def test_passive(self, passive, gravity):
+ """Tests passive."""
+ _, mjd, m, d = test_util.fixture("pendula.xml", passive=passive, gravity=gravity, kick=True, applied=True)
+
+ for arr in (d.qfrc_spring, d.qfrc_damper, d.qfrc_gravcomp, d.qfrc_passive):
+ arr.zero_()
+
+ mjwarp.passive(m, d)
+
+ _assert_eq(d.qfrc_spring.numpy()[0], mjd.qfrc_spring, "qfrc_spring")
+ _assert_eq(d.qfrc_damper.numpy()[0], mjd.qfrc_damper, "qfrc_damper")
+ _assert_eq(d.qfrc_gravcomp.numpy()[0], mjd.qfrc_gravcomp, "qfrc_gravcomp")
+ _assert_eq(d.qfrc_passive.numpy()[0], mjd.qfrc_passive, "qfrc_passive")
+
+ @parameterized.parameters(
+ (1, 0, 0, 0, 0),
+ (0, 1, 0, 0, 0),
+ (0, 0, 1, 0, 0),
+ (0, 0, 0, 1, 0),
+ (0, 0, 0, 0, 1),
+ (1, 1, 1, 1, 1),
+ )
+ def test_fluid(self, density, viscosity, wind0, wind1, wind2):
+ """Tests fluid model."""
+
+ _, mjd, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ keyframe=0,
+ )
+
+ for arr in (d.qfrc_passive, d.qfrc_fluid):
+ arr.zero_()
+
+ mjwarp.passive(m, d)
+
+ _assert_eq(d.qfrc_passive.numpy()[0], mjd.qfrc_passive, "qfrc_passive")
+ _assert_eq(d.qfrc_fluid.numpy()[0], mjd.qfrc_fluid, "qfrc_fluid")
+
+ @parameterized.parameters((True, True), (True, False), (False, True), (False, False))
+ def test_gravcomp(self, sparse, gravity):
+ """Tests gravity compensation."""
+
+ _, mjd, m, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ gravity=gravity,
+ sparse=sparse,
+ )
+
+ for arr in (d.qfrc_passive, d.qfrc_gravcomp, d.qfrc_actuator):
+ arr.zero_()
+
+ mjwarp.passive(m, d)
+ mjwarp.fwd_actuation(m, d)
+
+ _assert_eq(d.qfrc_passive.numpy()[0], mjd.qfrc_passive, "qfrc_passive")
+ _assert_eq(d.qfrc_gravcomp.numpy()[0], mjd.qfrc_gravcomp, "qfrc_gravcomp")
+ _assert_eq(d.qfrc_actuator.numpy()[0], mjd.qfrc_actuator, "qfrc_actuator")
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py
new file mode 100644
index 00000000..114ce4a8
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py
@@ -0,0 +1,903 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+from typing import Tuple
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec6
+
+
+@wp.func
+def _ray_map(pos: wp.vec3, mat: wp.mat33, pnt: wp.vec3, vec: wp.vec3) -> Tuple[wp.vec3, wp.vec3]:
+ """Maps ray to local geom frame coordinates.
+
+ Args:
+ pos: position of geom frame
+ mat: orientation of geom frame
+ pnt: starting point of ray in world coordinates
+ vec: direction of ray in world coordinates
+
+ Returns:
+ 3D point and 3D direction in local geom frame
+ """
+
+ matT = wp.transpose(mat)
+ lpnt = matT @ (pnt - pos)
+ lvec = matT @ vec
+
+ return lpnt, lvec
+
+
+@wp.func
+def _ray_eliminate(
+ # Model:
+ body_weldid: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ geom_group: wp.array(dtype=int),
+ geom_matid: wp.array(dtype=int), # kernel_analyzer: ignore
+ geom_rgba: wp.array(dtype=wp.vec4), # kernel_analyzer: ignore
+ mat_rgba: wp.array(dtype=wp.vec4), # kernel_analyzer: ignore
+ # In:
+ geomid: int,
+ geomgroup: vec6,
+ flg_static: bool,
+ bodyexclude: int,
+) -> bool:
+ """Eliminate ray."""
+ bodyid = geom_bodyid[geomid]
+ matid = geom_matid[geomid]
+
+ # body exclusion
+ if bodyid == bodyexclude:
+ return True
+
+ # invisible geom exclusion
+ if matid < 0 and geom_rgba[geomid][3] == 0.0:
+ return True
+
+ # invisible material exclusion
+ if matid >= 0:
+ if mat_rgba[matid][3] == 0.0:
+ return True
+
+ # static exclusion
+ if not flg_static and body_weldid[bodyid] == 0:
+ return True
+
+ # no geomgroup inclusion
+ if (
+ geomgroup[0] == -1
+ and geomgroup[1] == -1
+ and geomgroup[2] == -1
+ and geomgroup[3] == -1
+ and geomgroup[4] == -1
+ and geomgroup[5] == -1
+ ):
+ return False
+
+ # group inclusion/exclusion
+ groupid = wp.min(5, wp.max(0, geom_group[geomid]))
+
+ return geomgroup[groupid] == 0
+
+
+@wp.func
+def _ray_quad(a: float, b: float, c: float) -> Tuple[float, wp.vec2]:
+ """Compute solutions from quadratic: a*x^2 + 2*b*x + c = 0."""
+ det = b * b - a * c
+ if det < MJ_MINVAL:
+ return wp.inf, wp.vec2(wp.inf, wp.inf)
+ det = wp.sqrt(det)
+
+ # compute the two solutions
+ den = 1.0 / a
+ x0 = (-b - det) * den
+ x1 = (-b + det) * den
+ x = wp.vec2(x0, x1)
+
+ # finalize result
+ if x0 >= 0.0:
+ return x0, x
+ elif x1 >= 0.0:
+ return x1, x
+ else:
+ return wp.inf, x
+
+
+@wp.func
+def _ray_triangle(v0: wp.vec3, v1: wp.vec3, v2: wp.vec3, pnt: wp.vec3, vec: wp.vec3, b0: wp.vec3, b1: wp.vec3) -> float:
+ """Returns the distance at which a ray intersects with a triangle."""
+ dif0 = v0 - pnt
+ dif1 = v1 - pnt
+ dif2 = v2 - pnt
+
+ # project difference vectors in normal plane
+ planar_00 = wp.dot(dif0, b0)
+ planar_01 = wp.dot(dif0, b1)
+ planar_10 = wp.dot(dif1, b0)
+ planar_11 = wp.dot(dif1, b1)
+ planar_20 = wp.dot(dif2, b0)
+ planar_21 = wp.dot(dif2, b1)
+
+ # reject if on the same side of any coordinate axis
+ if (
+ (planar_00 > 0.0 and planar_10 > 0.0 and planar_20 > 0.0)
+ or (planar_00 < 0.0 and planar_10 < 0.0 and planar_20 < 0.0)
+ or (planar_01 > 0.0 and planar_11 > 0.0 and planar_21 > 0.0)
+ or (planar_01 < 0.0 and planar_11 < 0.0 and planar_21 < 0.0)
+ ):
+ return float(wp.inf)
+
+ # determine if origin is inside planar projection of triangle
+ # A = (p0-p2, p1-p2), b = -p2, solve A*t = b
+ A00 = planar_00 - planar_20
+ A10 = planar_10 - planar_20
+ A01 = planar_01 - planar_21
+ A11 = planar_11 - planar_21
+
+ b = wp.vec2(-planar_20, -planar_21)
+
+ det = A00 * A11 - A10 * A01
+ if wp.abs(det) < MJ_MINVAL:
+ return float(wp.inf)
+
+ t0 = (A11 * b[0] - A10 * b[1]) / det
+ t1 = (-A01 * b[0] + A00 * b[1]) / det
+
+ # check if outside
+ if t0 < 0.0 or t1 < 0.0 or t0 + t1 > 1.0:
+ return float(wp.inf)
+
+ # intersect ray with plane of triangle
+ dif0 = v0 - v2
+ dif1 = v1 - v2
+ dif2 = pnt - v2
+ nrm = wp.cross(dif0, dif1) # normal to triangle plane
+ denom = wp.dot(vec, nrm)
+ if wp.abs(denom) < MJ_MINVAL:
+ return float(wp.inf)
+
+ dist = -wp.dot(dif2, nrm) / denom
+ return wp.where(dist >= 0.0, dist, float(wp.inf))
+
+
+@wp.func
+def _ray_plane(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float:
+ """Returns the distance at which a ray intersects with a plane."""
+
+ # map to local frame
+ lpnt, lvec = _ray_map(pos, mat, pnt, vec)
+
+ # z-vec not pointing towards front face: reject
+ if lvec[2] > -MJ_MINVAL:
+ return wp.inf
+
+ # intersection with plane
+ x = -lpnt[2] / lvec[2]
+ if x < 0.0:
+ return wp.inf
+
+ p = wp.vec2(
+ lpnt[0] + x * lvec[0],
+ lpnt[1] + x * lvec[1],
+ )
+
+ # accept only within rendered rectangle
+ if (size[0] <= 0.0 or wp.abs(p[0]) <= size[0]) and (size[1] <= 0.0 or wp.abs(p[1]) <= size[1]):
+ return x
+ else:
+ return wp.inf
+
+
+@wp.func
+def _ray_sphere(pos: wp.vec3, dist_sqr: float, pnt: wp.vec3, vec: wp.vec3) -> float:
+ """Returns the distance at which a ray intersects with a sphere."""
+ dif = pnt - pos
+
+ a = wp.dot(vec, vec)
+ b = wp.dot(vec, dif)
+ c = wp.dot(dif, dif) - dist_sqr
+
+ sol, _ = _ray_quad(a, b, c)
+ return sol
+
+
+@wp.func
+def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float:
+ """Returns the distance at which a ray intersects with a capsule."""
+
+ # bounding sphere test
+ ssz = size[0] + size[1]
+ if _ray_sphere(pos, ssz * ssz, pnt, vec) < 0.0:
+ return wp.inf
+
+ # map to local frame
+ lpnt, lvec = _ray_map(pos, mat, pnt, vec)
+
+ # init solution
+ x = -1.0
+
+ # cylinder round side: (x * lvec + lpnt)' * (x * lvec + lpnt) = size[0] * size[0]
+ sq_size0 = size[0] * size[0]
+ a = lvec[0] * lvec[0] + lvec[1] * lvec[1]
+ b = lvec[0] * lpnt[0] + lvec[1] * lpnt[1]
+ c = lpnt[0] * lpnt[0] + lpnt[1] * lpnt[1] - sq_size0
+
+ # solve a * x^2 + 2 * b * x + c = 0
+ sol, xx = _ray_quad(a, b, c)
+
+ # make sure round solution is between flat sides
+ if sol >= 0.0 and wp.abs(lpnt[2] + sol * vec[2]) <= size[1]:
+ if x < 0.0 or sol < x:
+ x = sol
+
+ # top cap
+ ldif = wp.vec3(lpnt[0], lpnt[1], lpnt[2] - size[1])
+ a += lvec[2] * lvec[2]
+ b = wp.dot(lvec, ldif)
+ c = wp.dot(ldif, ldif) - sq_size0
+ _, xx = _ray_quad(a, b, c)
+
+ # accept only top half of sphere
+ for i in range(2):
+ if xx[i] >= 0.0 and lpnt[2] + xx[i] * lvec[2] >= size[1]:
+ if x < 0.0 or xx[i] < x:
+ x = xx[i]
+
+ # bottom cap
+ ldif = wp.vec3(ldif[0], ldif[1], lpnt[2] + size[1])
+ b = wp.dot(lvec, ldif)
+ c = wp.dot(ldif, ldif) - sq_size0
+ _, xx = _ray_quad(a, b, c)
+
+ # accept only bottom half of sphere
+ for i in range(2):
+ if xx[i] >= 0.0 and lpnt[2] + xx[i] * lvec[2] <= -size[1]:
+ if x < 0.0 or xx[i] < x:
+ x = xx[i]
+
+ return x
+
+
+@wp.func
+def _ray_ellipsoid(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float:
+ """Returns the distance at which a ray intersects with an ellipsoid."""
+
+ # map to local frame
+ lpnt, lvec = _ray_map(pos, mat, pnt, vec)
+
+ # invert size^2
+ s = wp.vec3(
+ 1.0 / (size[0] * size[0]),
+ 1.0 / (size[1] * size[1]),
+ 1.0 / (size[2] * size[2]),
+ )
+
+ # (x * lvec + lpnt)' * diag(1 / size^2) * (x * lvec + lpnt) = 1
+ slvec = wp.cw_mul(s, lvec)
+ a = wp.dot(slvec, lvec)
+ b = wp.dot(slvec, lpnt)
+ c = wp.dot(wp.cw_mul(s, lpnt), lpnt) - 1.0
+
+ # solve a * x^2 + 2 * b * x + c = 0
+ sol, _ = _ray_quad(a, b, c)
+ return sol
+
+
+@wp.func
+def _ray_cylinder(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float:
+ """Returns the distance at which a ray intersects with a cylinder."""
+ # bounding sphere test
+ ssz = size[0] * size[0] + size[1] * size[1]
+ if _ray_sphere(pos, ssz, pnt, vec) < 0.0:
+ return wp.inf
+
+ # map to local frame
+ lpnt, lvec = _ray_map(pos, mat, pnt, vec)
+
+ # init solution
+ x = wp.inf
+
+ # flat sides
+ if wp.abs(lvec[2]) > MJ_MINVAL:
+ for side in range(-1, 2, 2):
+ # solution of: lpnt[2] + x * lvec[2] = side * height_size
+ sol = (float(side) * size[1] - lpnt[2]) / lvec[2]
+
+ # process if non-negative
+ if sol >= 0.0:
+ # intersection with horizontal face
+ p = wp.vec2(
+ lpnt[0] + sol * lvec[0],
+ lpnt[1] + sol * lvec[1],
+ )
+
+ # accept within radius
+ if wp.dot(p, p) <= size[0] * size[0]:
+ if x < 0.0 or sol < x:
+ x = sol
+
+ # (x * lvec + lpnt)' * (x * lvec + lpnt) = size[0] * size[0]
+ a = lvec[0] * lvec[0] + lvec[1] * lvec[1]
+ b = lvec[0] * lpnt[0] + lvec[1] * lpnt[1]
+ c = lpnt[0] * lpnt[0] + lpnt[1] * lpnt[1] - size[0] * size[0]
+
+ # solve a * x^2 + 2 * b * x + c = 0
+ sol, _ = _ray_quad(a, b, c)
+
+ # make sure round solution is between flat sides
+ if sol >= 0.0 and wp.abs(lpnt[2] + sol * lvec[2]) <= size[1]:
+ if x < 0.0 or sol < x:
+ x = sol
+
+ return x
+
+
+_IFACE = wp.types.matrix((3, 2), dtype=int)(1, 2, 0, 2, 0, 1)
+
+
+@wp.func
+def _ray_box(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, vec6]:
+ """Returns the distance at which a ray intersects with a box."""
+ all = vec6(-1.0, -1.0, -1.0, -1.0, -1.0, -1.0)
+
+ # bounding sphere test
+ ssz = wp.dot(size, size)
+ if _ray_sphere(pos, ssz, pnt, vec) < 0.0:
+ return wp.inf, all
+
+ # map to local frame
+ lpnt, lvec = _ray_map(pos, mat, pnt, vec)
+
+ # init solution
+ x = wp.inf
+
+ # loop over axes with non-zero vec
+ for i in range(3):
+ if wp.abs(lvec[i]) > MJ_MINVAL:
+ for side in range(-1, 2, 2):
+ # solution of: lpnt[i] + x * lvec[i] = side * size[i]
+ sol = (float(side) * size[i] - lpnt[i]) / lvec[i]
+
+ # process if non-negative
+ if sol >= 0.0:
+ id0 = _IFACE[i][0]
+ id1 = _IFACE[i][1]
+
+ # intersection with face
+ p0 = lpnt[id0] + sol * lvec[id0]
+ p1 = lpnt[id1] + sol * lvec[id1]
+
+ # accept within rectangle
+ if (wp.abs(p0) <= size[id0]) and (wp.abs(p1) <= size[id1]):
+ # update
+ if (x < 0.0) or (sol < x):
+ x = sol
+
+ # save in all
+ all[2 * i + (side + 1) / 2] = sol
+
+ return x, all
+
+
+@wp.func
+def _ray_hfield(
+ # Model:
+ geom_type: wp.array(dtype=int),
+ geom_dataid: wp.array(dtype=int),
+ hfield_adr: wp.array(dtype=int),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ hfield_data: wp.array(dtype=float),
+ # In:
+ pos: wp.vec3,
+ mat: wp.mat33,
+ pnt: wp.vec3,
+ vec: wp.vec3,
+ id: int,
+):
+ # check geom type
+ if geom_type[id] != int(GeomType.HFIELD.value):
+ return wp.inf
+
+ # hfield id and dimensions
+ hid = geom_dataid[id]
+ nrow = hfield_nrow[hid]
+ ncol = hfield_ncol[hid]
+
+ size = hfield_size[hid]
+ adr = hfield_adr[hid]
+
+ mat_col = wp.vec3(mat[0, 2], mat[1, 2], mat[2, 2])
+
+ # compute size and pos of base box
+ base_scale = size[3] * 0.5
+ base_size = wp.vec3(size[0], size[1], base_scale)
+ base_pos = pos + mat_col * base_scale
+
+ # compute size and pos of top box
+ top_scale = size[2] * 0.5
+ top_size = wp.vec3(size[0], size[1], top_scale)
+ top_pos = pos + mat_col * top_scale
+
+ # init: intersection with base box
+ x, _ = _ray_box(base_pos, mat, base_size, pnt, vec)
+
+ # check top box: done if no intersection
+ top_intersect, all = _ray_box(top_pos, mat, top_size, pnt, vec)
+
+ if top_intersect < 0.0:
+ return x
+
+ # map to local frame
+ lpnt, lvec = _ray_map(pos, mat, pnt, vec)
+
+ # construct basis vectors of normal plane
+ b0 = wp.vec3(1.0, 1.0, 1.0)
+
+ if wp.abs(lvec[0]) >= wp.abs(lvec[1]) and wp.abs(lvec[0]) >= wp.abs(lvec[2]):
+ b0[0] = 0.0
+ elif wp.abs(lvec[1]) >= wp.abs(lvec[2]):
+ b0[1] = 0.0
+ else:
+ b0[2] = 0.0
+ b1 = b0 + lvec * -wp.dot(lvec, b0) / wp.dot(lvec, lvec)
+ b1 = wp.normalize(b1)
+
+ b2 = wp.cross(b1, lvec)
+ b2 = wp.normalize(b2)
+
+ # find ray segment intersecting top box
+ seg = wp.vec2(0.0, top_intersect)
+ for i in range(6):
+ if all[i] > seg[1]:
+ seg[0] = top_intersect
+ seg[1] = all[i]
+
+ # project segment endpoints in horizontal plane, discretize
+ dx = (2.0 * size[0]) / float(ncol - 1)
+ dy = (2.0 * size[1]) / float(nrow - 1)
+ SX = wp.vec2((lpnt[0] * seg[0] * lvec[0] + size[0]) / dx, (lpnt[0] * seg[1] * lvec[0] + size[0]) / dx)
+ SY = wp.vec2((lpnt[1] + seg[0] * lvec[1] + size[1]) / dy, (lpnt[1] + seg[1] * lvec[1] + size[1]) / dy)
+
+ # compute ranges, with +1 padding
+ cmin = wp.max(0, int(wp.floor(wp.min(SX[0], SX[1])) - 1.0))
+ cmax = wp.min(ncol - 1, int(wp.ceil(wp.max(SX[0], SX[1])) + 1.0))
+ rmin = wp.max(0, int(wp.floor(wp.min(SY[0], SY[1])) - 1.0))
+ rmax = wp.min(nrow - 1, int(wp.ceil(wp.max(SY[0], SY[1])) + 1.0))
+
+ # check triangles within bounds
+ for r in range(rmin, rmax):
+ for c in range(cmin, cmax):
+ # first triangle
+ v0 = wp.vec3(dx * float(c) - size[0], dy * float(r) - size[1], hfield_data[adr + r * ncol + c] * size[2])
+ v1 = wp.vec3(
+ dx * float(c + 1) - size[0], dy * float(r + 1) - size[1], hfield_data[adr + (r + 1) * ncol + (c + 1)] * size[2]
+ )
+ v2 = wp.vec3(dx * float(c + 1) - size[0], dy * float(r) - size[1], hfield_data[adr + r * ncol + (c + 1)] * size[2])
+ sol = _ray_triangle(v0, v1, v2, pnt, vec, b0, b1)
+ if sol >= 0.0 and (x < 0.0 or sol < x):
+ x = sol
+
+ # second triangle
+ v0 = wp.vec3(dx * float(c) - size[0], dy * float(r) - size[1], hfield_data[adr + r * ncol + c] * size[2])
+ v1 = wp.vec3(
+ dx * float(c + 1) - size[0], dy * float(r + 1) - size[1], hfield_data[adr + (r + 1) * ncol + (c + 1)] * size[2]
+ )
+ v2 = wp.vec3(dx * float(c) - size[0], dy * float(r + 1) - size[1], hfield_data[adr + (r + 1) * ncol + c] * size[2])
+ sol = _ray_triangle(v0, v1, v2, pnt, vec, b0, b1)
+ if sol >= 0.0 and (x < 0.0 or sol < x):
+ x = sol
+
+ # check viable sides of top box
+ for i in range(4):
+ if all[i] >= 0.0 and (all[i] < x or x < 0.0):
+ # normalized height of intersection point
+ z = (lpnt[2] + all[i] * lvec[2]) / size[2]
+
+ # rectangle points: y, y0, z0, z1
+ # side normal to x-axis
+ if i < 2:
+ y = (lpnt[1] + all[i] * lvec[1] + size[1]) / dy
+ y0 = wp.max(0.0, wp.min(float(nrow - 2), wp.floor(y)))
+ if i == 1:
+ z0 = hfield_data[adr + int(wp.round(y0)) * nrow + ncol - 1]
+ z1 = hfield_data[adr + int(wp.round(y0 + 1.0)) * nrow + ncol - 1]
+ else:
+ z0 = hfield_data[adr + int(wp.round(y0)) * nrow]
+ z1 = hfield_data[adr + int(wp.round(y0 + 1.0)) * nrow]
+ # side normal to y-axis
+ else:
+ y = (lpnt[0] + all[i] * lvec[0] + size[0]) / dx
+ y0 = wp.max(0.0, wp.min(float(ncol - 2), wp.floor(y)))
+ if i == 3:
+ z0 = hfield_data[adr + int(wp.round(y0)) + (nrow - 1) * ncol]
+ z1 = hfield_data[adr + int(wp.round(y0 + 1.0)) + (nrow - 1) * ncol]
+ else:
+ z0 = hfield_data[adr + int(wp.round(y0))]
+ z1 = hfield_data[adr + int(wp.round(y0 + 1.0))]
+
+ # check if point is below line segments
+ if z < z0 * (y0 + 1.0 - y) + z1 * (y - y0):
+ x = all[i]
+
+ return x
+
+
+@wp.func
+def _ray_mesh(
+ # Model:
+ nmeshface: int,
+ mesh_vertadr: wp.array(dtype=int),
+ mesh_vert: wp.array(dtype=wp.vec3),
+ mesh_faceadr: wp.array(dtype=int),
+ mesh_face: wp.array(dtype=wp.vec3i),
+ # In:
+ data_id: int,
+ pos: wp.vec3,
+ mat: wp.mat33,
+ pnt: wp.vec3,
+ vec: wp.vec3,
+) -> float:
+ """Returns the distance and geomid for ray mesh intersections."""
+ pnt, vec = _ray_map(pos, mat, pnt, vec)
+
+ # compute orthogonal basis vectors
+ if wp.abs(vec[0]) < wp.abs(vec[1]):
+ if wp.abs(vec[0]) < wp.abs(vec[2]):
+ b0 = wp.vec3(0.0, vec[2], -vec[1])
+ else:
+ b0 = wp.vec3(vec[1], -vec[0], 0.0)
+ else:
+ if wp.abs(vec[1]) < wp.abs(vec[2]):
+ b0 = wp.vec3(-vec[2], 0.0, vec[0])
+ else:
+ b0 = wp.vec3(vec[1], -vec[0], 0.0)
+
+ # normalize first vector
+ b0 = wp.normalize(b0)
+
+ # compute second vector as cross product
+ b1 = wp.cross(vec, b0)
+ b1 = wp.normalize(b1)
+
+ min_dist = float(wp.inf)
+
+ # get mesh vertex data range
+ vert_start = mesh_vertadr[data_id]
+
+ # get mesh face and vertex data
+ face_start = mesh_faceadr[data_id]
+
+ if data_id + 1 < mesh_faceadr.shape[0]:
+ face_end = mesh_faceadr[data_id + 1]
+ else:
+ face_end = nmeshface
+
+ # iterate through all faces
+ for i in range(face_start, face_end):
+ # get vertices for this face
+ v_idx = mesh_face[i]
+
+ # create triangle struct
+ v0 = mesh_vert[vert_start + v_idx.x]
+ v1 = mesh_vert[vert_start + v_idx.y]
+ v2 = mesh_vert[vert_start + v_idx.z]
+
+ # calculate intersection
+ dist = _ray_triangle(v0, v1, v2, pnt, vec, b0, b1)
+ if dist < min_dist:
+ min_dist = dist
+
+ return min_dist
+
+
+@wp.func
+def ray_geom(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3, geomtype: int) -> float:
+ """Returns distance along ray to intersection with geom, or infinity if none."""
+
+ # TODO(team): static loop unrolling to remove unnecessary branching
+ if geomtype == int(GeomType.PLANE.value):
+ return _ray_plane(pos, mat, size, pnt, vec)
+ elif geomtype == int(GeomType.SPHERE.value):
+ return _ray_sphere(pos, size[0] * size[0], pnt, vec)
+ elif geomtype == int(GeomType.CAPSULE.value):
+ return _ray_capsule(pos, mat, size, pnt, vec)
+ elif geomtype == int(GeomType.ELLIPSOID.value):
+ return _ray_ellipsoid(pos, mat, size, pnt, vec)
+ elif geomtype == int(GeomType.CYLINDER.value):
+ return _ray_cylinder(pos, mat, size, pnt, vec)
+ elif geomtype == int(GeomType.BOX.value):
+ dist, _ = _ray_box(pos, mat, size, pnt, vec)
+ return dist
+ else:
+ return wp.inf
+
+
+@wp.func
+def _ray_geom_mesh(
+ # Model:
+ nmeshface: int,
+ body_weldid: wp.array(dtype=int),
+ geom_type: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ geom_dataid: wp.array(dtype=int),
+ geom_group: wp.array(dtype=int),
+ geom_matid: wp.array2d(dtype=int),
+ geom_size: wp.array2d(dtype=wp.vec3),
+ geom_rgba: wp.array2d(dtype=wp.vec4),
+ hfield_adr: wp.array(dtype=int),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ hfield_data: wp.array(dtype=float),
+ mesh_vertadr: wp.array(dtype=int),
+ mesh_vert: wp.array(dtype=wp.vec3),
+ mesh_faceadr: wp.array(dtype=int),
+ mesh_face: wp.array(dtype=wp.vec3i),
+ mat_rgba: wp.array2d(dtype=wp.vec4),
+ # Data in:
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ # In:
+ worldid: int,
+ pnt: wp.vec3,
+ vec: wp.vec3,
+ geomgroup: vec6,
+ flg_static: bool,
+ bodyexclude: int,
+ geomid: int,
+) -> float:
+ if not _ray_eliminate(
+ body_weldid,
+ geom_bodyid,
+ geom_group,
+ geom_matid[worldid],
+ geom_rgba[worldid],
+ mat_rgba[worldid],
+ geomid,
+ geomgroup,
+ flg_static,
+ bodyexclude,
+ ):
+ pos = geom_xpos_in[worldid, geomid]
+ mat = geom_xmat_in[worldid, geomid]
+ type = geom_type[geomid]
+
+ if type == int(GeomType.MESH.value):
+ return _ray_mesh(
+ nmeshface,
+ mesh_vertadr,
+ mesh_vert,
+ mesh_faceadr,
+ mesh_face,
+ geom_dataid[geomid],
+ pos,
+ mat,
+ pnt,
+ vec,
+ )
+ elif type == int(GeomType.HFIELD.value):
+ return _ray_hfield(
+ geom_type,
+ geom_dataid,
+ hfield_adr,
+ hfield_nrow,
+ hfield_ncol,
+ hfield_size,
+ hfield_data,
+ pos,
+ mat,
+ pnt,
+ vec,
+ geomid,
+ )
+ else:
+ return ray_geom(pos, mat, geom_size[worldid, geomid], pnt, vec, type)
+ else:
+ return wp.inf
+
+
+@wp.kernel
+def _ray(
+ # Model:
+ ngeom: int,
+ nmeshface: int,
+ body_weldid: wp.array(dtype=int),
+ geom_type: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ geom_dataid: wp.array(dtype=int),
+ geom_group: wp.array(dtype=int),
+ geom_matid: wp.array2d(dtype=int),
+ geom_size: wp.array2d(dtype=wp.vec3),
+ geom_rgba: wp.array2d(dtype=wp.vec4),
+ hfield_adr: wp.array(dtype=int),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ hfield_data: wp.array(dtype=float),
+ mesh_vertadr: wp.array(dtype=int),
+ mesh_vert: wp.array(dtype=wp.vec3),
+ mesh_faceadr: wp.array(dtype=int),
+ mesh_face: wp.array(dtype=wp.vec3i),
+ mat_rgba: wp.array2d(dtype=wp.vec4),
+ # Data in:
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ # In:
+ pnt: wp.array2d(dtype=wp.vec3),
+ vec: wp.array2d(dtype=wp.vec3),
+ geomgroup: vec6,
+ flg_static: bool,
+ bodyexclude: wp.array(dtype=int),
+ # Out:
+ dist_out: wp.array(dtype=float, ndim=2),
+ geomid_out: wp.array(dtype=int, ndim=2),
+):
+ worldid, rayid, tid = wp.tid()
+
+ num_threads = wp.block_dim()
+
+ min_dist = float(wp.inf)
+ min_geomid = int(-1)
+
+ upper = ((ngeom + num_threads - 1) // num_threads) * num_threads
+ for geomid in range(tid, upper, num_threads):
+ if geomid < ngeom:
+ dist = _ray_geom_mesh(
+ nmeshface,
+ body_weldid,
+ geom_type,
+ geom_bodyid,
+ geom_dataid,
+ geom_group,
+ geom_matid,
+ geom_size,
+ geom_rgba,
+ hfield_adr,
+ hfield_nrow,
+ hfield_ncol,
+ hfield_size,
+ hfield_data,
+ mesh_vertadr,
+ mesh_vert,
+ mesh_faceadr,
+ mesh_face,
+ mat_rgba,
+ geom_xpos_in,
+ geom_xmat_in,
+ worldid,
+ pnt[worldid, rayid],
+ vec[worldid, rayid],
+ geomgroup,
+ flg_static,
+ bodyexclude[rayid],
+ geomid,
+ )
+ else:
+ dist = wp.inf
+
+ tile_dist = wp.tile(dist)
+ local_min_geomid = wp.tile_argmin(tile_dist)
+ local_min_dist = tile_dist[local_min_geomid[0]]
+
+ tile_geomid = wp.tile(geomid)
+
+ if local_min_dist < min_dist:
+ min_dist = local_min_dist
+ min_geomid = tile_geomid[local_min_geomid[0]]
+
+ if wp.isinf(min_dist):
+ dist_out[worldid, rayid] = -1.0
+ else:
+ dist_out[worldid, rayid] = min_dist
+ geomid_out[worldid, rayid] = min_geomid
+
+
+def ray(
+ m: Model,
+ d: Data,
+ pnt: wp.array2d(dtype=wp.vec3),
+ vec: wp.array2d(dtype=wp.vec3),
+ geomgroup: vec6 = None,
+ flg_static: bool = True,
+ bodyexclude: int = -1,
+) -> tuple[wp.array2d(dtype=float), wp.array2d(dtype=int)]:
+ """Returns the distance at which rays intersect with primitive geoms.
+
+ Args:
+ m (Model): The model containing kinematic and dynamic information (device).
+ d (Data): The data object containing the current state and output arrays (device).
+ pnt (wp.array2d(dtype=wp.vec3)): Ray origin points.
+ vec (wp.array2d(dtype=wp.vec3)): Ray directions.
+ geomgroup (vec6, optional): Group inclusion/exclusion mask.
+ If all are wp.inf, ignore.
+ flg_static (bool, optional): If True, allows rays to intersect with static geoms.
+ Defaults to True.
+ bodyexclude (int, optional): Ignore geoms on specified body id (-1 to disable).
+ Defaults to -1.
+
+ Returns:
+ wp.array2d(dtype=float): Distances from ray origins to geom surfaces.
+ wp.array2d(dtype=int): IDs of intersected geoms (-1 if none).
+ """
+
+ assert pnt.shape[0] == vec.shape[0]
+ assert d.ray_dist.shape[1] == d.ray_geomid.shape[1]
+ assert pnt.shape[0] == d.ray_dist.shape[1]
+
+ if geomgroup is None:
+ geomgroup = vec6(-1, -1, -1, -1, -1, -1)
+
+ d.ray_bodyexclude.fill_(bodyexclude)
+
+ rays(m, d, pnt, vec, geomgroup, flg_static, d.ray_bodyexclude, d.ray_dist, d.ray_geomid)
+
+ return d.ray_dist, d.ray_geomid
+
+
+def rays(
+ m: Model,
+ d: Data,
+ pnt: wp.array2d(dtype=wp.vec3),
+ vec: wp.array2d(dtype=wp.vec3),
+ geomgroup: vec6,
+ flg_static: bool,
+ bodyexclude: wp.array(dtype=int),
+ dist: wp.array2d(dtype=wp.vec3),
+ geomid: wp.array2d(dtype=int),
+):
+ wp.launch_tiled(
+ _ray,
+ dim=(d.nworld, pnt.shape[1]),
+ inputs=[
+ m.ngeom,
+ m.nmeshface,
+ m.body_weldid,
+ m.geom_type,
+ m.geom_bodyid,
+ m.geom_dataid,
+ m.geom_group,
+ m.geom_matid,
+ m.geom_size,
+ m.geom_rgba,
+ m.hfield_adr,
+ m.hfield_nrow,
+ m.hfield_ncol,
+ m.hfield_size,
+ m.hfield_data,
+ m.mesh_vertadr,
+ m.mesh_vert,
+ m.mesh_faceadr,
+ m.mesh_face,
+ m.mat_rgba,
+ d.geom_xpos,
+ d.geom_xmat,
+ pnt,
+ vec,
+ geomgroup,
+ flg_static,
+ bodyexclude,
+ dist,
+ geomid,
+ ],
+ block_dim=m.block_dim.ray,
+ )
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray_test.py
new file mode 100644
index 00000000..a81dda7e
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray_test.py
@@ -0,0 +1,312 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 ray functions."""
+
+import mujoco
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+
+import mujoco_warp as mjwarp
+
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec6
+
+# tolerance for difference between MuJoCo and MJX ray calculations - mostly
+# due to float precision
+_TOLERANCE = 5e-5
+
+
+def _assert_eq(a, b, name):
+ tol = _TOLERANCE * 10 # avoid test noise
+ err_msg = f"mismatch: {name}"
+ np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+
+
+class RayTest(absltest.TestCase):
+ def test_ray_nothing(self):
+ """Tests that ray returns -1 when nothing is hit."""
+ mjm, mjd, m, d = test_util.fixture("ray.xml")
+
+ pnt = wp.array([wp.vec3(12.146, 1.865, 3.895)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.vec3(0.0, 0.0, -1.0)], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0] # Extract from [[-1]]
+ dist_np = dist.numpy()[0, 0] # Extract from [[-1.]]
+ _assert_eq(geomid_np, -1, "geom_id")
+ _assert_eq(dist_np, -1, "dist")
+
+ def test_ray_plane(self):
+ """Tests ray<>plane matches MuJoCo."""
+ mjm, mjd, m, d = test_util.fixture("ray.xml")
+
+ # looking down at a slight angle
+ pnt = wp.array([wp.vec3(2.0, 1.0, 3.0)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, 0, "geom_id")
+ pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0]
+ unused = np.zeros(1, dtype=np.int32)
+ mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused)
+ _assert_eq(dist_np, mj_dist, "dist")
+
+ # looking on wrong side of plane
+ pnt = wp.array([wp.vec3(0.0, 0.0, -0.5)], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, -1, "geom_id")
+ _assert_eq(dist_np, -1, "dist")
+
+ def test_ray_sphere(self):
+ """Tests ray<>sphere matches MuJoCo."""
+ mjm, mjd, m, d = test_util.fixture("ray.xml")
+
+ # looking down at sphere at a slight angle
+ pnt = wp.array([wp.vec3(0.0, 0.0, 1.6)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, 1, "geom_id")
+ pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0]
+ unused = np.zeros(1, dtype=np.int32)
+ mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused)
+ _assert_eq(dist_np, mj_dist, "dist")
+
+ def test_ray_capsule(self):
+ """Tests ray<>capsule matches MuJoCo."""
+ mjm, mjd, m, d = test_util.fixture("ray.xml")
+
+ # looking down at capsule at a slight angle
+ pnt = wp.array([wp.vec3(0.5, 1.0, 1.6)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(0.0, 0.05, -1.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, 2, "geom_id")
+ pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0]
+ unused = np.zeros(1, dtype=np.int32)
+ mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused)
+ _assert_eq(dist_np, mj_dist, "dist")
+
+ # looking up at capsule from below
+ pnt = wp.array([wp.vec3(-0.5, 1.0, 0.05)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(0.0, 0.05, 1.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, 2, "geom_id")
+ pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0]
+ unused = np.zeros(1, dtype=np.int32)
+ mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused)
+ _assert_eq(dist_np, mj_dist, "dist")
+
+ # looking at cylinder of capsule from the side
+ pnt = wp.array([wp.vec3(0.0, 1.0, 0.75)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(1.0, 0.0, 0.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, 2, "geom_id")
+ pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0]
+ unused = np.zeros(1, dtype=np.int32)
+ mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused)
+ _assert_eq(dist_np, mj_dist, "dist")
+
+ def test_ray_cylinder(self):
+ """Tests ray<>cylinder matches MuJoCo."""
+ mjm, mjd, m, d = test_util.fixture("ray.xml")
+
+ pnt = wp.array([wp.vec3(2.0, 0.0, 0.05)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(0.0, 0.05, 1.0))], dtype=wp.vec3).reshape((1, 1))
+
+ mj_geomid = np.zeros(1, dtype=np.int32)
+ mj_dist = mujoco.mj_ray(mjm, mjd, pnt.numpy()[0, 0], vec.numpy()[0, 0], None, 1, -1, mj_geomid)
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+
+ _assert_eq(geomid.numpy()[0, 0], mj_geomid[0], "geomid")
+ _assert_eq(dist.numpy()[0, 0], mj_dist, "dist")
+
+ def test_ray_box(self):
+ """Tests ray<>box matches MuJoCo."""
+ mjm, mjd, m, d = test_util.fixture("ray.xml")
+
+ # looking down at box at a slight angle
+ pnt = wp.array([wp.vec3(1.0, 0.0, 1.6)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(0.0, 0.05, -1.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, 3, "geom_id")
+ pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0]
+ unused = np.zeros(1, dtype=np.int32)
+ mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused)
+ _assert_eq(dist_np, mj_dist, "dist")
+
+ # looking up at box from below
+ pnt = wp.array([wp.vec3(1.0, 0.0, 0.05)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(0.0, 0.05, 1.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, 3, "geom_id")
+ pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0]
+ unused = np.zeros(1, dtype=np.int32)
+ mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused)
+ _assert_eq(dist_np, mj_dist, "dist")
+
+ def test_ray_mesh(self):
+ """Tests ray<>mesh matches MuJoCo."""
+ mjm, mjd, m, d = test_util.fixture("ray.xml")
+
+ # look at the tetrahedron
+ pnt = wp.array([wp.vec3(2.0, 2.0, 2.0)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(-1.0, -1.0, -1.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, 4, "geom_id")
+
+ pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0]
+ unused = np.zeros(1, dtype=np.int32)
+ mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused)
+ _assert_eq(dist_np, mj_dist, "dist-tetrahedron")
+
+ # look away from the dodecahedron
+ pnt = wp.array([wp.vec3(4.0, 2.0, 2.0)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(2.0, 1.0, 1.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ _assert_eq(geomid_np, -1, "geom_id")
+
+ # look at the dodecahedron
+ pnt = wp.array([wp.vec3(4.0, 2.0, 2.0)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(-2.0, -1.0, -1.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, 5, "geom_id")
+
+ pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0]
+ unused = np.zeros(1, dtype=np.int32)
+ mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused)
+ _assert_eq(dist_np, mj_dist, "dist-dodecahedron")
+
+ def test_ray_hfield(self):
+ mjm, mjd, m, d = test_util.fixture("ray.xml")
+
+ pnt = wp.array([wp.vec3(0.0, 2.0, 2.0)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.vec3(0.0, 0.0, -1.0)], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+
+ mj_geomid = np.zeros(1, dtype=np.int32)
+ mj_dist = mujoco.mj_ray(mjm, mjd, pnt.numpy()[0, 0], vec.numpy()[0, 0], None, 1, -1, mj_geomid)
+
+ _assert_eq(dist.numpy()[0, 0], mj_dist, "dist")
+ _assert_eq(geomid.numpy()[0, 0], mj_geomid[0], "geomid")
+
+ def test_ray_geomgroup(self):
+ """Tests ray geomgroup filter."""
+ mjm, mjd, m, d = test_util.fixture("ray.xml")
+
+ # hits plane with geom_group[0] = 1
+ pnt = wp.array([wp.vec3(2.0, 1.0, 3.0)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1))
+ geomgroup = vec6(1, 0, 0, 0, 0, 0)
+ dist, geomid = mjwarp.ray(m, d, pnt, vec, geomgroup=geomgroup)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, 0, "geom_id")
+
+ pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0]
+ unused = np.zeros(1, dtype=np.int32)
+ mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused)
+ _assert_eq(dist_np, mj_dist, "dist")
+
+ # nothing hit with geom_group[0] = 0
+ pnt = wp.array([wp.vec3(2.0, 1.0, 3.0)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1))
+ geomgroup = vec6(0, 0, 0, 0, 0, 0)
+ dist, geomid = mjwarp.ray(m, d, pnt, vec, geomgroup=geomgroup)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, -1, "geom_id")
+ _assert_eq(dist_np, -1, "dist")
+
+ def test_ray_flg_static(self):
+ """Tests ray flg_static filter."""
+ mjm, mjd, m, d = test_util.fixture("ray.xml")
+
+ # nothing hit with flg_static = False
+ pnt = wp.array([wp.vec3(2.0, 1.0, 3.0)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec, flg_static=False)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, -1, "geom_id")
+ _assert_eq(dist_np, -1, "dist")
+
+ def test_ray_bodyexclude(self):
+ """Tests ray bodyexclude filter."""
+ mjm, mjd, m, d = test_util.fixture("ray.xml")
+
+ # nothing hit with bodyexclude = 0 (world body)
+ pnt = wp.array([wp.vec3(2.0, 1.0, 3.0)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec, bodyexclude=0)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, -1, "geom_id")
+ _assert_eq(dist_np, -1, "dist")
+
+ def test_ray_invisible(self):
+ """Tests ray doesn't hit transparent geoms."""
+ mjm, mjd, m, d = test_util.fixture("ray.xml")
+
+ # nothing hit with transparent geoms
+ m.geom_rgba = wp.array2d([[wp.vec4(0.0, 0.0, 0.0, 0.0)] * 8], dtype=wp.vec4)
+ mujoco.mj_forward(mjm, mjd)
+
+ pnt = wp.array([wp.vec3(2.0, 1.0, 3.0)], dtype=wp.vec3).reshape((1, 1))
+ vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1))
+ dist, geomid = mjwarp.ray(m, d, pnt, vec)
+ wp.synchronize()
+ geomid_np = geomid.numpy()[0, 0]
+ dist_np = dist.numpy()[0, 0]
+ _assert_eq(geomid_np, -1, "geom_id")
+ _assert_eq(dist_np, -1, "dist")
+
+
+if __name__ == "__main__":
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py
new file mode 100644
index 00000000..ecf5f322
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py
@@ -0,0 +1,2113 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+from typing import Any, Tuple
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src import math
+from mujoco.mjx.third_party.mujoco_warp._src import ray
+from mujoco.mjx.third_party.mujoco_warp._src import smooth
+from mujoco.mjx.third_party.mujoco_warp._src import support
+from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
+from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType
+from mujoco.mjx.third_party.mujoco_warp._src.types import ConstraintType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import DataType
+from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import JointType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+from mujoco.mjx.third_party.mujoco_warp._src.types import ObjType
+from mujoco.mjx.third_party.mujoco_warp._src.types import SensorType
+from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec6
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel
+
+
+@wp.func
+def _write_scalar(
+ # Model:
+ sensor_datatype: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_cutoff: wp.array(dtype=float),
+ # In:
+ sensorid: int,
+ sensor: Any,
+ # Out:
+ out: wp.array(dtype=float),
+):
+ adr = sensor_adr[sensorid]
+ cutoff = sensor_cutoff[sensorid]
+
+ if cutoff > 0.0:
+ datatype = sensor_datatype[sensorid]
+ if datatype == int(DataType.REAL.value):
+ out[adr] = wp.clamp(sensor, -cutoff, cutoff)
+ elif datatype == int(DataType.POSITIVE.value):
+ out[adr] = wp.min(sensor, cutoff)
+ else:
+ out[adr] = sensor
+
+
+@wp.func
+def _write_vector(
+ # Model:
+ sensor_datatype: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_cutoff: wp.array(dtype=float),
+ # In:
+ sensorid: int,
+ sensordim: int,
+ sensor: Any,
+ # Out:
+ out: wp.array(dtype=float),
+):
+ adr = sensor_adr[sensorid]
+ cutoff = sensor_cutoff[sensorid]
+
+ if cutoff > 0.0:
+ datatype = sensor_datatype[sensorid]
+ if datatype == int(DataType.REAL.value):
+ for i in range(sensordim):
+ out[adr + i] = wp.clamp(sensor[i], -cutoff, cutoff)
+ elif datatype == int(DataType.POSITIVE.value):
+ for i in range(sensordim):
+ out[adr + i] = wp.min(sensor[i], cutoff)
+ else:
+ for i in range(sensordim):
+ out[adr + i] = sensor[i]
+
+
+@wp.func
+def _magnetometer(
+ # Model:
+ opt_magnetic: wp.array(dtype=wp.vec3),
+ # Data in:
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ # In:
+ worldid: int,
+ objid: int,
+) -> wp.vec3:
+ magnetic = opt_magnetic[worldid]
+ return wp.transpose(site_xmat_in[worldid, objid]) @ magnetic
+
+
+@wp.func
+def _cam_projection(
+ # Model:
+ cam_fovy: wp.array(dtype=float),
+ cam_resolution: wp.array(dtype=wp.vec2i),
+ cam_sensorsize: wp.array(dtype=wp.vec2),
+ cam_intrinsic: wp.array(dtype=wp.vec4),
+ # Data in:
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ cam_xpos_in: wp.array2d(dtype=wp.vec3),
+ cam_xmat_in: wp.array2d(dtype=wp.mat33),
+ # In:
+ worldid: int,
+ objid: int,
+ refid: int,
+) -> wp.vec2:
+ sensorsize = cam_sensorsize[refid]
+ intrinsic = cam_intrinsic[refid]
+ fovy = cam_fovy[refid]
+ res = cam_resolution[refid]
+
+ target_xpos = site_xpos_in[worldid, objid]
+ xpos = cam_xpos_in[worldid, refid]
+ xmat = cam_xmat_in[worldid, refid]
+
+ translation = wp.mat44f(1.0, 0.0, 0.0, -xpos[0], 0.0, 1.0, 0.0, -xpos[1], 0.0, 0.0, 1.0, -xpos[2], 0.0, 0.0, 0.0, 1.0)
+ rotation = wp.mat44f(
+ xmat[0, 0], xmat[1, 0], xmat[2, 0], 0.0,
+ xmat[0, 1], xmat[1, 1], xmat[2, 1], 0.0,
+ xmat[0, 2], xmat[1, 2], xmat[2, 2], 0.0,
+ 0.0, 0.0, 0.0, 1.0,
+ ) # fmt: skip
+
+ # focal transformation matrix (3 x 4)
+ if sensorsize[0] != 0.0 and sensorsize[1] != 0.0:
+ fx = intrinsic[0] / (sensorsize[0] + MJ_MINVAL) * float(res[0])
+ fy = intrinsic[1] / (sensorsize[1] + MJ_MINVAL) * float(res[1])
+ focal = wp.mat44f(-fx, 0.0, 0.0, 0.0, 0.0, fy, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0)
+ else:
+ f = 0.5 / wp.tan(fovy * wp.static(wp.pi / 360.0)) * float(res[1])
+ focal = wp.mat44f(-f, 0.0, 0.0, 0.0, 0.0, f, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0)
+
+ # image matrix (3 x 3)
+ image = wp.mat44f(
+ 1.0, 0.0, 0.5 * float(res[0]), 0.0, 0.0, 1.0, 0.5 * float(res[1]), 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0
+ )
+
+ # projection matrix (3 x 4): product of all 4 matrices
+ # TODO(team): compute proj directly
+ proj = image @ focal @ rotation @ translation
+
+ # projection matrix multiples homogeneous [x, y, z, 1] vectors
+ pos_hom = wp.vec4(target_xpos[0], target_xpos[1], target_xpos[2], 1.0)
+
+ # project world coordinates into pixel space, see:
+ # https://en.wikipedia.org/wiki/3D_projection#Mathematical_formula
+ pixel_coord_hom = proj @ pos_hom
+
+ # avoid dividing by tiny numbers
+ denom = pixel_coord_hom[2]
+ if wp.abs(denom) < MJ_MINVAL:
+ denom = wp.clamp(denom, -MJ_MINVAL, MJ_MINVAL)
+
+ # compute projection
+ return wp.vec2f(pixel_coord_hom[0], pixel_coord_hom[1]) / denom
+
+
+@wp.kernel
+def _sensor_rangefinder_init(
+ # Model:
+ sensor_objid: wp.array(dtype=int),
+ sensor_rangefinder_adr: wp.array(dtype=int),
+ # Data in:
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ # Data out:
+ sensor_rangefinder_pnt_out: wp.array2d(dtype=wp.vec3),
+ sensor_rangefinder_vec_out: wp.array2d(dtype=wp.vec3),
+):
+ worldid, rfid = wp.tid()
+ sensorid = sensor_rangefinder_adr[rfid]
+ objid = sensor_objid[sensorid]
+ site_xpos = site_xpos_in[worldid, objid]
+ site_xmat = site_xmat_in[worldid, objid]
+
+ sensor_rangefinder_pnt_out[worldid, rfid] = site_xpos
+ sensor_rangefinder_vec_out[worldid, rfid] = wp.vec3(site_xmat[0, 2], site_xmat[1, 2], site_xmat[2, 2])
+
+
+@wp.func
+def _joint_pos(jnt_qposadr: wp.array(dtype=int), qpos_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float:
+ return qpos_in[worldid, jnt_qposadr[objid]]
+
+
+@wp.func
+def _tendon_pos(ten_length_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float:
+ return ten_length_in[worldid, objid]
+
+
+@wp.func
+def _actuator_pos(actuator_length_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float:
+ return actuator_length_in[worldid, objid]
+
+
+@wp.func
+def _ball_quat(jnt_qposadr: wp.array(dtype=int), qpos_in: wp.array2d(dtype=float), worldid: int, objid: int) -> wp.quat:
+ adr = jnt_qposadr[objid]
+ quat = wp.quat(
+ qpos_in[worldid, adr + 0],
+ qpos_in[worldid, adr + 1],
+ qpos_in[worldid, adr + 2],
+ qpos_in[worldid, adr + 3],
+ )
+ return wp.normalize(quat)
+
+
+@wp.kernel
+def _limit_pos_zero(
+ # Model:
+ sensor_adr: wp.array(dtype=int),
+ sensor_limitpos_adr: wp.array(dtype=int),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, limitposid = wp.tid()
+ sensordata_out[worldid, sensor_adr[sensor_limitpos_adr[limitposid]]] = 0.0
+
+
+@wp.kernel
+def _limit_pos(
+ # Model:
+ sensor_datatype: wp.array(dtype=int),
+ sensor_objid: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_cutoff: wp.array(dtype=float),
+ sensor_limitpos_adr: wp.array(dtype=int),
+ # Data in:
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nl_in: wp.array(dtype=int),
+ efc_type_in: wp.array2d(dtype=int),
+ efc_id_in: wp.array2d(dtype=int),
+ efc_pos_in: wp.array2d(dtype=float),
+ efc_margin_in: wp.array2d(dtype=float),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, efcid, limitposid = wp.tid()
+
+ ne = ne_in[worldid]
+ nf = nf_in[worldid]
+ nl = nl_in[worldid]
+
+ # skip if not limit
+ if efcid < ne + nf or efcid >= ne + nf + nl:
+ return
+
+ sensorid = sensor_limitpos_adr[limitposid]
+ if efc_id_in[worldid, efcid] == sensor_objid[sensorid]:
+ efc_type = efc_type_in[worldid, efcid]
+ if efc_type == int(ConstraintType.LIMIT_JOINT.value) or efc_type == int(ConstraintType.LIMIT_TENDON.value):
+ val = efc_pos_in[worldid, efcid] - efc_margin_in[worldid, efcid]
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, sensordata_out[worldid])
+
+
+@wp.func
+def _frame_pos(
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xmat_in: wp.array2d(dtype=wp.mat33),
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ ximat_in: wp.array2d(dtype=wp.mat33),
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ cam_xpos_in: wp.array2d(dtype=wp.vec3),
+ cam_xmat_in: wp.array2d(dtype=wp.mat33),
+ # In:
+ worldid: int,
+ objid: int,
+ objtype: int,
+ refid: int,
+ reftype: int,
+) -> wp.vec3:
+ if objtype == int(ObjType.BODY.value):
+ xpos = xipos_in[worldid, objid]
+ elif objtype == int(ObjType.XBODY.value):
+ xpos = xpos_in[worldid, objid]
+ elif objtype == int(ObjType.GEOM.value):
+ xpos = geom_xpos_in[worldid, objid]
+ elif objtype == int(ObjType.SITE.value):
+ xpos = site_xpos_in[worldid, objid]
+ elif objtype == int(ObjType.CAMERA.value):
+ xpos = cam_xpos_in[worldid, objid]
+ else: # UNKNOWN
+ xpos = wp.vec3(0.0)
+
+ if refid == -1:
+ return xpos
+
+ if reftype == int(ObjType.BODY.value):
+ xpos_ref = xipos_in[worldid, refid]
+ xmat_ref = ximat_in[worldid, refid]
+ elif objtype == int(ObjType.XBODY.value):
+ xpos_ref = xpos_in[worldid, refid]
+ xmat_ref = xmat_in[worldid, refid]
+ elif reftype == int(ObjType.GEOM.value):
+ xpos_ref = geom_xpos_in[worldid, refid]
+ xmat_ref = geom_xmat_in[worldid, refid]
+ elif reftype == int(ObjType.SITE.value):
+ xpos_ref = site_xpos_in[worldid, refid]
+ xmat_ref = site_xmat_in[worldid, refid]
+ elif reftype == int(ObjType.CAMERA.value):
+ xpos_ref = cam_xpos_in[worldid, refid]
+ xmat_ref = cam_xmat_in[worldid, refid]
+
+ else: # UNKNOWN
+ xpos_ref = wp.vec3(0.0)
+ xmat_ref = wp.identity(3, wp.float32)
+
+ return wp.transpose(xmat_ref) @ (xpos - xpos_ref)
+
+
+@wp.func
+def _frame_axis(
+ # Data in:
+ xmat_in: wp.array2d(dtype=wp.mat33),
+ ximat_in: wp.array2d(dtype=wp.mat33),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ cam_xmat_in: wp.array2d(dtype=wp.mat33),
+ # In:
+ worldid: int,
+ objid: int,
+ objtype: int,
+ refid: int,
+ reftype: int,
+ frame_axis: int,
+) -> wp.vec3:
+ if objtype == int(ObjType.BODY.value):
+ xmat = ximat_in[worldid, objid]
+ axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
+ elif objtype == int(ObjType.XBODY.value):
+ xmat = xmat_in[worldid, objid]
+ axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
+ elif objtype == int(ObjType.GEOM.value):
+ xmat = geom_xmat_in[worldid, objid]
+ axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
+ elif objtype == int(ObjType.SITE.value):
+ xmat = site_xmat_in[worldid, objid]
+ axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
+ elif objtype == int(ObjType.CAMERA.value):
+ xmat = cam_xmat_in[worldid, objid]
+ axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
+ else: # UNKNOWN
+ axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
+
+ if refid == -1:
+ return axis
+
+ if reftype == int(ObjType.BODY.value):
+ xmat_ref = ximat_in[worldid, refid]
+ elif reftype == int(ObjType.XBODY.value):
+ xmat_ref = xmat_in[worldid, refid]
+ elif reftype == int(ObjType.GEOM.value):
+ xmat_ref = geom_xmat_in[worldid, refid]
+ elif reftype == int(ObjType.SITE.value):
+ xmat_ref = site_xmat_in[worldid, refid]
+ elif reftype == int(ObjType.CAMERA.value):
+ xmat_ref = cam_xmat_in[worldid, refid]
+ else: # UNKNOWN
+ xmat_ref = wp.identity(3, dtype=wp.float32)
+
+ return wp.transpose(xmat_ref) @ axis
+
+
+@wp.func
+def _frame_quat(
+ # Model:
+ body_iquat: wp.array2d(dtype=wp.quat),
+ geom_bodyid: wp.array(dtype=int),
+ geom_quat: wp.array2d(dtype=wp.quat),
+ site_bodyid: wp.array(dtype=int),
+ site_quat: wp.array2d(dtype=wp.quat),
+ cam_bodyid: wp.array(dtype=int),
+ cam_quat: wp.array2d(dtype=wp.quat),
+ # Data in:
+ xquat_in: wp.array2d(dtype=wp.quat),
+ # In:
+ worldid: int,
+ objid: int,
+ objtype: int,
+ refid: int,
+ reftype: int,
+) -> wp.quat:
+ if objtype == int(ObjType.BODY.value):
+ quat = math.mul_quat(xquat_in[worldid, objid], body_iquat[worldid, objid])
+ elif objtype == int(ObjType.XBODY.value):
+ quat = xquat_in[worldid, objid]
+ elif objtype == int(ObjType.GEOM.value):
+ quat = math.mul_quat(xquat_in[worldid, geom_bodyid[objid]], geom_quat[worldid, objid])
+ elif objtype == int(ObjType.SITE.value):
+ quat = math.mul_quat(xquat_in[worldid, site_bodyid[objid]], site_quat[worldid, objid])
+ elif objtype == int(ObjType.CAMERA.value):
+ quat = math.mul_quat(xquat_in[worldid, cam_bodyid[objid]], cam_quat[worldid, objid])
+ else: # UNKNOWN
+ quat = wp.quat(1.0, 0.0, 0.0, 0.0)
+
+ if refid == -1:
+ return quat
+
+ if reftype == int(ObjType.BODY.value):
+ refquat = math.mul_quat(xquat_in[worldid, refid], body_iquat[worldid, refid])
+ elif reftype == int(ObjType.XBODY.value):
+ refquat = xquat_in[worldid, refid]
+ elif reftype == int(ObjType.GEOM.value):
+ refquat = math.mul_quat(xquat_in[worldid, geom_bodyid[refid]], geom_quat[worldid, refid])
+ elif reftype == int(ObjType.SITE.value):
+ refquat = math.mul_quat(xquat_in[worldid, site_bodyid[refid]], site_quat[worldid, refid])
+ elif reftype == int(ObjType.CAMERA.value):
+ refquat = math.mul_quat(xquat_in[worldid, cam_bodyid[refid]], cam_quat[worldid, refid])
+ else: # UNKNOWN
+ refquat = wp.quat(1.0, 0.0, 0.0, 0.0)
+
+ return math.mul_quat(math.quat_inv(refquat), quat)
+
+
+@wp.func
+def _subtree_com(subtree_com_in: wp.array2d(dtype=wp.vec3), worldid: int, objid: int) -> wp.vec3:
+ return subtree_com_in[worldid, objid]
+
+
+@wp.func
+def _clock(time_in: wp.array(dtype=float), worldid: int) -> float:
+ return time_in[worldid]
+
+
+@wp.kernel
+def _sensor_pos(
+ # Model:
+ opt_magnetic: wp.array(dtype=wp.vec3),
+ body_iquat: wp.array2d(dtype=wp.quat),
+ jnt_qposadr: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ geom_quat: wp.array2d(dtype=wp.quat),
+ site_bodyid: wp.array(dtype=int),
+ site_quat: wp.array2d(dtype=wp.quat),
+ cam_bodyid: wp.array(dtype=int),
+ cam_quat: wp.array2d(dtype=wp.quat),
+ cam_fovy: wp.array(dtype=float),
+ cam_resolution: wp.array(dtype=wp.vec2i),
+ cam_sensorsize: wp.array(dtype=wp.vec2),
+ cam_intrinsic: wp.array(dtype=wp.vec4),
+ sensor_type: wp.array(dtype=int),
+ sensor_datatype: wp.array(dtype=int),
+ sensor_objtype: wp.array(dtype=int),
+ sensor_objid: wp.array(dtype=int),
+ sensor_reftype: wp.array(dtype=int),
+ sensor_refid: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_cutoff: wp.array(dtype=float),
+ sensor_pos_adr: wp.array(dtype=int),
+ rangefinder_sensor_adr: wp.array(dtype=int),
+ # Data in:
+ time_in: wp.array(dtype=float),
+ energy_in: wp.array(dtype=wp.vec2),
+ qpos_in: wp.array2d(dtype=float),
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xquat_in: wp.array2d(dtype=wp.quat),
+ xmat_in: wp.array2d(dtype=wp.mat33),
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ ximat_in: wp.array2d(dtype=wp.mat33),
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ cam_xpos_in: wp.array2d(dtype=wp.vec3),
+ cam_xmat_in: wp.array2d(dtype=wp.mat33),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ actuator_length_in: wp.array2d(dtype=float),
+ ten_length_in: wp.array2d(dtype=float),
+ sensor_rangefinder_dist_in: wp.array2d(dtype=float),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, posid = wp.tid()
+ sensorid = sensor_pos_adr[posid]
+ sensortype = sensor_type[sensorid]
+ objid = sensor_objid[sensorid]
+ out = sensordata_out[worldid]
+
+ if sensortype == int(SensorType.MAGNETOMETER.value):
+ vec3 = _magnetometer(opt_magnetic, site_xmat_in, worldid, objid)
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+ elif sensortype == int(SensorType.CAMPROJECTION.value):
+ refid = sensor_refid[sensorid]
+ vec2 = _cam_projection(
+ cam_fovy, cam_resolution, cam_sensorsize, cam_intrinsic, site_xpos_in, cam_xpos_in, cam_xmat_in, worldid, objid, refid
+ )
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 2, vec2, out)
+ elif sensortype == int(SensorType.RANGEFINDER.value):
+ val = sensor_rangefinder_dist_in[worldid, rangefinder_sensor_adr[sensorid]]
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out)
+ elif sensortype == int(SensorType.JOINTPOS.value):
+ val = _joint_pos(jnt_qposadr, qpos_in, worldid, objid)
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out)
+ elif sensortype == int(SensorType.TENDONPOS.value):
+ val = _tendon_pos(ten_length_in, worldid, objid)
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out)
+ elif sensortype == int(SensorType.ACTUATORPOS.value):
+ val = _actuator_pos(actuator_length_in, worldid, objid)
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out)
+ elif sensortype == int(SensorType.BALLQUAT.value):
+ quat = _ball_quat(jnt_qposadr, qpos_in, worldid, objid)
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 4, quat, out)
+ elif sensortype == int(SensorType.FRAMEPOS.value):
+ objtype = sensor_objtype[sensorid]
+ refid = sensor_refid[sensorid]
+ reftype = sensor_reftype[sensorid]
+ vec3 = _frame_pos(
+ xpos_in,
+ xmat_in,
+ xipos_in,
+ ximat_in,
+ geom_xpos_in,
+ geom_xmat_in,
+ site_xpos_in,
+ site_xmat_in,
+ cam_xpos_in,
+ cam_xmat_in,
+ worldid,
+ objid,
+ objtype,
+ refid,
+ reftype,
+ )
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+ elif (
+ sensortype == int(SensorType.FRAMEXAXIS.value)
+ or sensortype == int(SensorType.FRAMEYAXIS.value)
+ or sensortype == int(SensorType.FRAMEZAXIS.value)
+ ):
+ objtype = sensor_objtype[sensorid]
+ refid = sensor_refid[sensorid]
+ reftype = sensor_reftype[sensorid]
+ if sensortype == int(SensorType.FRAMEXAXIS.value):
+ axis = 0
+ elif sensortype == int(SensorType.FRAMEYAXIS.value):
+ axis = 1
+ elif sensortype == int(SensorType.FRAMEZAXIS.value):
+ axis = 2
+ vec3 = _frame_axis(
+ ximat_in, xmat_in, geom_xmat_in, site_xmat_in, cam_xmat_in, worldid, objid, objtype, refid, reftype, axis
+ )
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+ elif sensortype == int(SensorType.FRAMEQUAT.value):
+ objtype = sensor_objtype[sensorid]
+ refid = sensor_refid[sensorid]
+ reftype = sensor_reftype[sensorid]
+ quat = _frame_quat(
+ body_iquat,
+ geom_bodyid,
+ geom_quat,
+ site_bodyid,
+ site_quat,
+ cam_bodyid,
+ cam_quat,
+ xquat_in,
+ worldid,
+ objid,
+ objtype,
+ refid,
+ reftype,
+ )
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 4, quat, out)
+ elif sensortype == int(SensorType.SUBTREECOM.value):
+ vec3 = _subtree_com(subtree_com_in, worldid, objid)
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+ elif sensortype == int(SensorType.E_POTENTIAL.value):
+ val = energy_in[worldid][0]
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out)
+ elif sensortype == int(SensorType.E_KINETIC.value):
+ val = energy_in[worldid][1]
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out)
+ elif sensortype == int(SensorType.CLOCK.value):
+ val = _clock(time_in, worldid)
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out)
+
+
+@event_scope
+def sensor_pos(m: Model, d: Data):
+ """Compute position-dependent sensor values."""
+
+ if m.opt.disableflags & DisableBit.SENSOR:
+ return
+
+ # rangefinder
+ if m.sensor_rangefinder_adr.size > 0:
+ # get position and direction
+ wp.launch(
+ _sensor_rangefinder_init,
+ dim=(d.nworld, m.sensor_rangefinder_adr.size),
+ inputs=[
+ m.sensor_objid,
+ m.sensor_rangefinder_adr,
+ d.site_xpos,
+ d.site_xmat,
+ ],
+ outputs=[
+ d.sensor_rangefinder_pnt,
+ d.sensor_rangefinder_vec,
+ ],
+ )
+
+ # get distances
+ ray.rays(
+ m,
+ d,
+ d.sensor_rangefinder_pnt,
+ d.sensor_rangefinder_vec,
+ vec6(wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf),
+ True,
+ m.sensor_rangefinder_bodyid,
+ d.sensor_rangefinder_dist,
+ d.sensor_rangefinder_geomid,
+ )
+
+ if m.sensor_e_potential:
+ energy_pos(m, d)
+
+ if m.sensor_e_kinetic:
+ energy_vel(m, d)
+
+ wp.launch(
+ _sensor_pos,
+ dim=(d.nworld, m.sensor_pos_adr.size),
+ inputs=[
+ m.opt.magnetic,
+ m.body_iquat,
+ m.jnt_qposadr,
+ m.geom_bodyid,
+ m.geom_quat,
+ m.site_bodyid,
+ m.site_quat,
+ m.cam_bodyid,
+ m.cam_quat,
+ m.cam_fovy,
+ m.cam_resolution,
+ m.cam_sensorsize,
+ m.cam_intrinsic,
+ m.sensor_type,
+ m.sensor_datatype,
+ m.sensor_objtype,
+ m.sensor_objid,
+ m.sensor_reftype,
+ m.sensor_refid,
+ m.sensor_adr,
+ m.sensor_cutoff,
+ m.sensor_pos_adr,
+ m.rangefinder_sensor_adr,
+ d.time,
+ d.energy,
+ d.qpos,
+ d.xpos,
+ d.xquat,
+ d.xmat,
+ d.xipos,
+ d.ximat,
+ d.geom_xpos,
+ d.geom_xmat,
+ d.site_xpos,
+ d.site_xmat,
+ d.cam_xpos,
+ d.cam_xmat,
+ d.subtree_com,
+ d.actuator_length,
+ d.ten_length,
+ d.sensor_rangefinder_dist,
+ ],
+ outputs=[d.sensordata],
+ )
+
+ # jointlimitpos and tendonlimitpos
+ wp.launch(
+ _limit_pos_zero,
+ dim=(d.nworld, m.sensor_limitpos_adr.size),
+ inputs=[m.sensor_adr, m.sensor_limitpos_adr],
+ outputs=[d.sensordata],
+ )
+
+ wp.launch(
+ _limit_pos,
+ dim=(d.nworld, d.njmax, m.sensor_limitpos_adr.size),
+ inputs=[
+ m.sensor_datatype,
+ m.sensor_objid,
+ m.sensor_adr,
+ m.sensor_cutoff,
+ m.sensor_limitpos_adr,
+ d.ne,
+ d.nf,
+ d.nl,
+ d.efc.type,
+ d.efc.id,
+ d.efc.pos,
+ d.efc.margin,
+ ],
+ outputs=[
+ d.sensordata,
+ ],
+ )
+
+
+@wp.func
+def _velocimeter(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ # Data in:
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ worldid: int,
+ objid: int,
+) -> wp.vec3:
+ bodyid = site_bodyid[objid]
+ pos = site_xpos_in[worldid, objid]
+ rot = site_xmat_in[worldid, objid]
+ cvel = cvel_in[worldid, bodyid]
+ ang = wp.spatial_top(cvel)
+ lin = wp.spatial_bottom(cvel)
+ subtree_com = subtree_com_in[worldid, body_rootid[bodyid]]
+ dif = pos - subtree_com
+ return wp.transpose(rot) @ (lin - wp.cross(dif, ang))
+
+
+@wp.func
+def _gyro(
+ # Model:
+ site_bodyid: wp.array(dtype=int),
+ # Data in:
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ worldid: int,
+ objid: int,
+) -> wp.vec3:
+ bodyid = site_bodyid[objid]
+ rot = site_xmat_in[worldid, objid]
+ cvel = cvel_in[worldid, bodyid]
+ ang = wp.spatial_top(cvel)
+ return wp.transpose(rot) @ ang
+
+
+@wp.func
+def _joint_vel(jnt_dofadr: wp.array(dtype=int), qvel_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float:
+ return qvel_in[worldid, jnt_dofadr[objid]]
+
+
+@wp.func
+def _tendon_vel(ten_velocity_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float:
+ return ten_velocity_in[worldid, objid]
+
+
+@wp.func
+def _actuator_vel(actuator_velocity_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float:
+ return actuator_velocity_in[worldid, objid]
+
+
+@wp.func
+def _ball_ang_vel(jnt_dofadr: wp.array(dtype=int), qvel_in: wp.array2d(dtype=float), worldid: int, objid: int) -> wp.vec3:
+ adr = jnt_dofadr[objid]
+ return wp.vec3(qvel_in[worldid, adr + 0], qvel_in[worldid, adr + 1], qvel_in[worldid, adr + 2])
+
+
+@wp.kernel
+def _limit_vel_zero(
+ # Model:
+ sensor_adr: wp.array(dtype=int),
+ sensor_limitvel_adr: wp.array(dtype=int),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, limitvelid = wp.tid()
+ sensordata_out[worldid, sensor_adr[sensor_limitvel_adr[limitvelid]]] = 0.0
+
+
+@wp.kernel
+def _limit_vel(
+ # Model:
+ sensor_datatype: wp.array(dtype=int),
+ sensor_objid: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_cutoff: wp.array(dtype=float),
+ sensor_limitvel_adr: wp.array(dtype=int),
+ # Data in:
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nl_in: wp.array(dtype=int),
+ efc_type_in: wp.array2d(dtype=int),
+ efc_id_in: wp.array2d(dtype=int),
+ efc_vel_in: wp.array2d(dtype=float),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, efcid, limitvelid = wp.tid()
+
+ ne = ne_in[worldid]
+ nf = nf_in[worldid]
+ nl = nl_in[worldid]
+
+ # skip if not limit
+ if efcid < ne + nf or efcid >= ne + nf + nl:
+ return
+
+ sensorid = sensor_limitvel_adr[limitvelid]
+ if efc_id_in[worldid, efcid] == sensor_objid[sensorid]:
+ efc_type = efc_type_in[worldid, efcid]
+ if efc_type == int(ConstraintType.LIMIT_JOINT.value) or efc_type == int(ConstraintType.LIMIT_TENDON.value):
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, efc_vel_in[worldid, efcid], sensordata_out[worldid])
+
+
+@wp.func
+def _cvel_offset(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ cam_bodyid: wp.array(dtype=int),
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ cam_xpos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ worldid: int,
+ objtype: int,
+ objid: int,
+) -> Tuple[wp.spatial_vector, wp.vec3]:
+ if objtype == int(ObjType.BODY.value):
+ pos = xipos_in[worldid, objid]
+ bodyid = objid
+ elif objtype == int(ObjType.XBODY.value):
+ pos = xpos_in[worldid, objid]
+ bodyid = objid
+ elif objtype == int(ObjType.GEOM.value):
+ pos = geom_xpos_in[worldid, objid]
+ bodyid = geom_bodyid[objid]
+ elif objtype == int(ObjType.SITE.value):
+ pos = site_xpos_in[worldid, objid]
+ bodyid = site_bodyid[objid]
+ elif objtype == int(ObjType.CAMERA.value):
+ pos = cam_xpos_in[worldid, objid]
+ bodyid = cam_bodyid[objid]
+ else: # UNKNOWN
+ pos = wp.vec3(0.0)
+ bodyid = 0
+
+ return cvel_in[worldid, bodyid], pos - subtree_com_in[worldid, body_rootid[bodyid]]
+
+
+@wp.func
+def _frame_linvel(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ cam_bodyid: wp.array(dtype=int),
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xmat_in: wp.array2d(dtype=wp.mat33),
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ ximat_in: wp.array2d(dtype=wp.mat33),
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ cam_xpos_in: wp.array2d(dtype=wp.vec3),
+ cam_xmat_in: wp.array2d(dtype=wp.mat33),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ worldid: int,
+ objid: int,
+ objtype: int,
+ refid: int,
+ reftype: int,
+) -> wp.vec3:
+ if objtype == int(ObjType.BODY.value):
+ xpos = xipos_in[worldid, objid]
+ elif objtype == int(ObjType.XBODY.value):
+ xpos = xpos_in[worldid, objid]
+ elif objtype == int(ObjType.GEOM.value):
+ xpos = geom_xpos_in[worldid, objid]
+ elif objtype == int(ObjType.SITE.value):
+ xpos = site_xpos_in[worldid, objid]
+ elif objtype == int(ObjType.CAMERA.value):
+ xpos = cam_xpos_in[worldid, objid]
+ else: # UNKNOWN
+ xpos = wp.vec3(0.0)
+
+ if reftype == int(ObjType.BODY.value):
+ xposref = xipos_in[worldid, refid]
+ xmatref = ximat_in[worldid, refid]
+ elif reftype == int(ObjType.XBODY.value):
+ xposref = xpos_in[worldid, refid]
+ xmatref = xmat_in[worldid, refid]
+ elif reftype == int(ObjType.GEOM.value):
+ xposref = geom_xpos_in[worldid, refid]
+ xmatref = geom_xmat_in[worldid, refid]
+ elif reftype == int(ObjType.SITE.value):
+ xposref = site_xpos_in[worldid, refid]
+ xmatref = site_xmat_in[worldid, refid]
+ elif reftype == int(ObjType.CAMERA.value):
+ xposref = cam_xpos_in[worldid, refid]
+ xmatref = cam_xmat_in[worldid, refid]
+ else: # UNKNOWN
+ xposref = wp.vec3(0.0)
+ xmatref = wp.identity(3, dtype=float)
+
+ cvel, offset = _cvel_offset(
+ body_rootid,
+ geom_bodyid,
+ site_bodyid,
+ cam_bodyid,
+ xpos_in,
+ xipos_in,
+ geom_xpos_in,
+ site_xpos_in,
+ cam_xpos_in,
+ subtree_com_in,
+ cvel_in,
+ worldid,
+ objtype,
+ objid,
+ )
+ cvelref, offsetref = _cvel_offset(
+ body_rootid,
+ geom_bodyid,
+ site_bodyid,
+ cam_bodyid,
+ xpos_in,
+ xipos_in,
+ geom_xpos_in,
+ site_xpos_in,
+ cam_xpos_in,
+ subtree_com_in,
+ cvel_in,
+ worldid,
+ reftype,
+ refid,
+ )
+ clinvel = wp.spatial_bottom(cvel)
+ cangvel = wp.spatial_top(cvel)
+ cangvelref = wp.spatial_top(cvelref)
+ xlinvel = clinvel - wp.cross(offset, cangvel)
+
+ if refid > -1:
+ clinvelref = wp.spatial_bottom(cvelref)
+ xlinvelref = clinvelref - wp.cross(offsetref, cangvelref)
+ rvec = xpos - xposref
+ rel_vel = xlinvel - xlinvelref + wp.cross(rvec, cangvelref)
+ return wp.transpose(xmatref) @ rel_vel
+ else:
+ return xlinvel
+
+
+@wp.func
+def _frame_angvel(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ cam_bodyid: wp.array(dtype=int),
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xmat_in: wp.array2d(dtype=wp.mat33),
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ ximat_in: wp.array2d(dtype=wp.mat33),
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ cam_xpos_in: wp.array2d(dtype=wp.vec3),
+ cam_xmat_in: wp.array2d(dtype=wp.mat33),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ worldid: int,
+ objid: int,
+ objtype: int,
+ refid: int,
+ reftype: int,
+) -> wp.vec3:
+ cvel, _ = _cvel_offset(
+ body_rootid,
+ geom_bodyid,
+ site_bodyid,
+ cam_bodyid,
+ xpos_in,
+ xipos_in,
+ geom_xpos_in,
+ site_xpos_in,
+ cam_xpos_in,
+ subtree_com_in,
+ cvel_in,
+ worldid,
+ objtype,
+ objid,
+ )
+ cangvel = wp.spatial_top(cvel)
+
+ if refid > -1:
+ if reftype == int(ObjType.BODY.value):
+ xmatref = ximat_in[worldid, refid]
+ elif reftype == int(ObjType.XBODY.value):
+ xmatref = xmat_in[worldid, refid]
+ elif reftype == int(ObjType.GEOM.value):
+ xmatref = geom_xmat_in[worldid, refid]
+ elif reftype == int(ObjType.SITE.value):
+ xmatref = site_xmat_in[worldid, refid]
+ elif reftype == int(ObjType.CAMERA.value):
+ xmatref = cam_xmat_in[worldid, refid]
+ else: # UNKNOWN
+ xmatref = wp.identity(3, dtype=float)
+
+ cvelref, _ = _cvel_offset(
+ body_rootid,
+ geom_bodyid,
+ site_bodyid,
+ cam_bodyid,
+ xpos_in,
+ xipos_in,
+ geom_xpos_in,
+ site_xpos_in,
+ cam_xpos_in,
+ subtree_com_in,
+ cvel_in,
+ worldid,
+ reftype,
+ refid,
+ )
+ cangvelref = wp.spatial_top(cvelref)
+
+ return wp.transpose(xmatref) @ (cangvel - cangvelref)
+ else:
+ return cangvel
+
+
+@wp.func
+def _subtree_linvel(subtree_linvel_in: wp.array2d(dtype=wp.vec3), worldid: int, objid: int) -> wp.vec3:
+ return subtree_linvel_in[worldid, objid]
+
+
+@wp.func
+def _subtree_angmom(subtree_angmom_in: wp.array2d(dtype=wp.vec3), worldid: int, objid: int) -> wp.vec3:
+ return subtree_angmom_in[worldid, objid]
+
+
+@wp.kernel
+def _sensor_vel(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ cam_bodyid: wp.array(dtype=int),
+ sensor_type: wp.array(dtype=int),
+ sensor_datatype: wp.array(dtype=int),
+ sensor_objtype: wp.array(dtype=int),
+ sensor_objid: wp.array(dtype=int),
+ sensor_reftype: wp.array(dtype=int),
+ sensor_refid: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_cutoff: wp.array(dtype=float),
+ sensor_vel_adr: wp.array(dtype=int),
+ # Data in:
+ qvel_in: wp.array2d(dtype=float),
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xmat_in: wp.array2d(dtype=wp.mat33),
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ ximat_in: wp.array2d(dtype=wp.mat33),
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ cam_xpos_in: wp.array2d(dtype=wp.vec3),
+ cam_xmat_in: wp.array2d(dtype=wp.mat33),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ ten_velocity_in: wp.array2d(dtype=float),
+ actuator_velocity_in: wp.array2d(dtype=float),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ subtree_linvel_in: wp.array2d(dtype=wp.vec3),
+ subtree_angmom_in: wp.array2d(dtype=wp.vec3),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, velid = wp.tid()
+ sensorid = sensor_vel_adr[velid]
+ sensortype = sensor_type[sensorid]
+ objid = sensor_objid[sensorid]
+ out = sensordata_out[worldid]
+
+ if sensortype == int(SensorType.VELOCIMETER.value):
+ vec3 = _velocimeter(body_rootid, site_bodyid, site_xpos_in, site_xmat_in, subtree_com_in, cvel_in, worldid, objid)
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+ elif sensortype == int(SensorType.GYRO.value):
+ vec3 = _gyro(site_bodyid, site_xmat_in, cvel_in, worldid, objid)
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+ elif sensortype == int(SensorType.JOINTVEL.value):
+ val = _joint_vel(jnt_dofadr, qvel_in, worldid, objid)
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out)
+ elif sensortype == int(SensorType.TENDONVEL.value):
+ val = _tendon_vel(ten_velocity_in, worldid, objid)
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out)
+ elif sensortype == int(SensorType.ACTUATORVEL.value):
+ val = _actuator_vel(actuator_velocity_in, worldid, objid)
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out)
+ elif sensortype == int(SensorType.BALLANGVEL.value):
+ vec3 = _ball_ang_vel(jnt_dofadr, qvel_in, worldid, objid)
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+ elif sensortype == int(SensorType.FRAMELINVEL.value):
+ objtype = sensor_objtype[sensorid]
+ refid = sensor_refid[sensorid]
+ reftype = sensor_reftype[sensorid]
+ frame_linvel = _frame_linvel(
+ body_rootid,
+ geom_bodyid,
+ site_bodyid,
+ cam_bodyid,
+ xpos_in,
+ xmat_in,
+ xipos_in,
+ ximat_in,
+ geom_xpos_in,
+ geom_xmat_in,
+ site_xpos_in,
+ site_xmat_in,
+ cam_xpos_in,
+ cam_xmat_in,
+ subtree_com_in,
+ cvel_in,
+ worldid,
+ objid,
+ objtype,
+ refid,
+ reftype,
+ )
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, frame_linvel, out)
+ elif sensortype == int(SensorType.FRAMEANGVEL.value):
+ objtype = sensor_objtype[sensorid]
+ refid = sensor_refid[sensorid]
+ reftype = sensor_reftype[sensorid]
+ frame_angvel = _frame_angvel(
+ body_rootid,
+ geom_bodyid,
+ site_bodyid,
+ cam_bodyid,
+ xpos_in,
+ xmat_in,
+ xipos_in,
+ ximat_in,
+ geom_xpos_in,
+ geom_xmat_in,
+ site_xpos_in,
+ site_xmat_in,
+ cam_xpos_in,
+ cam_xmat_in,
+ subtree_com_in,
+ cvel_in,
+ worldid,
+ objid,
+ objtype,
+ refid,
+ reftype,
+ )
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, frame_angvel, out)
+ elif sensortype == int(SensorType.SUBTREELINVEL.value):
+ vec3 = _subtree_linvel(subtree_linvel_in, worldid, objid)
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+ elif sensortype == int(SensorType.SUBTREEANGMOM.value):
+ vec3 = _subtree_angmom(subtree_angmom_in, worldid, objid)
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+
+
+@event_scope
+def sensor_vel(m: Model, d: Data):
+ """Compute velocity-dependent sensor values."""
+
+ if m.opt.disableflags & DisableBit.SENSOR:
+ return
+
+ if m.sensor_subtree_vel:
+ smooth.subtree_vel(m, d)
+
+ wp.launch(
+ _sensor_vel,
+ dim=(d.nworld, m.sensor_vel_adr.size),
+ inputs=[
+ m.body_rootid,
+ m.jnt_dofadr,
+ m.geom_bodyid,
+ m.site_bodyid,
+ m.cam_bodyid,
+ m.sensor_type,
+ m.sensor_datatype,
+ m.sensor_objtype,
+ m.sensor_objid,
+ m.sensor_reftype,
+ m.sensor_refid,
+ m.sensor_adr,
+ m.sensor_cutoff,
+ m.sensor_vel_adr,
+ d.qvel,
+ d.xpos,
+ d.xmat,
+ d.xipos,
+ d.ximat,
+ d.geom_xpos,
+ d.geom_xmat,
+ d.site_xpos,
+ d.site_xmat,
+ d.cam_xpos,
+ d.cam_xmat,
+ d.subtree_com,
+ d.ten_velocity,
+ d.actuator_velocity,
+ d.cvel,
+ d.subtree_linvel,
+ d.subtree_angmom,
+ ],
+ outputs=[d.sensordata],
+ )
+
+ wp.launch(
+ _limit_vel_zero,
+ dim=(d.nworld, m.sensor_limitvel_adr.size),
+ inputs=[m.sensor_adr, m.sensor_limitvel_adr],
+ outputs=[d.sensordata],
+ )
+
+ wp.launch(
+ _limit_vel,
+ dim=(d.nworld, d.njmax, m.sensor_limitvel_adr.size),
+ inputs=[
+ m.sensor_datatype,
+ m.sensor_objid,
+ m.sensor_adr,
+ m.sensor_cutoff,
+ m.sensor_limitvel_adr,
+ d.ne,
+ d.nf,
+ d.nl,
+ d.efc.type,
+ d.efc.id,
+ d.efc.vel,
+ ],
+ outputs=[
+ d.sensordata,
+ ],
+ )
+
+
+@wp.func
+def _accelerometer(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ # Data in:
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ cacc_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ worldid: int,
+ objid: int,
+) -> wp.vec3:
+ bodyid = site_bodyid[objid]
+ rot = site_xmat_in[worldid, objid]
+ rotT = wp.transpose(rot)
+ cvel = cvel_in[worldid, bodyid]
+ cvel_top = wp.spatial_top(cvel)
+ cvel_bottom = wp.spatial_bottom(cvel)
+ cacc = cacc_in[worldid, bodyid]
+ cacc_top = wp.spatial_top(cacc)
+ cacc_bottom = wp.spatial_bottom(cacc)
+ dif = site_xpos_in[worldid, objid] - subtree_com_in[worldid, body_rootid[bodyid]]
+ ang = rotT @ cvel_top
+ lin = rotT @ (cvel_bottom - wp.cross(dif, cvel_top))
+ acc = rotT @ (cacc_bottom - wp.cross(dif, cacc_top))
+ correction = wp.cross(ang, lin)
+ return acc + correction
+
+
+@wp.func
+def _force(
+ # Model:
+ site_bodyid: wp.array(dtype=int),
+ # Data in:
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ cfrc_int_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ worldid: int,
+ objid: int,
+) -> wp.vec3:
+ bodyid = site_bodyid[objid]
+ cfrc_int = cfrc_int_in[worldid, bodyid]
+ site_xmat = site_xmat_in[worldid, objid]
+ return wp.transpose(site_xmat) @ wp.spatial_bottom(cfrc_int)
+
+
+@wp.func
+def _torque(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ # Data in:
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cfrc_int_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ worldid: int,
+ objid: int,
+) -> wp.vec3:
+ bodyid = site_bodyid[objid]
+ cfrc_int = cfrc_int_in[worldid, bodyid]
+ site_xmat = site_xmat_in[worldid, objid]
+ dif = site_xpos_in[worldid, objid] - subtree_com_in[worldid, body_rootid[bodyid]]
+ return wp.transpose(site_xmat) @ (wp.spatial_top(cfrc_int) - wp.cross(dif, wp.spatial_bottom(cfrc_int)))
+
+
+@wp.func
+def _actuator_force(actuator_force_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float:
+ return actuator_force_in[worldid, objid]
+
+
+@wp.func
+def _joint_actuator_force(
+ # Model:
+ jnt_dofadr: wp.array(dtype=int),
+ # Data in:
+ qfrc_actuator_in: wp.array2d(dtype=float),
+ # In:
+ worldid: int,
+ objid: int,
+) -> float:
+ return qfrc_actuator_in[worldid, jnt_dofadr[objid]]
+
+
+@wp.kernel
+def _tendon_actuator_force_zero(
+ # Model:
+ sensor_adr: wp.array(dtype=int),
+ sensor_tendonactfrc_adr: wp.array(dtype=int),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, tenactfrcid = wp.tid()
+ sensorid = sensor_tendonactfrc_adr[tenactfrcid]
+ adr = sensor_adr[sensorid]
+ sensordata_out[worldid, adr] = 0.0
+
+
+@wp.kernel
+def _tendon_actuator_force(
+ # Model:
+ actuator_trntype: wp.array(dtype=int),
+ actuator_trnid: wp.array(dtype=wp.vec2i),
+ sensor_objid: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_tendonactfrc_adr: wp.array(dtype=int),
+ # Data in:
+ actuator_force_in: wp.array2d(dtype=float),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, tenactfrcid, actid = wp.tid()
+ sensorid = sensor_tendonactfrc_adr[tenactfrcid]
+
+ if actuator_trntype[actid] == int(TrnType.TENDON.value) and actuator_trnid[actid][0] == sensor_objid[sensorid]:
+ adr = sensor_adr[sensorid]
+ sensordata_out[worldid, adr] += actuator_force_in[worldid, actid]
+
+
+@wp.kernel
+def _tendon_actuator_force_cutoff(
+ # Model:
+ sensor_datatype: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_cutoff: wp.array(dtype=float),
+ sensor_tendonactfrc_adr: wp.array(dtype=int),
+ # Data in:
+ sensordata_in: wp.array2d(dtype=float),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, tenactfrcid = wp.tid()
+ sensorid = sensor_tendonactfrc_adr[tenactfrcid]
+ adr = sensor_adr[sensorid]
+ val = sensordata_in[worldid, adr]
+
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, sensordata_out[worldid])
+
+
+@wp.kernel
+def _limit_frc_zero(
+ # Model:
+ sensor_adr: wp.array(dtype=int),
+ sensor_limitfrc_adr: wp.array(dtype=int),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, limitfrcid = wp.tid()
+ sensordata_out[worldid, sensor_adr[sensor_limitfrc_adr[limitfrcid]]] = 0.0
+
+
+@wp.kernel
+def _limit_frc(
+ # Model:
+ sensor_datatype: wp.array(dtype=int),
+ sensor_objid: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_cutoff: wp.array(dtype=float),
+ sensor_limitfrc_adr: wp.array(dtype=int),
+ # Data in:
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nl_in: wp.array(dtype=int),
+ efc_type_in: wp.array2d(dtype=int),
+ efc_id_in: wp.array2d(dtype=int),
+ efc_force_in: wp.array2d(dtype=float),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, efcid, limitfrcid = wp.tid()
+
+ ne = ne_in[worldid]
+ nf = nf_in[worldid]
+ nl = nl_in[worldid]
+
+ # skip if not limit
+ if efcid < ne + nf or efcid >= ne + nf + nl:
+ return
+
+ sensorid = sensor_limitfrc_adr[limitfrcid]
+ if efc_id_in[worldid, efcid] == sensor_objid[sensorid]:
+ efc_type = efc_type_in[worldid, efcid]
+ if efc_type == int(ConstraintType.LIMIT_JOINT.value) or efc_type == int(ConstraintType.LIMIT_TENDON.value):
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, efc_force_in[worldid, efcid], sensordata_out[worldid])
+
+
+@wp.func
+def _framelinacc(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ cam_bodyid: wp.array(dtype=int),
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ cam_xpos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ cacc_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ worldid: int,
+ objid: int,
+ objtype: int,
+) -> wp.vec3:
+ if objtype == int(ObjType.BODY.value):
+ bodyid = objid
+ pos = xipos_in[worldid, objid]
+ elif objtype == int(ObjType.XBODY.value):
+ bodyid = objid
+ pos = xpos_in[worldid, objid]
+ elif objtype == int(ObjType.GEOM.value):
+ bodyid = geom_bodyid[objid]
+ pos = geom_xpos_in[worldid, objid]
+ elif objtype == int(ObjType.SITE.value):
+ bodyid = site_bodyid[objid]
+ pos = site_xpos_in[worldid, objid]
+ elif objtype == int(ObjType.CAMERA.value):
+ bodyid = cam_bodyid[objid]
+ pos = cam_xpos_in[worldid, objid]
+ else: # UNKNOWN
+ bodyid = 0
+ pos = wp.vec3(0.0)
+
+ cacc = cacc_in[worldid, bodyid]
+ cvel = cvel_in[worldid, bodyid]
+ offset = pos - subtree_com_in[worldid, body_rootid[bodyid]]
+ ang = wp.spatial_top(cvel)
+ lin = wp.spatial_bottom(cvel) - wp.cross(offset, ang)
+ acc = wp.spatial_bottom(cacc) - wp.cross(offset, wp.spatial_top(cacc))
+ correction = wp.cross(ang, lin)
+
+ return acc + correction
+
+
+@wp.func
+def _frameangacc(
+ # Model:
+ geom_bodyid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ cam_bodyid: wp.array(dtype=int),
+ # Data in:
+ cacc_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ worldid: int,
+ objid: int,
+ objtype: int,
+) -> wp.vec3:
+ if objtype == int(ObjType.BODY.value) or objtype == int(ObjType.XBODY.value):
+ bodyid = objid
+ elif objtype == int(ObjType.GEOM.value):
+ bodyid = geom_bodyid[objid]
+ elif objtype == int(ObjType.SITE.value):
+ bodyid = site_bodyid[objid]
+ elif objtype == int(ObjType.CAMERA.value):
+ bodyid = cam_bodyid[objid]
+ else: # UNKNOWN
+ bodyid = 0
+
+ return wp.spatial_top(cacc_in[worldid, bodyid])
+
+
+@wp.kernel
+def _sensor_acc(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ cam_bodyid: wp.array(dtype=int),
+ sensor_type: wp.array(dtype=int),
+ sensor_datatype: wp.array(dtype=int),
+ sensor_objtype: wp.array(dtype=int),
+ sensor_objid: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_cutoff: wp.array(dtype=float),
+ sensor_acc_adr: wp.array(dtype=int),
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ cam_xpos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ actuator_force_in: wp.array2d(dtype=float),
+ qfrc_actuator_in: wp.array2d(dtype=float),
+ cacc_in: wp.array2d(dtype=wp.spatial_vector),
+ cfrc_int_in: wp.array2d(dtype=wp.spatial_vector),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, accid = wp.tid()
+ sensorid = sensor_acc_adr[accid]
+ sensortype = sensor_type[sensorid]
+ objid = sensor_objid[sensorid]
+ out = sensordata_out[worldid]
+
+ if sensortype == int(SensorType.ACCELEROMETER.value):
+ vec3 = _accelerometer(
+ body_rootid, site_bodyid, site_xpos_in, site_xmat_in, subtree_com_in, cvel_in, cacc_in, worldid, objid
+ )
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+ elif sensortype == int(SensorType.FORCE.value):
+ vec3 = _force(site_bodyid, site_xmat_in, cfrc_int_in, worldid, objid)
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+ elif sensortype == int(SensorType.TORQUE.value):
+ vec3 = _torque(body_rootid, site_bodyid, site_xpos_in, site_xmat_in, subtree_com_in, cfrc_int_in, worldid, objid)
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+ elif sensortype == int(SensorType.ACTUATORFRC.value):
+ val = _actuator_force(actuator_force_in, worldid, objid)
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out)
+ elif sensortype == int(SensorType.JOINTACTFRC.value):
+ val = _joint_actuator_force(jnt_dofadr, qfrc_actuator_in, worldid, objid)
+ _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out)
+ elif sensortype == int(SensorType.FRAMELINACC.value):
+ objtype = sensor_objtype[sensorid]
+ vec3 = _framelinacc(
+ body_rootid,
+ geom_bodyid,
+ site_bodyid,
+ cam_bodyid,
+ xpos_in,
+ xipos_in,
+ geom_xpos_in,
+ site_xpos_in,
+ cam_xpos_in,
+ subtree_com_in,
+ cvel_in,
+ cacc_in,
+ worldid,
+ objid,
+ objtype,
+ )
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+ elif sensortype == int(SensorType.FRAMEANGACC.value):
+ objtype = sensor_objtype[sensorid]
+ vec3 = _frameangacc(
+ geom_bodyid,
+ site_bodyid,
+ cam_bodyid,
+ cacc_in,
+ worldid,
+ objid,
+ objtype,
+ )
+ _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
+
+
+@wp.kernel
+def _sensor_touch_zero(
+ # Model:
+ sensor_adr: wp.array(dtype=int),
+ sensor_touch_adr: wp.array(dtype=int),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ worldid, sensortouchadrid = wp.tid()
+ sensorid = sensor_touch_adr[sensortouchadrid]
+ adr = sensor_adr[sensorid]
+ sensordata_out[worldid, adr] = 0.0
+
+
+@wp.kernel
+def _sensor_touch(
+ # Model:
+ opt_cone: int,
+ geom_bodyid: wp.array(dtype=int),
+ site_type: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ site_size: wp.array(dtype=wp.vec3),
+ sensor_objid: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_touch_adr: wp.array(dtype=int),
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ contact_pos_in: wp.array(dtype=wp.vec3),
+ contact_frame_in: wp.array(dtype=wp.mat33),
+ contact_dim_in: wp.array(dtype=int),
+ contact_geom_in: wp.array(dtype=wp.vec2i),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ contact_worldid_in: wp.array(dtype=int),
+ efc_force_in: wp.array2d(dtype=float),
+ # Data out:
+ sensordata_out: wp.array2d(dtype=float),
+):
+ conid, sensortouchadrid = wp.tid()
+
+ if conid > ncon_in[0]:
+ return
+
+ sensorid = sensor_touch_adr[sensortouchadrid]
+
+ objid = sensor_objid[sensorid]
+ bodyid = site_bodyid[objid]
+
+ # find contact in sensor zone, add normal force
+
+ # contacting bodies
+ geom = contact_geom_in[conid]
+ conbody = wp.vec2i(geom_bodyid[geom[0]], geom_bodyid[geom[1]])
+
+ # select contacts involving sensorized body
+ worldid = contact_worldid_in[conid]
+ efc_address0 = contact_efc_address_in[conid, 0]
+ if efc_address0 >= 0 and (bodyid == conbody[0] or bodyid == conbody[1]):
+ # get contact normal force
+ normalforce = efc_force_in[worldid, efc_address0]
+
+ if opt_cone == int(ConeType.PYRAMIDAL.value):
+ dim = contact_dim_in[conid]
+ for i in range(1, 2 * (dim - 1)):
+ normalforce += efc_force_in[worldid, contact_efc_address_in[conid, i]]
+
+ if normalforce <= 0.0:
+ return
+
+ # convert contact normal force to global frame, normalize
+ frame = contact_frame_in[conid]
+ conray = wp.vec3(frame[0, 0], frame[0, 1], frame[0, 2]) * normalforce
+ conray, _ = math.normalize_with_norm(conray)
+
+ # flip ray direction if sensor is on body2
+ if bodyid == conbody[1]:
+ conray = -conray
+
+ # add if ray-zone intersection (always true when contact.pos inside zone)
+ if (
+ ray.ray_geom(
+ site_xpos_in[worldid, objid],
+ site_xmat_in[worldid, objid],
+ site_size[objid],
+ contact_pos_in[conid],
+ conray,
+ site_type[objid],
+ )
+ >= 0.0
+ ):
+ adr = sensor_adr[sensorid]
+ wp.atomic_add(sensordata_out[worldid], adr, normalforce)
+
+
+@event_scope
+def sensor_acc(m: Model, d: Data):
+ """Compute acceleration-dependent sensor values."""
+ if m.opt.disableflags & DisableBit.SENSOR:
+ return
+
+ wp.launch(
+ _sensor_touch_zero,
+ dim=(d.nworld, m.sensor_touch_adr.size),
+ inputs=[
+ m.sensor_adr,
+ m.sensor_touch_adr,
+ ],
+ outputs=[
+ d.sensordata,
+ ],
+ )
+
+ wp.launch(
+ _sensor_touch,
+ dim=(d.nconmax, m.sensor_touch_adr.size),
+ inputs=[
+ m.opt.cone,
+ m.geom_bodyid,
+ m.site_type,
+ m.site_bodyid,
+ m.site_size,
+ m.sensor_objid,
+ m.sensor_adr,
+ m.sensor_touch_adr,
+ d.ncon,
+ d.site_xpos,
+ d.site_xmat,
+ d.contact.pos,
+ d.contact.frame,
+ d.contact.dim,
+ d.contact.geom,
+ d.contact.efc_address,
+ d.contact.worldid,
+ d.efc.force,
+ ],
+ outputs=[
+ d.sensordata,
+ ],
+ )
+
+ if m.sensor_rne_postconstraint:
+ smooth.rne_postconstraint(m, d)
+
+ wp.launch(
+ _sensor_acc,
+ dim=(d.nworld, m.sensor_acc_adr.size),
+ inputs=[
+ m.body_rootid,
+ m.jnt_dofadr,
+ m.geom_bodyid,
+ m.site_bodyid,
+ m.cam_bodyid,
+ m.sensor_type,
+ m.sensor_datatype,
+ m.sensor_objtype,
+ m.sensor_objid,
+ m.sensor_adr,
+ m.sensor_cutoff,
+ m.sensor_acc_adr,
+ d.xpos,
+ d.xipos,
+ d.geom_xpos,
+ d.site_xpos,
+ d.site_xmat,
+ d.cam_xpos,
+ d.subtree_com,
+ d.cvel,
+ d.actuator_force,
+ d.qfrc_actuator,
+ d.cacc,
+ d.cfrc_int,
+ ],
+ outputs=[d.sensordata],
+ )
+
+ wp.launch(
+ _tendon_actuator_force_zero,
+ dim=(d.nworld, m.sensor_tendonactfrc_adr.size),
+ inputs=[
+ m.sensor_adr,
+ m.sensor_tendonactfrc_adr,
+ ],
+ outputs=[
+ d.sensordata,
+ ],
+ )
+
+ wp.launch(
+ _tendon_actuator_force,
+ dim=(d.nworld, m.sensor_tendonactfrc_adr.size, m.nu),
+ inputs=[
+ m.actuator_trntype,
+ m.actuator_trnid,
+ m.sensor_objid,
+ m.sensor_adr,
+ m.sensor_tendonactfrc_adr,
+ d.actuator_force,
+ ],
+ outputs=[
+ d.sensordata,
+ ],
+ )
+
+ wp.launch(
+ _tendon_actuator_force_cutoff,
+ dim=(d.nworld, m.sensor_tendonactfrc_adr.size),
+ inputs=[
+ m.sensor_datatype,
+ m.sensor_adr,
+ m.sensor_cutoff,
+ m.sensor_tendonactfrc_adr,
+ d.sensordata,
+ ],
+ outputs=[d.sensordata],
+ )
+
+ wp.launch(
+ _limit_frc_zero,
+ dim=(d.nworld, m.sensor_limitfrc_adr.size),
+ inputs=[m.sensor_adr, m.sensor_limitfrc_adr],
+ outputs=[d.sensordata],
+ )
+
+ wp.launch(
+ _limit_frc,
+ dim=(d.nworld, d.njmax, m.sensor_limitfrc_adr.size),
+ inputs=[
+ m.sensor_datatype,
+ m.sensor_objid,
+ m.sensor_adr,
+ m.sensor_cutoff,
+ m.sensor_limitfrc_adr,
+ d.ne,
+ d.nf,
+ d.nl,
+ d.efc.type,
+ d.efc.id,
+ d.efc.force,
+ ],
+ outputs=[
+ d.sensordata,
+ ],
+ )
+
+
+@wp.kernel
+def _energy_pos_zero(
+ # Data out:
+ energy_out: wp.array(dtype=wp.vec2),
+):
+ worldid = wp.tid()
+ energy_out[worldid][0] = 0.0
+
+
+@wp.kernel
+def _energy_pos_gravity(
+ # Model:
+ opt_gravity: wp.array(dtype=wp.vec3),
+ body_mass: wp.array2d(dtype=float),
+ # Data in:
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ # Data out:
+ energy_out: wp.array(dtype=wp.vec2),
+):
+ worldid, bodyid = wp.tid()
+ gravity = opt_gravity[worldid]
+ bodyid += 1 # skip world body
+
+ energy = wp.vec2(
+ body_mass[worldid, bodyid] * wp.dot(gravity, xipos_in[worldid, bodyid]),
+ 0.0,
+ )
+
+ wp.atomic_sub(energy_out, worldid, energy)
+
+
+@wp.kernel
+def _energy_pos_passive_joint(
+ # Model:
+ qpos_spring: wp.array2d(dtype=float),
+ jnt_type: wp.array(dtype=int),
+ jnt_qposadr: wp.array(dtype=int),
+ jnt_stiffness: wp.array2d(dtype=float),
+ # Data in:
+ qpos_in: wp.array2d(dtype=float),
+ # Data out:
+ energy_out: wp.array(dtype=wp.vec2),
+):
+ worldid, jntid = wp.tid()
+ stiffness = jnt_stiffness[worldid, jntid]
+
+ if stiffness == 0.0:
+ return
+
+ padr = jnt_qposadr[jntid]
+ jnttype = jnt_type[jntid]
+
+ if jnttype == int(JointType.FREE.value):
+ dif0 = wp.vec3(
+ qpos_in[worldid, padr + 0] - qpos_spring[worldid, padr + 0],
+ qpos_in[worldid, padr + 1] - qpos_spring[worldid, padr + 1],
+ qpos_in[worldid, padr + 2] - qpos_spring[worldid, padr + 2],
+ )
+
+ # convert quaternion difference into angular "velocity"
+ quat1 = wp.quat(
+ qpos_in[worldid, padr + 3],
+ qpos_in[worldid, padr + 4],
+ qpos_in[worldid, padr + 5],
+ qpos_in[worldid, padr + 6],
+ )
+ quat1 = wp.normalize(quat1)
+
+ quat_spring = wp.quat(
+ qpos_spring[worldid, padr + 3],
+ qpos_spring[worldid, padr + 4],
+ qpos_spring[worldid, padr + 5],
+ qpos_spring[worldid, padr + 6],
+ )
+
+ dif1 = math.quat_sub(quat1, quat_spring)
+
+ energy = wp.vec2(
+ 0.5 * stiffness * (wp.dot(dif0, dif0) + wp.dot(dif1, dif1)),
+ 0.0,
+ )
+
+ wp.atomic_add(energy_out, worldid, energy)
+
+ elif jnttype == int(JointType.BALL.value):
+ quat = wp.quat(
+ qpos_in[worldid, padr + 0],
+ qpos_in[worldid, padr + 1],
+ qpos_in[worldid, padr + 2],
+ qpos_in[worldid, padr + 3],
+ )
+ quat = wp.normalize(quat)
+
+ quat_spring = wp.quat(
+ qpos_spring[worldid, padr + 0],
+ qpos_spring[worldid, padr + 1],
+ qpos_spring[worldid, padr + 2],
+ qpos_spring[worldid, padr + 3],
+ )
+
+ dif = math.quat_sub(quat, quat_spring)
+ energy = wp.vec2(
+ 0.5 * stiffness * wp.dot(dif, dif),
+ 0.0,
+ )
+ wp.atomic_add(energy_out, worldid, energy)
+ elif jnttype == int(JointType.SLIDE.value) or jnttype == int(JointType.HINGE.value):
+ dif_ = qpos_in[worldid, padr] - qpos_spring[worldid, padr]
+ energy = wp.vec2(
+ 0.5 * stiffness * dif_ * dif_,
+ 0.0,
+ )
+ wp.atomic_add(energy_out, worldid, energy)
+
+
+@wp.kernel
+def _energy_pos_passive_tendon(
+ # Model:
+ tendon_stiffness: wp.array2d(dtype=float),
+ tendon_lengthspring: wp.array2d(dtype=wp.vec2),
+ # Data in:
+ ten_length_in: wp.array2d(dtype=float),
+ # Data out:
+ energy_out: wp.array(dtype=wp.vec2),
+):
+ worldid, tenid = wp.tid()
+
+ stiffness = tendon_stiffness[worldid, tenid]
+
+ if stiffness == 0.0:
+ return
+
+ length = ten_length_in[worldid, tenid]
+
+ # compute spring displacement
+ lengthspring = tendon_lengthspring[worldid, tenid]
+ lower = lengthspring[0]
+ upper = lengthspring[1]
+
+ if length > upper:
+ displacement = upper - length
+ elif length < lower:
+ displacement = lower - length
+ else:
+ displacement = 0.0
+
+ energy = wp.vec2(0.5 * stiffness * displacement * displacement, 0.0)
+ wp.atomic_add(energy_out, worldid, energy)
+
+
+def energy_pos(m: Model, d: Data):
+ """Position-dependent energy (potential)."""
+ wp.launch(_energy_pos_zero, dim=(d.nworld,), outputs=[d.energy])
+
+ # init potential energy: -sum_i(body_i.mass * dot(gravity, body_i.pos))
+ if not m.opt.disableflags & DisableBit.GRAVITY:
+ wp.launch(
+ _energy_pos_gravity, dim=(d.nworld, m.nbody - 1), inputs=[m.opt.gravity, m.body_mass, d.xipos], outputs=[d.energy]
+ )
+
+ if not m.opt.disableflags & DisableBit.PASSIVE:
+ # add joint-level springs
+ wp.launch(
+ _energy_pos_passive_joint,
+ dim=(d.nworld, m.njnt),
+ inputs=[
+ m.qpos_spring,
+ m.jnt_type,
+ m.jnt_qposadr,
+ m.jnt_stiffness,
+ d.qpos,
+ ],
+ outputs=[d.energy],
+ )
+
+ # add tendon-level springs
+ if m.ntendon:
+ wp.launch(
+ _energy_pos_passive_tendon,
+ dim=(d.nworld, m.ntendon),
+ inputs=[
+ m.tendon_stiffness,
+ m.tendon_lengthspring,
+ d.ten_length,
+ ],
+ outputs=[d.energy],
+ )
+
+ # TODO(team): flex
+
+
+@cache_kernel
+def _energy_vel_kinetic(nv: int):
+ @nested_kernel
+ def energy_vel_kinetic(
+ # Data in:
+ qvel_in: wp.array2d(dtype=float),
+ # In:
+ Mqvel: wp.array2d(dtype=float),
+ # Out:
+ energy_out: wp.array(dtype=wp.vec2),
+ ):
+ worldid = wp.tid()
+
+ qvel_tile = wp.tile_load(qvel_in[worldid], shape=wp.static(nv))
+ Mqvel_tile = wp.tile_load(Mqvel[worldid], shape=wp.static(nv))
+
+ # qvel * (M @ qvel)
+ qvelMqvel_tile = wp.tile_map(wp.mul, qvel_tile, Mqvel_tile)
+
+ # sum(qvel * (M @ qvel))
+ quadratic_tile = wp.tile_reduce(wp.add, qvelMqvel_tile)
+
+ energy_out[worldid][1] = 0.5 * quadratic_tile[0]
+
+ return energy_vel_kinetic
+
+
+def energy_vel(m: Model, d: Data):
+ """Velocity-dependent energy (kinetic)."""
+
+ # kinetic energy: 0.5 * qvel.T @ M @ qvel
+
+ # M @ qvel
+ skip = wp.zeros(d.nworld, dtype=bool)
+ support.mul_m(m, d, d.efc.mv, d.qvel, skip)
+
+ wp.launch_tiled(
+ _energy_vel_kinetic(m.nv),
+ dim=(d.nworld,),
+ inputs=[d.qvel, d.efc.mv],
+ outputs=[d.energy],
+ block_dim=m.block_dim.energy_vel_kinetic,
+ )
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py
new file mode 100644
index 00000000..458dea1f
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py
@@ -0,0 +1,411 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 sensor functions."""
+
+import mujoco
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+from absl.testing import parameterized
+
+import mujoco_warp as mjwarp
+
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+
+# tolerance for difference between MuJoCo and MJWarp calculations - mostly
+# due to float precision
+_TOLERANCE = 5e-5
+
+
+def _assert_eq(a, b, name):
+ tol = _TOLERANCE * 10 # avoid test noise
+ err_msg = f"mismatch: {name}"
+ np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+
+
+class SensorTest(parameterized.TestCase):
+ def test_sensor(self):
+ """Test sensors."""
+ mjm, mjd, m, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ keyframe=0,
+ kick=True,
+ )
+
+ d.sensordata.zero_()
+
+ mjwarp.sensor_pos(m, d)
+ mjwarp.sensor_vel(m, d)
+ mjwarp.sensor_acc(m, d)
+
+ _assert_eq(d.sensordata.numpy()[0], mjd.sensordata, "sensordata")
+
+ def test_rangefinder(self):
+ """Test rangefinder."""
+ for keyframe in range(2):
+ _, mjd, m, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ )
+
+ d.sensordata.zero_()
+
+ mjwarp.sensor_pos(m, d)
+
+ _assert_eq(d.sensordata.numpy()[0], mjd.sensordata, "sensordata")
+
+ def test_touch_sensor(self):
+ """Test touch sensor."""
+ for keyframe in range(2):
+ _, mjd, m, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ keyframe=keyframe,
+ )
+
+ d.sensordata.zero_()
+
+ mjwarp.sensor_acc(m, d)
+
+ _assert_eq(d.sensordata.numpy()[0], mjd.sensordata, "sensordata")
+
+ def test_tendon_sensor(self):
+ """Test tendon sensors."""
+ _, mjd, m, d = test_util.fixture("tendon/fixed.xml", keyframe=0, sparse=False)
+
+ d.sensordata.zero_()
+
+ mjwarp.sensor_pos(m, d)
+ mjwarp.sensor_vel(m, d)
+
+ _assert_eq(d.sensordata.numpy()[0], mjd.sensordata, "sensordata")
+
+ @parameterized.parameters("humanoid/humanoid.xml", "constraints.xml")
+ def test_energy(self, xml):
+ mjm, mjd, m, d = test_util.fixture(xml, constraint=False, kick=True)
+
+ d.energy.zero_()
+
+ mujoco.mj_energyPos(mjm, mjd)
+ mjwarp.energy_pos(m, d)
+
+ _assert_eq(d.energy.numpy()[0][0], mjd.energy[0], "potential energy")
+
+ mujoco.mj_energyVel(mjm, mjd)
+ mjwarp.energy_vel(m, d)
+
+ _assert_eq(d.energy.numpy()[0][1], mjd.energy[1], "kinetic energy")
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py
new file mode 100644
index 00000000..b22ee540
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py
@@ -0,0 +1,3128 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src import math
+from mujoco.mjx.third_party.mujoco_warp._src import support
+from mujoco.mjx.third_party.mujoco_warp._src import util_misc
+from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
+from mujoco.mjx.third_party.mujoco_warp._src.types import CamLightType
+from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import JointType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+from mujoco.mjx.third_party.mujoco_warp._src.types import ObjType
+from mujoco.mjx.third_party.mujoco_warp._src.types import TileSet
+from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType
+from mujoco.mjx.third_party.mujoco_warp._src.types import WrapType
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec5
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec10
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec11
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel
+
+wp.set_module_options({"enable_backward": False})
+
+
+@wp.kernel
+def _kinematics_root(
+ # Data out:
+ xpos_out: wp.array2d(dtype=wp.vec3),
+ xquat_out: wp.array2d(dtype=wp.quat),
+ xmat_out: wp.array2d(dtype=wp.mat33),
+ xipos_out: wp.array2d(dtype=wp.vec3),
+ ximat_out: wp.array2d(dtype=wp.mat33),
+):
+ worldid = wp.tid()
+ xpos_out[worldid, 0] = wp.vec3(0.0)
+ xquat_out[worldid, 0] = wp.quat(1.0, 0.0, 0.0, 0.0)
+ xipos_out[worldid, 0] = wp.vec3(0.0)
+ xmat_out[worldid, 0] = wp.identity(n=3, dtype=wp.float32)
+ ximat_out[worldid, 0] = wp.identity(n=3, dtype=wp.float32)
+
+
+@wp.kernel
+def _kinematics_level(
+ # Model:
+ qpos0: wp.array2d(dtype=float),
+ body_parentid: wp.array(dtype=int),
+ body_jntnum: wp.array(dtype=int),
+ body_jntadr: wp.array(dtype=int),
+ body_pos: wp.array2d(dtype=wp.vec3),
+ body_quat: wp.array2d(dtype=wp.quat),
+ body_ipos: wp.array2d(dtype=wp.vec3),
+ body_iquat: wp.array2d(dtype=wp.quat),
+ jnt_type: wp.array(dtype=int),
+ jnt_qposadr: wp.array(dtype=int),
+ jnt_pos: wp.array2d(dtype=wp.vec3),
+ jnt_axis: wp.array2d(dtype=wp.vec3),
+ # Data in:
+ qpos_in: wp.array2d(dtype=float),
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xquat_in: wp.array2d(dtype=wp.quat),
+ xmat_in: wp.array2d(dtype=wp.mat33),
+ # In:
+ body_tree_: wp.array(dtype=int),
+ # Data out:
+ xpos_out: wp.array2d(dtype=wp.vec3),
+ xquat_out: wp.array2d(dtype=wp.quat),
+ xmat_out: wp.array2d(dtype=wp.mat33),
+ xipos_out: wp.array2d(dtype=wp.vec3),
+ ximat_out: wp.array2d(dtype=wp.mat33),
+ xanchor_out: wp.array2d(dtype=wp.vec3),
+ xaxis_out: wp.array2d(dtype=wp.vec3),
+):
+ worldid, nodeid = wp.tid()
+ bodyid = body_tree_[nodeid]
+ jntadr = body_jntadr[bodyid]
+ jntnum = body_jntnum[bodyid]
+ qpos = qpos_in[worldid]
+
+ if jntnum == 0:
+ # no joints - apply fixed translation and rotation relative to parent
+ pid = body_parentid[bodyid]
+ xpos = (xmat_in[worldid, pid] * body_pos[worldid, bodyid]) + xpos_in[worldid, pid]
+ xquat = math.mul_quat(xquat_in[worldid, pid], body_quat[worldid, bodyid])
+ elif jntnum == 1 and jnt_type[jntadr] == wp.static(JointType.FREE.value):
+ # free joint
+ qadr = jnt_qposadr[jntadr]
+ xpos = wp.vec3(qpos[qadr], qpos[qadr + 1], qpos[qadr + 2])
+ xquat = wp.quat(qpos[qadr + 3], qpos[qadr + 4], qpos[qadr + 5], qpos[qadr + 6])
+ xquat = wp.normalize(xquat)
+ xanchor_out[worldid, jntadr] = xpos
+ xaxis_out[worldid, jntadr] = jnt_axis[worldid, jntadr]
+ else:
+ # regular or no joints
+ # apply fixed translation and rotation relative to parent
+ pid = body_parentid[bodyid]
+ xpos = (xmat_in[worldid, pid] * body_pos[worldid, bodyid]) + xpos_in[worldid, pid]
+ xquat = math.mul_quat(xquat_in[worldid, pid], body_quat[worldid, bodyid])
+
+ for _ in range(jntnum):
+ qadr = jnt_qposadr[jntadr]
+ jnt_type_ = jnt_type[jntadr]
+ jnt_axis_ = jnt_axis[worldid, jntadr]
+ xanchor = math.rot_vec_quat(jnt_pos[worldid, jntadr], xquat) + xpos
+ xaxis = math.rot_vec_quat(jnt_axis_, xquat)
+
+ if jnt_type_ == wp.static(JointType.BALL.value):
+ qloc = wp.quat(
+ qpos[qadr + 0],
+ qpos[qadr + 1],
+ qpos[qadr + 2],
+ qpos[qadr + 3],
+ )
+ qloc = wp.normalize(qloc)
+ xquat = math.mul_quat(xquat, qloc)
+ # correct for off-center rotation
+ xpos = xanchor - math.rot_vec_quat(jnt_pos[worldid, jntadr], xquat)
+ elif jnt_type_ == wp.static(JointType.SLIDE.value):
+ xpos += xaxis * (qpos[qadr] - qpos0[worldid, qadr])
+ elif jnt_type_ == wp.static(JointType.HINGE.value):
+ qpos0_ = qpos0[worldid, qadr]
+ qloc_ = math.axis_angle_to_quat(jnt_axis_, qpos[qadr] - qpos0_)
+ xquat = math.mul_quat(xquat, qloc_)
+ # correct for off-center rotation
+ xpos = xanchor - math.rot_vec_quat(jnt_pos[worldid, jntadr], xquat)
+
+ xanchor_out[worldid, jntadr] = xanchor
+ xaxis_out[worldid, jntadr] = xaxis
+ jntadr += 1
+
+ xpos_out[worldid, bodyid] = xpos
+ xquat_out[worldid, bodyid] = wp.normalize(xquat)
+ xmat_out[worldid, bodyid] = math.quat_to_mat(xquat)
+ xipos_out[worldid, bodyid] = xpos + math.rot_vec_quat(body_ipos[worldid, bodyid], xquat)
+ ximat_out[worldid, bodyid] = math.quat_to_mat(math.mul_quat(xquat, body_iquat[worldid, bodyid]))
+
+
+@wp.kernel
+def _geom_local_to_global(
+ # Model:
+ geom_bodyid: wp.array(dtype=int),
+ geom_pos: wp.array2d(dtype=wp.vec3),
+ geom_quat: wp.array2d(dtype=wp.quat),
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xquat_in: wp.array2d(dtype=wp.quat),
+ geom_skip_in: wp.array(dtype=bool),
+ # Data out:
+ geom_skip_out: wp.array(dtype=bool),
+ geom_xpos_out: wp.array2d(dtype=wp.vec3),
+ geom_xmat_out: wp.array2d(dtype=wp.mat33),
+):
+ worldid, geomid = wp.tid()
+ bodyid = geom_bodyid[geomid]
+ if not geom_skip_in[geomid]:
+ # Calculate only if necessary
+ xpos = xpos_in[worldid, bodyid]
+ xquat = xquat_in[worldid, bodyid]
+ geom_xpos_out[worldid, geomid] = xpos + math.rot_vec_quat(geom_pos[worldid, geomid], xquat)
+ geom_xmat_out[worldid, geomid] = math.quat_to_mat(math.mul_quat(xquat, geom_quat[worldid, geomid]))
+
+ if bodyid == 0:
+ # static geom pose are calculated only once
+ geom_skip_out[geomid] = True
+
+
+@wp.kernel
+def _site_local_to_global(
+ # Model:
+ site_bodyid: wp.array(dtype=int),
+ site_pos: wp.array2d(dtype=wp.vec3),
+ site_quat: wp.array2d(dtype=wp.quat),
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xquat_in: wp.array2d(dtype=wp.quat),
+ # Data out:
+ site_xpos_out: wp.array2d(dtype=wp.vec3),
+ site_xmat_out: wp.array2d(dtype=wp.mat33),
+):
+ worldid, siteid = wp.tid()
+ bodyid = site_bodyid[siteid]
+ xpos = xpos_in[worldid, bodyid]
+ xquat = xquat_in[worldid, bodyid]
+ site_xpos_out[worldid, siteid] = xpos + math.rot_vec_quat(site_pos[worldid, siteid], xquat)
+ site_xmat_out[worldid, siteid] = math.quat_to_mat(math.mul_quat(xquat, site_quat[worldid, siteid]))
+
+
+@wp.kernel
+def _flex_vertices(
+ # Model:
+ flex_vertbodyid: wp.array(dtype=int),
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ # Data out:
+ flexvert_xpos_out: wp.array2d(dtype=wp.vec3),
+):
+ worldid, vertid = wp.tid()
+ flexvert_xpos_out[worldid, vertid] = xpos_in[worldid, flex_vertbodyid[vertid]]
+
+
+@wp.kernel
+def _flex_edges(
+ # Model:
+ body_dofadr: wp.array(dtype=int),
+ flex_vertadr: wp.array(dtype=int),
+ flex_vertbodyid: wp.array(dtype=int),
+ flex_edge: wp.array(dtype=wp.vec2i),
+ # Data in:
+ qvel_in: wp.array2d(dtype=float),
+ flexvert_xpos_in: wp.array2d(dtype=wp.vec3),
+ # Data out:
+ flexedge_length_out: wp.array2d(dtype=float),
+ flexedge_velocity_out: wp.array2d(dtype=float),
+):
+ worldid, edgeid = wp.tid()
+ f = 0 # TODO(quaglino): get f from edgeid
+ vbase = flex_vertadr[f]
+ v = flex_edge[edgeid]
+ pos1 = flexvert_xpos_in[worldid, vbase + v[0]]
+ pos2 = flexvert_xpos_in[worldid, vbase + v[1]]
+ vec = pos2 - pos1
+ vecnorm = wp.length(vec)
+ flexedge_length_out[worldid, edgeid] = vecnorm
+ # TODO(quaglino): use Jacobian
+ i = body_dofadr[flex_vertbodyid[vbase + v[0]]]
+ j = body_dofadr[flex_vertbodyid[vbase + v[1]]]
+ vel1 = wp.vec3(qvel_in[worldid, i], qvel_in[worldid, i + 1], qvel_in[worldid, i + 2])
+ vel2 = wp.vec3(qvel_in[worldid, j], qvel_in[worldid, j + 1], qvel_in[worldid, j + 2])
+ flexedge_velocity_out[worldid, edgeid] = wp.dot(vel2 - vel1, vec) / vecnorm
+
+
+@wp.kernel
+def _mocap(
+ # Model:
+ body_ipos: wp.array2d(dtype=wp.vec3),
+ body_iquat: wp.array2d(dtype=wp.quat),
+ mocap_bodyid: wp.array(dtype=int),
+ # Data in:
+ mocap_pos_in: wp.array2d(dtype=wp.vec3),
+ mocap_quat_in: wp.array2d(dtype=wp.quat),
+ # Data out:
+ xpos_out: wp.array2d(dtype=wp.vec3),
+ xquat_out: wp.array2d(dtype=wp.quat),
+ xmat_out: wp.array2d(dtype=wp.mat33),
+ xipos_out: wp.array2d(dtype=wp.vec3),
+ ximat_out: wp.array2d(dtype=wp.mat33),
+):
+ worldid, mocapid = wp.tid()
+ bodyid = mocap_bodyid[mocapid]
+ mocap_quat = wp.normalize(mocap_quat_in[worldid, mocapid])
+ xpos = mocap_pos_in[worldid, mocapid]
+ xpos_out[worldid, bodyid] = xpos
+ xquat_out[worldid, bodyid] = mocap_quat
+ xmat_out[worldid, bodyid] = math.quat_to_mat(mocap_quat)
+ xipos_out[worldid, bodyid] = xpos + math.rot_vec_quat(body_ipos[worldid, bodyid], mocap_quat)
+ ximat_out[worldid, bodyid] = math.quat_to_mat(math.mul_quat(mocap_quat, body_iquat[worldid, bodyid]))
+
+
+@event_scope
+def kinematics(m: Model, d: Data):
+ """
+ Computes forward kinematics for all bodies, sites, geoms, and flexible elements.
+
+ This function updates the global positions and orientations of all bodies, as well as the
+ derived positions and orientations of geoms, sites, and flexible elements, based on the
+ current joint positions and any attached mocap bodies.
+ """
+ wp.launch(_kinematics_root, dim=(d.nworld), inputs=[], outputs=[d.xpos, d.xquat, d.xmat, d.xipos, d.ximat])
+
+ for i in range(1, len(m.body_tree)):
+ body_tree = m.body_tree[i]
+ wp.launch(
+ _kinematics_level,
+ dim=(d.nworld, body_tree.size),
+ inputs=[
+ m.qpos0,
+ m.body_parentid,
+ m.body_jntnum,
+ m.body_jntadr,
+ m.body_pos,
+ m.body_quat,
+ m.body_ipos,
+ m.body_iquat,
+ m.jnt_type,
+ m.jnt_qposadr,
+ m.jnt_pos,
+ m.jnt_axis,
+ d.qpos,
+ d.xpos,
+ d.xquat,
+ d.xmat,
+ body_tree,
+ ],
+ outputs=[d.xpos, d.xquat, d.xmat, d.xipos, d.ximat, d.xanchor, d.xaxis],
+ )
+
+ wp.launch(
+ _mocap,
+ dim=(d.nworld, m.nmocap),
+ inputs=[m.body_ipos, m.body_iquat, m.mocap_bodyid, d.mocap_pos, d.mocap_quat],
+ outputs=[d.xpos, d.xquat, d.xmat, d.xipos, d.ximat],
+ )
+
+ wp.launch(
+ _geom_local_to_global,
+ dim=(d.nworld, m.ngeom),
+ inputs=[m.geom_bodyid, m.geom_pos, m.geom_quat, d.xpos, d.xquat, d.geom_skip],
+ outputs=[d.geom_skip, d.geom_xpos, d.geom_xmat],
+ )
+
+ wp.launch(
+ _site_local_to_global,
+ dim=(d.nworld, m.nsite),
+ inputs=[m.site_bodyid, m.site_pos, m.site_quat, d.xpos, d.xquat],
+ outputs=[d.site_xpos, d.site_xmat],
+ )
+
+ wp.launch(_flex_vertices, dim=(d.nworld, m.nflexvert), inputs=[m.flex_vertbodyid, d.xpos], outputs=[d.flexvert_xpos])
+ wp.launch(
+ _flex_edges,
+ dim=(d.nworld, m.nflexedge),
+ inputs=[m.body_dofadr, m.flex_vertadr, m.flex_vertbodyid, m.flex_edge, d.qvel, d.flexvert_xpos],
+ outputs=[d.flexedge_length, d.flexedge_velocity],
+ )
+
+
+@wp.kernel
+def _subtree_com_init(
+ # Model:
+ body_mass: wp.array2d(dtype=float),
+ # Data in:
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ # Data out:
+ xipos_out: wp.array2d(dtype=wp.vec3),
+):
+ worldid, bodyid = wp.tid()
+ xipos_out[worldid, bodyid] = xipos_in[worldid, bodyid] * body_mass[worldid, bodyid]
+
+
+@wp.kernel
+def _subtree_com_acc(
+ # Model:
+ body_parentid: wp.array(dtype=int),
+ # Data in:
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ # In:
+ body_tree_: wp.array(dtype=int),
+ # Data out:
+ subtree_com_out: wp.array2d(dtype=wp.vec3),
+):
+ worldid, nodeid = wp.tid()
+ bodyid = body_tree_[nodeid]
+ pid = body_parentid[bodyid]
+ wp.atomic_add(subtree_com_out, worldid, pid, subtree_com_in[worldid, bodyid])
+
+
+@wp.kernel
+def _subtree_div(
+ # Model:
+ subtree_mass: wp.array2d(dtype=float),
+ # Data out:
+ subtree_com_out: wp.array2d(dtype=wp.vec3),
+):
+ worldid, bodyid = wp.tid()
+ subtree_com_out[worldid, bodyid] /= subtree_mass[worldid, bodyid]
+
+
+@wp.kernel
+def _cinert(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ body_mass: wp.array2d(dtype=float),
+ body_inertia: wp.array2d(dtype=wp.vec3),
+ # Data in:
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ ximat_in: wp.array2d(dtype=wp.mat33),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ # Data out:
+ cinert_out: wp.array2d(dtype=vec10),
+):
+ worldid, bodyid = wp.tid()
+ mat = ximat_in[worldid, bodyid]
+ inert = body_inertia[worldid, bodyid]
+ mass = body_mass[worldid, bodyid]
+ dif = xipos_in[worldid, bodyid] - subtree_com_in[worldid, body_rootid[bodyid]]
+ # express inertia in com-based frame (mju_inertCom)
+
+ res = vec10()
+ # res_rot = mat * diag(inert) * mat'
+ tmp = mat @ wp.diag(inert) @ wp.transpose(mat)
+ res[0] = tmp[0, 0]
+ res[1] = tmp[1, 1]
+ res[2] = tmp[2, 2]
+ res[3] = tmp[0, 1]
+ res[4] = tmp[0, 2]
+ res[5] = tmp[1, 2]
+ # res_rot -= mass * dif_cross * dif_cross
+ res[0] += mass * (dif[1] * dif[1] + dif[2] * dif[2])
+ res[1] += mass * (dif[0] * dif[0] + dif[2] * dif[2])
+ res[2] += mass * (dif[0] * dif[0] + dif[1] * dif[1])
+ res[3] -= mass * dif[0] * dif[1]
+ res[4] -= mass * dif[0] * dif[2]
+ res[5] -= mass * dif[1] * dif[2]
+ # res_tran = mass * dif
+ res[6] = mass * dif[0]
+ res[7] = mass * dif[1]
+ res[8] = mass * dif[2]
+ # res_mass = mass
+ res[9] = mass
+
+ cinert_out[worldid, bodyid] = res
+
+
+@wp.kernel
+def _cdof(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ jnt_type: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ jnt_bodyid: wp.array(dtype=int),
+ # Data in:
+ xmat_in: wp.array2d(dtype=wp.mat33),
+ xanchor_in: wp.array2d(dtype=wp.vec3),
+ xaxis_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ # Data out:
+ cdof_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ worldid, jntid = wp.tid()
+ bodyid = jnt_bodyid[jntid]
+ dofid = jnt_dofadr[jntid]
+ jnt_type_ = jnt_type[jntid]
+ xaxis = xaxis_in[worldid, jntid]
+ xmat = wp.transpose(xmat_in[worldid, bodyid])
+
+ # compute com-anchor vector
+ offset = subtree_com_in[worldid, body_rootid[bodyid]] - xanchor_in[worldid, jntid]
+
+ res = cdof_out[worldid]
+ if jnt_type_ == wp.static(JointType.FREE.value):
+ res[dofid + 0] = wp.spatial_vector(0.0, 0.0, 0.0, 1.0, 0.0, 0.0)
+ res[dofid + 1] = wp.spatial_vector(0.0, 0.0, 0.0, 0.0, 1.0, 0.0)
+ res[dofid + 2] = wp.spatial_vector(0.0, 0.0, 0.0, 0.0, 0.0, 1.0)
+ # I_3 rotation in child frame (assume no subsequent rotations)
+ res[dofid + 3] = wp.spatial_vector(xmat[0], wp.cross(xmat[0], offset))
+ res[dofid + 4] = wp.spatial_vector(xmat[1], wp.cross(xmat[1], offset))
+ res[dofid + 5] = wp.spatial_vector(xmat[2], wp.cross(xmat[2], offset))
+ elif jnt_type_ == wp.static(JointType.BALL.value): # ball
+ # I_3 rotation in child frame (assume no subsequent rotations)
+ res[dofid + 0] = wp.spatial_vector(xmat[0], wp.cross(xmat[0], offset))
+ res[dofid + 1] = wp.spatial_vector(xmat[1], wp.cross(xmat[1], offset))
+ res[dofid + 2] = wp.spatial_vector(xmat[2], wp.cross(xmat[2], offset))
+ elif jnt_type_ == wp.static(JointType.SLIDE.value):
+ res[dofid] = wp.spatial_vector(wp.vec3(0.0), xaxis)
+ elif jnt_type_ == wp.static(JointType.HINGE.value): # hinge
+ res[dofid] = wp.spatial_vector(xaxis, wp.cross(xaxis, offset))
+
+
+@event_scope
+def com_pos(m: Model, d: Data):
+ """
+ Computes subtree center of mass positions. Transforms inertia and motion to global frame
+ centered at subtree CoM.
+
+ Accumulates the mass-weighted positions up the kinematic tree, divides by total mass, and
+ computes composite inertias and motion degrees of freedom in the subtree CoM frame.
+ """
+ wp.launch(_subtree_com_init, dim=(d.nworld, m.nbody), inputs=[m.body_mass, d.xipos, d.subtree_com])
+
+ for i in reversed(range(len(m.body_tree))):
+ body_tree = m.body_tree[i]
+ wp.launch(
+ _subtree_com_acc,
+ dim=(d.nworld, body_tree.size),
+ inputs=[m.body_parentid, d.subtree_com, body_tree],
+ outputs=[d.subtree_com],
+ )
+
+ wp.launch(_subtree_div, dim=(d.nworld, m.nbody), inputs=[m.subtree_mass], outputs=[d.subtree_com])
+ wp.launch(
+ _cinert,
+ dim=(d.nworld, m.nbody),
+ inputs=[m.body_rootid, m.body_mass, m.body_inertia, d.xipos, d.ximat, d.subtree_com],
+ outputs=[d.cinert],
+ )
+ wp.launch(
+ _cdof,
+ dim=(d.nworld, m.njnt),
+ inputs=[m.body_rootid, m.jnt_type, m.jnt_dofadr, m.jnt_bodyid, d.xmat, d.xanchor, d.xaxis, d.subtree_com],
+ outputs=[d.cdof],
+ )
+
+
+@wp.kernel
+def _cam_local_to_global(
+ # Model:
+ cam_bodyid: wp.array(dtype=int),
+ cam_pos: wp.array2d(dtype=wp.vec3),
+ cam_quat: wp.array2d(dtype=wp.quat),
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xquat_in: wp.array2d(dtype=wp.quat),
+ # Data out:
+ cam_xpos_out: wp.array2d(dtype=wp.vec3),
+ cam_xmat_out: wp.array2d(dtype=wp.mat33),
+):
+ """Fixed cameras."""
+ worldid, camid = wp.tid()
+ bodyid = cam_bodyid[camid]
+ xpos = xpos_in[worldid, bodyid]
+ xquat = xquat_in[worldid, bodyid]
+ cam_xpos_out[worldid, camid] = xpos + math.rot_vec_quat(cam_pos[worldid, camid], xquat)
+ cam_xmat_out[worldid, camid] = math.quat_to_mat(math.mul_quat(xquat, cam_quat[worldid, camid]))
+
+
+@wp.kernel
+def _cam_fn(
+ # Model:
+ cam_mode: wp.array(dtype=int),
+ cam_bodyid: wp.array(dtype=int),
+ cam_targetbodyid: wp.array(dtype=int),
+ cam_poscom0: wp.array2d(dtype=wp.vec3),
+ cam_pos0: wp.array2d(dtype=wp.vec3),
+ cam_mat0: wp.array2d(dtype=wp.mat33),
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ # Data out:
+ cam_xpos_out: wp.array2d(dtype=wp.vec3),
+ cam_xmat_out: wp.array2d(dtype=wp.mat33),
+):
+ worldid, camid = wp.tid()
+ is_target_cam = (cam_mode[camid] == wp.static(CamLightType.TARGETBODY.value)) or (
+ cam_mode[camid] == wp.static(CamLightType.TARGETBODYCOM.value)
+ )
+ invalid_target = is_target_cam and (cam_targetbodyid[camid] < 0)
+ if invalid_target:
+ return
+ elif cam_mode[camid] == wp.static(CamLightType.TRACK.value):
+ cam_xmat_out[worldid, camid] = cam_mat0[worldid, camid]
+ body_xpos = xpos_in[worldid, cam_bodyid[camid]]
+ cam_xpos_out[worldid, camid] = body_xpos + cam_pos0[worldid, camid]
+ elif cam_mode[camid] == wp.static(CamLightType.TRACKCOM.value):
+ cam_xmat_out[worldid, camid] = cam_mat0[worldid, camid]
+ cam_xpos_out[worldid, camid] = subtree_com_in[worldid, cam_bodyid[camid]] + cam_poscom0[worldid, camid]
+ elif cam_mode[camid] == wp.static(CamLightType.TARGETBODY.value) or cam_mode[camid] == wp.static(
+ CamLightType.TARGETBODYCOM.value
+ ):
+ pos = xpos_in[worldid, cam_targetbodyid[camid]]
+ if cam_mode[camid] == wp.static(CamLightType.TARGETBODYCOM.value):
+ pos = subtree_com_in[worldid, cam_targetbodyid[camid]]
+ # zaxis = -desired camera direction, in global frame
+ mat_3 = wp.normalize(cam_xpos_out[worldid, camid] - pos)
+ # xaxis: orthogonal to zaxis and to (0,0,1)
+ mat_1 = wp.normalize(wp.cross(wp.vec3(0.0, 0.0, 1.0), mat_3))
+ mat_2 = wp.normalize(wp.cross(mat_3, mat_1))
+ # fmt: off
+ cam_xmat_out[worldid, camid] = wp.mat33(
+ mat_1[0], mat_2[0], mat_3[0],
+ mat_1[1], mat_2[1], mat_3[1],
+ mat_1[2], mat_2[2], mat_3[2]
+ )
+ # fmt: on
+
+
+@wp.kernel
+def _light_local_to_global(
+ # Model:
+ light_bodyid: wp.array(dtype=int),
+ light_pos: wp.array2d(dtype=wp.vec3),
+ light_dir: wp.array2d(dtype=wp.vec3),
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xquat_in: wp.array2d(dtype=wp.quat),
+ # Data out:
+ light_xpos_out: wp.array2d(dtype=wp.vec3),
+ light_xdir_out: wp.array2d(dtype=wp.vec3),
+):
+ """Fixed lights."""
+ worldid, lightid = wp.tid()
+ bodyid = light_bodyid[lightid]
+ xpos = xpos_in[worldid, bodyid]
+ xquat = xquat_in[worldid, bodyid]
+ light_xpos_out[worldid, lightid] = xpos + math.rot_vec_quat(light_pos[worldid, lightid], xquat)
+ light_xdir_out[worldid, lightid] = math.rot_vec_quat(light_dir[worldid, lightid], xquat)
+
+
+@wp.kernel
+def _light_fn(
+ # Model:
+ light_mode: wp.array(dtype=int),
+ light_bodyid: wp.array(dtype=int),
+ light_targetbodyid: wp.array(dtype=int),
+ light_poscom0: wp.array2d(dtype=wp.vec3),
+ light_pos0: wp.array2d(dtype=wp.vec3),
+ light_dir0: wp.array2d(dtype=wp.vec3),
+ # Data in:
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ light_xpos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ # Data out:
+ light_xpos_out: wp.array2d(dtype=wp.vec3),
+ light_xdir_out: wp.array2d(dtype=wp.vec3),
+):
+ worldid, lightid = wp.tid()
+ is_target_light = (light_mode[lightid] == wp.static(CamLightType.TARGETBODY.value)) or (
+ light_mode[lightid] == wp.static(CamLightType.TARGETBODYCOM.value)
+ )
+ invalid_target = is_target_light and (light_targetbodyid[lightid] < 0)
+ if invalid_target:
+ return
+ elif light_mode[lightid] == wp.static(CamLightType.TRACK.value):
+ light_xdir_out[worldid, lightid] = light_dir0[worldid, lightid]
+ body_xpos = xpos_in[worldid, light_bodyid[lightid]]
+ light_xpos_out[worldid, lightid] = body_xpos + light_pos0[worldid, lightid]
+ elif light_mode[lightid] == wp.static(CamLightType.TRACKCOM.value):
+ light_xdir_out[worldid, lightid] = light_dir0[worldid, lightid]
+ light_xpos_out[worldid, lightid] = subtree_com_in[worldid, light_bodyid[lightid]] + light_poscom0[worldid, lightid]
+ elif light_mode[lightid] == wp.static(CamLightType.TARGETBODY.value) or light_mode[lightid] == wp.static(
+ CamLightType.TARGETBODYCOM.value
+ ):
+ pos = xpos_in[worldid, light_targetbodyid[lightid]]
+ if light_mode[lightid] == wp.static(CamLightType.TARGETBODYCOM.value):
+ pos = subtree_com_in[worldid, light_targetbodyid[lightid]]
+ light_xdir_out[worldid, lightid] = pos - light_xpos_in[worldid, lightid]
+ light_xdir_out[worldid, lightid] = wp.normalize(light_xdir_out[worldid, lightid])
+
+
+@event_scope
+def camlight(m: Model, d: Data):
+ """
+ Computes camera and light positions and orientations.
+
+ Updates the global positions and orientations for all cameras and lights in the model,
+ including special handling for tracking and target modes.
+ """
+ wp.launch(
+ _cam_local_to_global,
+ dim=(d.nworld, m.ncam),
+ inputs=[m.cam_bodyid, m.cam_pos, m.cam_quat, d.xpos, d.xquat],
+ outputs=[d.cam_xpos, d.cam_xmat],
+ )
+ wp.launch(
+ _cam_fn,
+ dim=(d.nworld, m.ncam),
+ inputs=[m.cam_mode, m.cam_bodyid, m.cam_targetbodyid, m.cam_poscom0, m.cam_pos0, m.cam_mat0, d.xpos, d.subtree_com],
+ outputs=[d.cam_xpos, d.cam_xmat],
+ )
+ wp.launch(
+ _light_local_to_global,
+ dim=(d.nworld, m.nlight),
+ inputs=[m.light_bodyid, m.light_pos, m.light_dir, d.xpos, d.xquat],
+ outputs=[d.light_xpos, d.light_xdir],
+ )
+ wp.launch(
+ _light_fn,
+ dim=(d.nworld, m.nlight),
+ inputs=[
+ m.light_mode,
+ m.light_bodyid,
+ m.light_targetbodyid,
+ m.light_poscom0,
+ m.light_pos0,
+ m.light_dir0,
+ d.xpos,
+ d.light_xpos,
+ d.subtree_com,
+ ],
+ outputs=[d.light_xpos, d.light_xdir],
+ )
+
+
+@wp.kernel
+def _crb_accumulate(
+ # Model:
+ body_parentid: wp.array(dtype=int),
+ # Data in:
+ crb_in: wp.array2d(dtype=vec10),
+ # In:
+ body_tree_: wp.array(dtype=int),
+ # Data out:
+ crb_out: wp.array2d(dtype=vec10),
+):
+ worldid, nodeid = wp.tid()
+ bodyid = body_tree_[nodeid]
+ pid = body_parentid[bodyid]
+ if pid == 0:
+ return
+ wp.atomic_add(crb_out, worldid, pid, crb_in[worldid, bodyid])
+
+
+@wp.kernel
+def _qM_sparse(
+ # Model:
+ dof_bodyid: wp.array(dtype=int),
+ dof_parentid: wp.array(dtype=int),
+ dof_Madr: wp.array(dtype=int),
+ dof_armature: wp.array2d(dtype=float),
+ # Data in:
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ crb_in: wp.array2d(dtype=vec10),
+ # Data out:
+ qM_out: wp.array3d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ madr_ij = dof_Madr[dofid]
+ bodyid = dof_bodyid[dofid]
+
+ # init M(i,i) with armature inertia
+ qM_out[worldid, 0, madr_ij] = dof_armature[worldid, dofid]
+
+ # precompute buf = crb_body_i * cdof_i
+ buf = math.inert_vec(crb_in[worldid, bodyid], cdof_in[worldid, dofid])
+
+ # sparse backward pass over ancestors
+ while dofid >= 0:
+ qM_out[worldid, 0, madr_ij] += wp.dot(cdof_in[worldid, dofid], buf)
+ madr_ij += 1
+ dofid = dof_parentid[dofid]
+
+
+@wp.kernel
+def _qM_dense(
+ # Model:
+ dof_bodyid: wp.array(dtype=int),
+ dof_parentid: wp.array(dtype=int),
+ dof_armature: wp.array2d(dtype=float),
+ # Data in:
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ crb_in: wp.array2d(dtype=vec10),
+ # Data out:
+ qM_out: wp.array3d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ bodyid = dof_bodyid[dofid]
+ # init M(i,i) with armature inertia
+ M = dof_armature[worldid, dofid]
+
+ # precompute buf = crb_body_i * cdof_i
+ buf = math.inert_vec(crb_in[worldid, bodyid], cdof_in[worldid, dofid])
+ M += wp.dot(cdof_in[worldid, dofid], buf)
+
+ qM_out[worldid, dofid, dofid] = M
+
+ # sparse backward pass over ancestors
+ dofidi = dofid
+ dofid = dof_parentid[dofid]
+ while dofid >= 0:
+ qMij = wp.dot(cdof_in[worldid, dofid], buf)
+ qM_out[worldid, dofidi, dofid] += qMij
+ qM_out[worldid, dofid, dofidi] += qMij
+ dofid = dof_parentid[dofid]
+
+
+@event_scope
+def crb(m: Model, d: Data):
+ """
+ Computes composite rigid body inertias for each body and the joint-space inertia matrix.
+
+ Accumulates composite rigid body inertias up the kinematic tree and computes the
+ joint-space inertia matrix in either sparse or dense format, depending on model options.
+ """
+ wp.copy(d.crb, d.cinert)
+
+ for i in reversed(range(len(m.body_tree))):
+ body_tree = m.body_tree[i]
+ wp.launch(_crb_accumulate, dim=(d.nworld, body_tree.size), inputs=[m.body_parentid, d.crb, body_tree], outputs=[d.crb])
+
+ d.qM.zero_()
+ if m.opt.is_sparse:
+ wp.launch(
+ _qM_sparse,
+ dim=(d.nworld, m.nv),
+ inputs=[m.dof_bodyid, m.dof_parentid, m.dof_Madr, m.dof_armature, d.cdof, d.crb],
+ outputs=[d.qM],
+ )
+ else:
+ wp.launch(
+ _qM_dense, dim=(d.nworld, m.nv), inputs=[m.dof_bodyid, m.dof_parentid, m.dof_armature, d.cdof, d.crb], outputs=[d.qM]
+ )
+
+
+@wp.kernel
+def _tendon_armature(
+ # Model:
+ opt_is_sparse: bool,
+ dof_parentid: wp.array(dtype=int),
+ dof_Madr: wp.array(dtype=int),
+ tendon_armature: wp.array2d(dtype=float),
+ # Data in:
+ ten_J_in: wp.array3d(dtype=float),
+ # Data out:
+ qM_out: wp.array3d(dtype=float),
+):
+ worldid, tenid, dofid = wp.tid()
+
+ if opt_is_sparse:
+ madr_ij = dof_Madr[dofid]
+
+ armature = tendon_armature[worldid, tenid]
+
+ if armature == 0.0:
+ return
+
+ ten_Ji = ten_J_in[worldid, tenid, dofid]
+
+ if ten_Ji == 0.0:
+ return
+
+ # sparse backward pass over ancestors
+ dofidi = dofid
+ while dofid >= 0:
+ if dofid != dofidi:
+ ten_Jj = ten_J_in[worldid, tenid, dofid]
+ else:
+ ten_Jj = ten_Ji
+
+ qMij = armature * ten_Jj * ten_Ji
+
+ if opt_is_sparse:
+ wp.atomic_add(qM_out[worldid, 0], madr_ij, qMij)
+ madr_ij += 1
+ else:
+ wp.atomic_add(qM_out[worldid, dofidi], dofid, qMij)
+ if dofidi != dofid:
+ wp.atomic_add(qM_out[worldid, dofid], dofidi, qMij)
+
+ dofid = dof_parentid[dofid]
+
+
+@event_scope
+def tendon_armature(m: Model, d: Data):
+ """Add tendon armature to qM."""
+ wp.launch(
+ _tendon_armature,
+ dim=(d.nworld, m.ntendon, m.nv),
+ inputs=[m.opt.is_sparse, m.dof_parentid, m.dof_Madr, m.tendon_armature, d.ten_J],
+ outputs=[d.qM],
+ )
+
+
+@wp.kernel
+def _copy_CSR(
+ # Model:
+ mapM2M: wp.array(dtype=int),
+ # In:
+ M_in: wp.array3d(dtype=float),
+ # Out:
+ L_out: wp.array3d(dtype=float),
+):
+ worldid, ind = wp.tid()
+ L_out[worldid, 0, ind] = M_in[worldid, 0, mapM2M[ind]]
+
+
+@wp.kernel
+def _qLD_acc(
+ # Model:
+ M_rownnz: wp.array(dtype=int),
+ M_rowadr: wp.array(dtype=int),
+ # In:
+ qLD_updates_: wp.array(dtype=wp.vec3i),
+ L_in: wp.array3d(dtype=float),
+ # Out:
+ L_out: wp.array3d(dtype=float),
+):
+ worldid, nodeid = wp.tid()
+ update = qLD_updates_[nodeid]
+ i, k, Madr_ki = update[0], update[1], update[2]
+ Madr_i = M_rowadr[i] # Address of row being updated
+ diag_k = M_rowadr[k] + M_rownnz[k] - 1 # Address of diagonal element of k
+ # tmp = M(k,i) / M(k,k)
+ tmp = L_out[worldid, 0, Madr_ki] / L_out[worldid, 0, diag_k]
+ for j in range(M_rownnz[i]):
+ # M(i,j) -= M(k,j) * tmp
+ wp.atomic_sub(L_out[worldid, 0], Madr_i + j, L_in[worldid, 0, M_rowadr[k] + j] * tmp)
+ # M(k,i) = tmp
+ L_out[worldid, 0, Madr_ki] = tmp
+
+
+@wp.kernel
+def _qLDiag_div(
+ # Model:
+ M_rownnz: wp.array(dtype=int),
+ M_rowadr: wp.array(dtype=int),
+ # In:
+ L_in: wp.array3d(dtype=float),
+ # Out:
+ D_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ diag_i = M_rowadr[dofid] + M_rownnz[dofid] - 1 # Address of diagonal element of i
+ D_out[worldid, dofid] = 1.0 / L_in[worldid, 0, diag_i]
+
+
+def _factor_i_sparse(m: Model, d: Data, M: wp.array3d(dtype=float), L: wp.array3d(dtype=float), D: wp.array2d(dtype=float)):
+ """Sparse L'*D*L factorization of inertia-like matrix M, assumed spd."""
+ wp.launch(_copy_CSR, dim=(d.nworld, m.nC), inputs=[m.mapM2M, M], outputs=[L])
+
+ for i in reversed(range(len(m.qLD_updates))):
+ qLD_updates = m.qLD_updates[i]
+ wp.launch(_qLD_acc, dim=(d.nworld, qLD_updates.size), inputs=[m.M_rownnz, m.M_rowadr, qLD_updates, L], outputs=[L])
+
+ wp.launch(_qLDiag_div, dim=(d.nworld, m.nv), inputs=[m.M_rownnz, m.M_rowadr, L], outputs=[D])
+
+
+@cache_kernel
+def _tile_cholesky_factorize(tile: TileSet):
+ """Returns a kernel for dense Cholesky factorization of a tile."""
+
+ @nested_kernel
+ def cholesky_factorize(
+ # Data In:
+ qM_in: wp.array3d(dtype=float),
+ # In:
+ adr: wp.array(dtype=int),
+ # Out:
+ L_out: wp.array3d(dtype=float),
+ ):
+ worldid, nodeid = wp.tid()
+ TILE_SIZE = wp.static(tile.size)
+
+ dofid = adr[nodeid]
+ M_tile = wp.tile_load(qM_in[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid))
+ L_tile = wp.tile_cholesky(M_tile)
+ wp.tile_store(L_out[worldid], L_tile, offset=(dofid, dofid))
+
+ return cholesky_factorize
+
+
+def _factor_i_dense(m: Model, d: Data, M: wp.array, L: wp.array):
+ """Dense Cholesky factorization of inertia-like matrix M, assumed spd."""
+ for tile in m.qM_tiles:
+ wp.launch_tiled(
+ _tile_cholesky_factorize(tile),
+ dim=(d.nworld, tile.adr.size),
+ inputs=[M, tile.adr],
+ outputs=[L],
+ block_dim=m.block_dim.cholesky_factorize,
+ )
+
+
+@event_scope
+def factor_m(m: Model, d: Data):
+ """Factorization of inertia-like matrix M, assumed spd."""
+ if m.opt.is_sparse:
+ _factor_i_sparse(m, d, d.qM, d.qLD, d.qLDiagInv)
+ else:
+ _factor_i_dense(m, d, d.qM, d.qLD)
+
+
+@wp.kernel
+def _cacc_world(
+ # In:
+ gravity: wp.array(dtype=wp.vec3),
+ # Data out:
+ cacc_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ worldid = wp.tid()
+ cacc_out[worldid, 0] = wp.spatial_vector(wp.vec3(0.0), -gravity[worldid])
+
+
+def _rne_cacc_world(m: Model, d: Data):
+ if m.opt.disableflags & DisableBit.GRAVITY:
+ d.cacc.zero_()
+ else:
+ wp.launch(_cacc_world, dim=[d.nworld], inputs=[m.opt.gravity], outputs=[d.cacc])
+
+
+@wp.kernel
+def _cacc(
+ # Model:
+ body_parentid: wp.array(dtype=int),
+ body_dofnum: wp.array(dtype=int),
+ body_dofadr: wp.array(dtype=int),
+ # Data in:
+ qvel_in: wp.array2d(dtype=float),
+ qacc_in: wp.array2d(dtype=float),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ cdof_dot_in: wp.array2d(dtype=wp.spatial_vector),
+ cacc_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ body_tree_: wp.array(dtype=int),
+ flg_acc: bool,
+ # Data out:
+ cacc_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ worldid, nodeid = wp.tid()
+ bodyid = body_tree_[nodeid]
+ dofnum = body_dofnum[bodyid]
+ pid = body_parentid[bodyid]
+ dofadr = body_dofadr[bodyid]
+ local_cacc = cacc_in[worldid, pid]
+ for i in range(dofnum):
+ local_cacc += cdof_dot_in[worldid, dofadr + i] * qvel_in[worldid, dofadr + i]
+ if flg_acc:
+ local_cacc += cdof_in[worldid, dofadr + i] * qacc_in[worldid, dofadr + i]
+ cacc_out[worldid, bodyid] = local_cacc
+
+
+def _rne_cacc_forward(m: Model, d: Data, flg_acc: bool = False):
+ for body_tree in m.body_tree:
+ wp.launch(
+ _cacc,
+ dim=(d.nworld, body_tree.size),
+ inputs=[m.body_parentid, m.body_dofnum, m.body_dofadr, d.qvel, d.qacc, d.cdof, d.cdof_dot, d.cacc, body_tree, flg_acc],
+ outputs=[d.cacc],
+ )
+
+
+@wp.kernel
+def _cfrc(
+ # Data in:
+ cinert_in: wp.array2d(dtype=vec10),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ cacc_in: wp.array2d(dtype=wp.spatial_vector),
+ cfrc_ext_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ flg_cfrc_ext: bool,
+ # Data out:
+ cfrc_int_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ worldid, bodyid = wp.tid()
+ bodyid += 1 # skip world body
+ cacc = cacc_in[worldid, bodyid]
+ cinert = cinert_in[worldid, bodyid]
+ cvel = cvel_in[worldid, bodyid]
+ frc = math.inert_vec(cinert, cacc)
+ frc += math.motion_cross_force(cvel, math.inert_vec(cinert, cvel))
+ if flg_cfrc_ext:
+ frc -= cfrc_ext_in[worldid, bodyid]
+
+ cfrc_int_out[worldid, bodyid] = frc
+
+
+def _rne_cfrc(m: Model, d: Data, flg_cfrc_ext: bool = False):
+ wp.launch(
+ _cfrc, dim=[d.nworld, m.nbody - 1], inputs=[d.cinert, d.cvel, d.cacc, d.cfrc_ext, flg_cfrc_ext], outputs=[d.cfrc_int]
+ )
+
+
+@wp.kernel
+def _cfrc_backward(
+ # Model:
+ body_parentid: wp.array(dtype=int),
+ # Data in:
+ cfrc_int_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ body_tree_: wp.array(dtype=int),
+ # Data out:
+ cfrc_int_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ worldid, nodeid = wp.tid()
+ bodyid = body_tree_[nodeid]
+ pid = body_parentid[bodyid]
+ if bodyid != 0:
+ wp.atomic_add(cfrc_int_out[worldid], pid, cfrc_int_in[worldid, bodyid])
+
+
+def _rne_cfrc_backward(m: Model, d: Data):
+ for body_tree in reversed(m.body_tree):
+ wp.launch(
+ _cfrc_backward, dim=[d.nworld, body_tree.size], inputs=[m.body_parentid, d.cfrc_int, body_tree], outputs=[d.cfrc_int]
+ )
+
+
+@wp.kernel
+def _qfrc_bias(
+ # Model:
+ dof_bodyid: wp.array(dtype=int),
+ # Data in:
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ cfrc_int_in: wp.array2d(dtype=wp.spatial_vector),
+ # Data out:
+ qfrc_bias_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ bodyid = dof_bodyid[dofid]
+ qfrc_bias_out[worldid, dofid] = wp.dot(cdof_in[worldid, dofid], cfrc_int_in[worldid, bodyid])
+
+
+@event_scope
+def rne(m: Model, d: Data, flg_acc: bool = False):
+ """
+ Computes inverse dynamics using the recursive Newton-Euler algorithm.
+
+ Computes the bias forces (qfrc_bias) and internal forces (cfrc_int) for the current state,
+ including the effects of gravity and optionally joint accelerations.
+
+ Args:
+ m (Model): The model containing kinematic and dynamic information.
+ d (Data): The data object containing the current state and output arrays.
+ flg_acc (bool, optional): If True, includes joint accelerations in the computation.
+ Defaults to False.
+ """
+ _rne_cacc_world(m, d)
+ _rne_cacc_forward(m, d, flg_acc=flg_acc)
+ _rne_cfrc(m, d)
+ _rne_cfrc_backward(m, d)
+ wp.launch(_qfrc_bias, dim=[d.nworld, m.nv], inputs=[m.dof_bodyid, d.cdof, d.cfrc_int], outputs=[d.qfrc_bias])
+
+
+@wp.kernel
+def _cfrc_ext(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ # Data in:
+ xfrc_applied_in: wp.array2d(dtype=wp.spatial_vector),
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ # Data out:
+ cfrc_ext_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ worldid, bodyid = wp.tid()
+ if bodyid == 0:
+ cfrc_ext_out[worldid, 0] = wp.spatial_vector(0.0, 0.0, 0.0, 0.0, 0.0, 0.0)
+ else:
+ xfrc_applied = xfrc_applied_in[worldid, bodyid]
+ subtree_com = subtree_com_in[worldid, body_rootid[bodyid]]
+ xipos = xipos_in[worldid, bodyid]
+ cfrc_ext_out[worldid, bodyid] = support.transform_force(xfrc_applied, subtree_com - xipos)
+
+
+@wp.kernel
+def _cfrc_ext_equality(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ site_pos: wp.array2d(dtype=wp.vec3),
+ eq_obj1id: wp.array(dtype=int),
+ eq_obj2id: wp.array(dtype=int),
+ eq_objtype: wp.array(dtype=int),
+ eq_data: wp.array2d(dtype=vec11),
+ # Data in:
+ ne_connect_in: wp.array(dtype=int),
+ ne_weld_in: wp.array(dtype=int),
+ xpos_in: wp.array2d(dtype=wp.vec3),
+ xmat_in: wp.array2d(dtype=wp.mat33),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ efc_id_in: wp.array2d(dtype=int),
+ efc_force_in: wp.array2d(dtype=float),
+ # Data out:
+ cfrc_ext_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ worldid, eqid = wp.tid()
+
+ ne_connect = ne_connect_in[worldid]
+ ne_weld = ne_weld_in[worldid]
+ num_connect = ne_connect // 3
+
+ if eqid >= num_connect + ne_weld // 6:
+ return
+
+ is_connect = eqid < num_connect
+ if is_connect:
+ efcid = 3 * eqid
+ cfrc_torque = wp.vec3(0.0, 0.0, 0.0) # no torque from connect
+ else:
+ efcid = 6 * eqid - ne_connect
+ cfrc_torque = wp.vec3(efc_force_in[worldid, efcid + 3], efc_force_in[worldid, efcid + 4], efc_force_in[worldid, efcid + 5])
+
+ cfrc_force = wp.vec3(
+ efc_force_in[worldid, efcid + 0],
+ efc_force_in[worldid, efcid + 1],
+ efc_force_in[worldid, efcid + 2],
+ )
+
+ id = efc_id_in[worldid, efcid]
+ eq_data_ = eq_data[worldid, id]
+ body_semantic = eq_objtype[id] == wp.static(ObjType.BODY.value)
+
+ obj1 = eq_obj1id[id]
+ obj2 = eq_obj2id[id]
+
+ if body_semantic:
+ bodyid1 = obj1
+ bodyid2 = obj2
+ else:
+ bodyid1 = site_bodyid[obj1]
+ bodyid2 = site_bodyid[obj2]
+
+ # body 1
+ if bodyid1:
+ if body_semantic:
+ if is_connect:
+ offset = wp.vec3(eq_data_[0], eq_data_[1], eq_data_[2])
+ else:
+ offset = wp.vec3(eq_data_[3], eq_data_[4], eq_data_[5])
+ else:
+ offset = site_pos[worldid, obj1]
+
+ # transform point on body1: local -> global
+ pos = xmat_in[worldid, bodyid1] @ offset + xpos_in[worldid, bodyid1]
+
+ # subtree CoM-based torque_force vector
+ newpos = subtree_com_in[worldid, body_rootid[bodyid1]]
+
+ dif = newpos - pos
+ cfrc_com = wp.spatial_vector(cfrc_torque - wp.cross(dif, cfrc_force), cfrc_force)
+
+ # apply (opposite for body 1)
+ wp.atomic_add(cfrc_ext_out[worldid], bodyid1, cfrc_com)
+
+ # body 2
+ if bodyid2:
+ if body_semantic:
+ if is_connect:
+ offset = wp.vec3(eq_data_[3], eq_data_[4], eq_data_[5])
+ else:
+ offset = wp.vec3(eq_data_[0], eq_data_[1], eq_data_[2])
+ else:
+ offset = site_pos[worldid, obj2]
+
+ # transform point on body2: local -> global
+ pos = xmat_in[worldid, bodyid2] @ offset + xpos_in[worldid, bodyid2]
+
+ # subtree CoM-based torque_force vector
+ newpos = subtree_com_in[worldid, body_rootid[bodyid2]]
+
+ dif = newpos - pos
+ cfrc_com = wp.spatial_vector(cfrc_torque - wp.cross(dif, cfrc_force), cfrc_force)
+
+ # apply
+ wp.atomic_sub(cfrc_ext_out[worldid], bodyid2, cfrc_com)
+
+
+@wp.func
+def transform_force(force: wp.vec3, torque: wp.vec3, offset: wp.vec3) -> wp.spatial_vector:
+ torque -= wp.cross(offset, force)
+ return wp.spatial_vector(torque, force)
+
+
+@wp.kernel
+def _cfrc_ext_contact(
+ # Model:
+ opt_cone: int,
+ body_rootid: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ contact_pos_in: wp.array(dtype=wp.vec3),
+ contact_frame_in: wp.array(dtype=wp.mat33),
+ contact_friction_in: wp.array(dtype=vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_geom_in: wp.array(dtype=wp.vec2i),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ contact_worldid_in: wp.array(dtype=int),
+ efc_force_in: wp.array2d(dtype=float),
+ # Data out:
+ cfrc_ext_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ contactid = wp.tid()
+
+ if contactid >= ncon_in[0]:
+ return
+
+ geom = contact_geom_in[contactid]
+ id1 = geom_bodyid[geom[0]]
+ id2 = geom_bodyid[geom[1]]
+
+ if id1 == 0 and id2 == 0:
+ return
+
+ worldid = contact_worldid_in[contactid]
+
+ # contact force in world frame
+ force = support.contact_force_fn(
+ opt_cone,
+ ncon_in,
+ contact_frame_in,
+ contact_friction_in,
+ contact_dim_in,
+ contact_efc_address_in,
+ efc_force_in,
+ worldid,
+ contactid,
+ to_world_frame=True,
+ )
+
+ pos = contact_pos_in[contactid]
+
+ # contact force on bodies
+ if id1:
+ com1 = subtree_com_in[worldid, body_rootid[id1]]
+ wp.atomic_sub(cfrc_ext_out[worldid], id1, support.transform_force(force, com1 - pos))
+
+ if id2:
+ com2 = subtree_com_in[worldid, body_rootid[id2]]
+ wp.atomic_add(cfrc_ext_out[worldid], id2, support.transform_force(force, com2 - pos))
+
+
+@event_scope
+def rne_postconstraint(m: Model, d: Data):
+ """
+ Computes the recursive Newton-Euler algorithm after constraints are applied.
+
+ Computes cacc, cfrc_ext, and cfrc_int, including the effects of applied forces, equality
+ constraints, and contacts.
+ """
+ # cfrc_ext = perturb
+ wp.launch(
+ _cfrc_ext,
+ dim=(d.nworld, m.nbody),
+ inputs=[m.body_rootid, d.xfrc_applied, d.xipos, d.subtree_com],
+ outputs=[d.cfrc_ext],
+ )
+
+ wp.launch(
+ _cfrc_ext_equality,
+ dim=(d.nworld, m.neq),
+ inputs=[
+ m.body_rootid,
+ m.site_bodyid,
+ m.site_pos,
+ m.eq_obj1id,
+ m.eq_obj2id,
+ m.eq_objtype,
+ m.eq_data,
+ d.ne_connect,
+ d.ne_weld,
+ d.xpos,
+ d.xmat,
+ d.subtree_com,
+ d.efc.id,
+ d.efc.force,
+ ],
+ outputs=[d.cfrc_ext],
+ )
+
+ # cfrc_ext += contacts
+ wp.launch(
+ _cfrc_ext_contact,
+ dim=(d.nconmax,),
+ inputs=[
+ m.opt.cone,
+ m.body_rootid,
+ m.geom_bodyid,
+ d.ncon,
+ d.subtree_com,
+ d.contact.pos,
+ d.contact.frame,
+ d.contact.friction,
+ d.contact.dim,
+ d.contact.geom,
+ d.contact.efc_address,
+ d.contact.worldid,
+ d.efc.force,
+ ],
+ outputs=[d.cfrc_ext],
+ )
+
+ # forward pass over bodies: compute cacc, cfrc_int
+ _rne_cacc_world(m, d)
+ _rne_cacc_forward(m, d, flg_acc=True)
+
+ # cfrc_body = cinert * cacc + cvel x (cinert * cvel)
+ _rne_cfrc(m, d, flg_cfrc_ext=True)
+
+ # backward pass over bodies: accumulate cfrc_int from children
+ _rne_cfrc_backward(m, d)
+
+
+@wp.kernel
+def _tendon_dot(
+ # Model:
+ nv: int,
+ body_parentid: wp.array(dtype=int),
+ body_rootid: wp.array(dtype=int),
+ jnt_type: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ dof_bodyid: wp.array(dtype=int),
+ dof_jntid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ tendon_adr: wp.array(dtype=int),
+ tendon_num: wp.array(dtype=int),
+ tendon_armature: wp.array2d(dtype=float),
+ wrap_objid: wp.array(dtype=int),
+ wrap_prm: wp.array(dtype=float),
+ wrap_type: wp.array(dtype=int),
+ # Data in:
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ cdof_dot_in: wp.array2d(dtype=wp.spatial_vector),
+ # Data out:
+ ten_Jdot_out: wp.array3d(dtype=float),
+):
+ worldid, tenid = wp.tid()
+
+ armature = tendon_armature[worldid, tenid]
+ if armature == 0.0:
+ return
+
+ # fixed tendon has zero Jdot
+ adr = tendon_adr[tenid]
+ if wrap_type[adr] == int(WrapType.JOINT.value):
+ return
+
+ # process spatial tendon
+ divisor = float(1.0)
+ num = tendon_num[tenid]
+ j = int(0)
+ while j < num - 1:
+ # get 1st and 2nd object
+ type0 = wrap_type[adr + j + 0]
+ type1 = wrap_type[adr + j + 1]
+ id0 = wrap_objid[adr + j + 0]
+ id1 = wrap_objid[adr + j + 1]
+
+ # pulley
+ pulley = int(WrapType.PULLEY.value)
+ if (type0 == pulley) or (type1 == pulley):
+ # get divisor
+ if type0 == pulley:
+ divisor = wrap_prm[adr + j]
+
+ j += 1
+ continue
+
+ # init sequence; assume it start with site
+ wpnt0 = site_xpos_in[worldid, id0]
+
+ bodyid0 = site_bodyid[id0]
+ pos0 = site_xpos_in[worldid, id0]
+ cvel0 = cvel_in[worldid, bodyid0]
+ subtree_com0 = subtree_com_in[worldid, body_rootid[bodyid0]]
+ dif0 = pos0 - subtree_com0
+ wvel0 = wp.spatial_bottom(cvel0) - wp.cross(dif0, wp.spatial_top(cvel0))
+ wbody0 = site_bodyid[id0]
+
+ # second object is geom: process site-geom-site
+ if (type1 == int(WrapType.SPHERE.value)) or (type1 == int(WrapType.CYLINDER.value)):
+ # TODO(team): derivatives of util_misc.wrap
+ return
+
+ # complete sequence
+ wbody1 = site_bodyid[id1]
+ wpnt1 = site_xpos_in[worldid, id1]
+
+ bodyid1 = site_bodyid[id1]
+ pos1 = site_xpos_in[worldid, id1]
+ cvel1 = cvel_in[worldid, bodyid1]
+ subtree_com1 = subtree_com_in[worldid, body_rootid[bodyid1]]
+ dif1 = pos1 - subtree_com1
+ wvel1 = wp.spatial_bottom(cvel1) - wp.cross(dif1, wp.spatial_top(cvel1))
+
+ # accumulate moments if consecutive points are in different bodies
+ if wbody0 != wbody1:
+ # dpnt = 3D position difference, normalize
+ dpnt, norm = math.normalize_with_norm(wpnt1 - wpnt0)
+
+ # dvel = d / dt (dpnt)
+ dvel = wvel1 - wvel0
+ dot = wp.dot(dpnt, dvel)
+ dvel += dpnt * (-dot)
+ if norm > MJ_MINVAL:
+ dvel /= norm
+ else:
+ dvel = wp.vec3(0.0)
+
+ # get endpoint Jacobian time derivatives, subtract
+ # TODO(team): parallelize?
+ for i in range(nv):
+ jac1, _ = support.jac_dot(
+ body_parentid,
+ body_rootid,
+ jnt_type,
+ jnt_dofadr,
+ dof_bodyid,
+ dof_jntid,
+ subtree_com_in,
+ cdof_in,
+ cvel_in,
+ cdof_dot_in,
+ wpnt0,
+ wbody0,
+ i,
+ worldid,
+ )
+ jac2, _ = support.jac_dot(
+ body_parentid,
+ body_rootid,
+ jnt_type,
+ jnt_dofadr,
+ dof_bodyid,
+ dof_jntid,
+ subtree_com_in,
+ cdof_in,
+ cvel_in,
+ cdof_dot_in,
+ wpnt1,
+ wbody1,
+ i,
+ worldid,
+ )
+ jacdif = jac2 - jac1
+
+ # chain rule, first term: Jdot += d / dt (jac2 - jac1) * dpnt
+ Jdot = wp.dot(jacdif, dpnt)
+
+ # get endpoint Jacobians, subtract
+ jac1, _ = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ wpnt0,
+ wbody0,
+ i,
+ worldid,
+ )
+ jac2, _ = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ wpnt1,
+ wbody1,
+ i,
+ worldid,
+ )
+ jacdif = jac2 - jac1
+
+ # chain rule, second term: Jdot += (jac2 - jac1) * d / dt (dpnt)
+ Jdot += wp.dot(jacdif, dvel)
+
+ ten_Jdot_out[worldid, tenid, i] += Jdot / divisor
+
+ # TODO(team): j += 2 if geom wrapping
+ j += 1
+
+
+@wp.kernel
+def _tendon_bias_coef(
+ # Model:
+ tendon_armature: wp.array2d(dtype=float),
+ # Data in:
+ qvel_in: wp.array2d(dtype=float),
+ ten_Jdot_in: wp.array3d(dtype=float),
+ # Data out:
+ ten_bias_coef_out: wp.array2d(dtype=float),
+):
+ worldid, tenid, dofid = wp.tid()
+
+ armature = tendon_armature[worldid, tenid]
+ if armature == 0.0:
+ return
+
+ ten_Jdot = ten_Jdot_in[worldid, tenid, dofid]
+ if ten_Jdot == 0.0:
+ return
+
+ wp.atomic_add(ten_bias_coef_out[worldid], tenid, ten_Jdot * qvel_in[worldid, dofid])
+
+
+@wp.kernel
+def _tendon_bias_qfrc(
+ # Model:
+ tendon_armature: wp.array2d(dtype=float),
+ # Data in:
+ ten_J_in: wp.array3d(dtype=float),
+ ten_bias_coef_in: wp.array2d(dtype=float),
+ # Out:
+ qfrc_out: wp.array2d(dtype=float),
+):
+ worldid, tenid, dofid = wp.tid()
+
+ armature = tendon_armature[worldid, tenid]
+ if armature == 0.0:
+ return
+
+ ten_J = ten_J_in[worldid, tenid, dofid]
+ if ten_J == 0.0:
+ return
+
+ wp.atomic_add(qfrc_out[worldid], dofid, ten_J * armature * ten_bias_coef_in[worldid, tenid])
+
+
+@event_scope
+def tendon_bias(m: Model, d: Data, qfrc: wp.array2d(dtype=float)):
+ """Add bias force due to tendon armature."""
+ d.ten_Jdot.zero_()
+ wp.launch(
+ _tendon_dot,
+ dim=(d.nworld, m.ntendon),
+ inputs=[
+ m.nv,
+ m.body_parentid,
+ m.body_rootid,
+ m.jnt_type,
+ m.jnt_dofadr,
+ m.dof_bodyid,
+ m.dof_jntid,
+ m.site_bodyid,
+ m.tendon_adr,
+ m.tendon_num,
+ m.tendon_armature,
+ m.wrap_objid,
+ m.wrap_prm,
+ m.wrap_type,
+ d.site_xpos,
+ d.subtree_com,
+ d.cdof,
+ d.cvel,
+ d.cdof_dot,
+ ],
+ outputs=[
+ d.ten_Jdot,
+ ],
+ )
+
+ d.ten_bias_coef.zero_()
+ wp.launch(
+ _tendon_bias_coef,
+ dim=(d.nworld, m.ntendon, m.nv),
+ inputs=[
+ m.tendon_armature,
+ d.qvel,
+ d.ten_Jdot,
+ ],
+ outputs=[
+ d.ten_bias_coef,
+ ],
+ )
+
+ wp.launch(
+ _tendon_bias_qfrc,
+ dim=(d.nworld, m.ntendon, m.nv),
+ inputs=[
+ m.tendon_armature,
+ d.ten_J,
+ d.ten_bias_coef,
+ ],
+ outputs=[
+ qfrc,
+ ],
+ )
+
+
+@wp.kernel
+def _comvel_root(cvel_out: wp.array2d(dtype=wp.spatial_vector)):
+ worldid, elementid = wp.tid()
+ cvel_out[worldid, 0][elementid] = 0.0
+
+
+@wp.kernel
+def _comvel_level(
+ # Model:
+ body_parentid: wp.array(dtype=int),
+ body_jntnum: wp.array(dtype=int),
+ body_jntadr: wp.array(dtype=int),
+ body_dofadr: wp.array(dtype=int),
+ jnt_type: wp.array(dtype=int),
+ # Data in:
+ qvel_in: wp.array2d(dtype=float),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ body_tree_: wp.array(dtype=int),
+ # Data out:
+ cvel_out: wp.array2d(dtype=wp.spatial_vector),
+ cdof_dot_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ worldid, nodeid = wp.tid()
+ bodyid = body_tree_[nodeid]
+ dofid = body_dofadr[bodyid]
+ jntid = body_jntadr[bodyid]
+ jntnum = body_jntnum[bodyid]
+ pid = body_parentid[bodyid]
+
+ if jntnum == 0:
+ cvel_out[worldid, bodyid] = cvel_in[worldid, pid]
+ return
+
+ cvel = cvel_in[worldid, pid]
+ qvel = qvel_in[worldid]
+ cdof = cdof_in[worldid]
+
+ for j in range(jntid, jntid + jntnum):
+ jnttype = jnt_type[j]
+
+ if jnttype == wp.static(JointType.FREE.value):
+ cvel += cdof[dofid + 0] * qvel[dofid + 0]
+ cvel += cdof[dofid + 1] * qvel[dofid + 1]
+ cvel += cdof[dofid + 2] * qvel[dofid + 2]
+
+ cdof_dot_out[worldid, dofid + 3] = math.motion_cross(cvel, cdof[dofid + 3])
+ cdof_dot_out[worldid, dofid + 4] = math.motion_cross(cvel, cdof[dofid + 4])
+ cdof_dot_out[worldid, dofid + 5] = math.motion_cross(cvel, cdof[dofid + 5])
+
+ cvel += cdof[dofid + 3] * qvel[dofid + 3]
+ cvel += cdof[dofid + 4] * qvel[dofid + 4]
+ cvel += cdof[dofid + 5] * qvel[dofid + 5]
+
+ dofid += 6
+ elif jnttype == wp.static(JointType.BALL.value):
+ cdof_dot_out[worldid, dofid + 0] = math.motion_cross(cvel, cdof[dofid + 0])
+ cdof_dot_out[worldid, dofid + 1] = math.motion_cross(cvel, cdof[dofid + 1])
+ cdof_dot_out[worldid, dofid + 2] = math.motion_cross(cvel, cdof[dofid + 2])
+
+ cvel += cdof[dofid + 0] * qvel[dofid + 0]
+ cvel += cdof[dofid + 1] * qvel[dofid + 1]
+ cvel += cdof[dofid + 2] * qvel[dofid + 2]
+
+ dofid += 3
+ else:
+ cdof_dot_out[worldid, dofid] = math.motion_cross(cvel, cdof[dofid])
+ cvel += cdof[dofid] * qvel[dofid]
+
+ dofid += 1
+
+ cvel_out[worldid, bodyid] = cvel
+
+
+@event_scope
+def com_vel(m: Model, d: Data):
+ """
+ Computes the spatial velocities (cvel) and the derivative cdof_dot for all bodies.
+
+ Propagates velocities down the kinematic tree, updating the spatial velocity and
+ derivative for each body.
+ """
+ wp.launch(_comvel_root, dim=(d.nworld, 6), inputs=[], outputs=[d.cvel])
+
+ for body_tree in m.body_tree:
+ wp.launch(
+ _comvel_level,
+ dim=(d.nworld, body_tree.size),
+ inputs=[m.body_parentid, m.body_jntnum, m.body_jntadr, m.body_dofadr, m.jnt_type, d.qvel, d.cdof, d.cvel, body_tree],
+ outputs=[d.cvel, d.cdof_dot],
+ )
+
+
+@wp.kernel
+def _transmission(
+ # Model:
+ nv: int,
+ body_parentid: wp.array(dtype=int),
+ body_rootid: wp.array(dtype=int),
+ body_weldid: wp.array(dtype=int),
+ body_dofnum: wp.array(dtype=int),
+ body_dofadr: wp.array(dtype=int),
+ jnt_type: wp.array(dtype=int),
+ jnt_qposadr: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ dof_bodyid: wp.array(dtype=int),
+ dof_parentid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ site_quat: wp.array2d(dtype=wp.quat),
+ actuator_trntype: wp.array(dtype=int),
+ actuator_trnid: wp.array(dtype=wp.vec2i),
+ actuator_gear: wp.array2d(dtype=wp.spatial_vector),
+ actuator_cranklength: wp.array(dtype=float),
+ tendon_adr: wp.array(dtype=int),
+ tendon_num: wp.array(dtype=int),
+ wrap_objid: wp.array(dtype=int),
+ wrap_type: wp.array(dtype=int),
+ # Data in:
+ qpos_in: wp.array2d(dtype=float),
+ xquat_in: wp.array2d(dtype=wp.quat),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ site_xmat_in: wp.array2d(dtype=wp.mat33),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ ten_length_in: wp.array2d(dtype=float),
+ ten_J_in: wp.array3d(dtype=float),
+ # Data out:
+ actuator_length_out: wp.array2d(dtype=float),
+ actuator_moment_out: wp.array3d(dtype=float),
+):
+ worldid, actid = wp.tid()
+ trntype = actuator_trntype[actid]
+ gear = actuator_gear[worldid, actid]
+ if trntype == wp.static(TrnType.JOINT.value) or trntype == wp.static(TrnType.JOINTINPARENT.value):
+ qpos = qpos_in[worldid]
+ jntid = actuator_trnid[actid][0]
+ jnt_typ = jnt_type[jntid]
+ qadr = jnt_qposadr[jntid]
+ vadr = jnt_dofadr[jntid]
+ if jnt_typ == wp.static(JointType.FREE.value):
+ actuator_length_out[worldid, actid] = 0.0
+ if trntype == wp.static(TrnType.JOINTINPARENT.value):
+ quat = wp.normalize(
+ wp.quat(
+ qpos[qadr + 3],
+ qpos[qadr + 4],
+ qpos[qadr + 5],
+ qpos[qadr + 6],
+ )
+ )
+ quat_neg = math.quat_inv(quat)
+ gearaxis = math.rot_vec_quat(wp.spatial_bottom(gear), quat_neg)
+ actuator_moment_out[worldid, actid, vadr + 0] = gear[0]
+ actuator_moment_out[worldid, actid, vadr + 1] = gear[1]
+ actuator_moment_out[worldid, actid, vadr + 2] = gear[2]
+ actuator_moment_out[worldid, actid, vadr + 3] = gearaxis[0]
+ actuator_moment_out[worldid, actid, vadr + 4] = gearaxis[1]
+ actuator_moment_out[worldid, actid, vadr + 5] = gearaxis[2]
+ else:
+ for i in range(6):
+ actuator_moment_out[worldid, actid, vadr + i] = gear[i]
+ elif jnt_typ == wp.static(JointType.BALL.value):
+ q = wp.quat(qpos[qadr + 0], qpos[qadr + 1], qpos[qadr + 2], qpos[qadr + 3])
+ q = wp.normalize(q)
+ axis_angle = math.quat_to_vel(q)
+ gearaxis = wp.spatial_top(gear) # [:3]
+ if trntype == wp.static(TrnType.JOINTINPARENT.value):
+ quat_neg = math.quat_inv(q)
+ gearaxis = math.rot_vec_quat(gearaxis, quat_neg)
+ actuator_length_out[worldid, actid] = wp.dot(axis_angle, gearaxis)
+ for i in range(3):
+ actuator_moment_out[worldid, actid, vadr + i] = gearaxis[i]
+ elif jnt_typ == wp.static(JointType.SLIDE.value) or jnt_typ == wp.static(JointType.HINGE.value):
+ actuator_length_out[worldid, actid] = qpos[qadr] * gear[0]
+ actuator_moment_out[worldid, actid, vadr] = gear[0]
+ else:
+ wp.printf("unrecognized joint type")
+ elif trntype == wp.static(TrnType.SLIDERCRANK.value):
+ # get data
+ trnid = actuator_trnid[actid]
+ id = trnid[0]
+ idslider = trnid[1]
+ gear0 = gear[0]
+ rod = actuator_cranklength[actid]
+ site_xmat = site_xmat_in[worldid, idslider]
+ axis = wp.vec3(site_xmat[0, 2], site_xmat[1, 2], site_xmat[2, 2])
+ site_xpos_id = site_xpos_in[worldid, id]
+ site_xpos_idslider = site_xpos_in[worldid, idslider]
+ vec = site_xpos_id - site_xpos_idslider
+
+ # compute length and determinant
+ # length = a' * v - sqrt(det); det = (a' * v)^2 + r^2 - v' * v
+ av = wp.dot(vec, axis)
+ det = av * av + rod * rod - wp.dot(vec, vec)
+ ok = 1
+ if det <= 0.0:
+ ok = 0
+ sdet = 0.0
+ length = av
+ else:
+ sdet = wp.sqrt(det)
+ length = av - sdet
+
+ actuator_length_out[worldid, actid] = length * gear0
+
+ # compute derivatives of length w.r.t. vec and axis
+ if ok == 1:
+ scale = 1.0 - av / sdet
+ dldv = axis * scale + vec / sdet
+ dlda = vec * scale
+ else:
+ dldv = axis
+ dlda = vec
+
+ # apply chain rule
+ # TODO(team): parallelize?
+ for i in range(nv):
+ # get Jacobians of axis(jacA) and vec(jac)
+ # mj_jacPointAxis
+ jacp, jacr = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ site_xpos_idslider,
+ site_bodyid[idslider],
+ i,
+ worldid,
+ )
+ jacS = jacp
+ jacA = wp.cross(jacr, axis)
+
+ # mj_jacSite
+ jac, _ = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ site_xpos_id,
+ site_bodyid[id],
+ i,
+ worldid,
+ )
+ jac -= jacS
+
+ # apply the chain rule
+ moment = wp.dot(dlda, jacA) + wp.dot(dldv, jac)
+ actuator_moment_out[worldid, actid, i] = moment * gear0
+ elif trntype == wp.static(TrnType.TENDON.value):
+ tenid = actuator_trnid[actid][0]
+
+ gear0 = gear[0]
+ actuator_length_out[worldid, actid] = ten_length_in[worldid, tenid] * gear0
+
+ # fixed
+ adr = tendon_adr[tenid]
+ if wrap_type[adr] == wp.static(WrapType.JOINT.value):
+ ten_num = tendon_num[tenid]
+ for i in range(ten_num):
+ dofadr = jnt_dofadr[wrap_objid[adr + i]]
+ actuator_moment_out[worldid, actid, dofadr] = ten_J_in[worldid, tenid, dofadr] * gear0
+ else: # spatial
+ for dofadr in range(nv):
+ actuator_moment_out[worldid, actid, dofadr] = ten_J_in[worldid, tenid, dofadr] * gear0
+ elif trntype == wp.static(TrnType.BODY.value):
+ # cannot compute meaningful length, set to zero
+ actuator_length_out[worldid, actid] = 0.0
+
+ # initialize moment
+ for i in range(nv):
+ actuator_moment_out[worldid, actid, i] = 0.0
+
+ # moment computed by _transmission_body_moment and _transmission_body_moment_scale
+ elif trntype == int(TrnType.SITE.value):
+ trnid = actuator_trnid[actid]
+ siteid = trnid[0]
+ refid = trnid[1]
+
+ gear = actuator_gear[worldid, actid]
+ gear_translation = wp.spatial_top(gear)
+ gear_rotational = wp.spatial_bottom(gear)
+
+ # reference site undefined
+ if refid == -1:
+ # wrench: gear expressed in global frame
+ site_xmat = site_xmat_in[worldid, siteid]
+ wrench_translation = site_xmat @ gear_translation
+ wrench_rotation = site_xmat @ gear_rotational
+
+ # moment: global Jacobian projected on wrench
+ # TODO(team): parallelize
+ for i in range(nv):
+ jacp, jacr = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ site_xpos_in[worldid, siteid],
+ site_bodyid[siteid],
+ i,
+ worldid,
+ )
+ actuator_length_out[worldid, actid] = 0.0
+ actuator_moment_out[worldid, actid, i] = wp.dot(jacp, wrench_translation) + wp.dot(jacr, wrench_rotation)
+ # reference site defined
+ else:
+ # initialize last dof address for each body
+ bodyid = site_bodyid[siteid]
+ bodyrefid = site_bodyid[refid]
+ b0 = body_weldid[bodyid]
+ b1 = body_weldid[bodyrefid]
+ dofadr0 = body_dofadr[b0] + body_dofnum[b0] - 1
+ dofadr1 = body_dofadr[b1] + body_dofnum[b1] - 1
+
+ # find common ancestral dof, if any
+ dofadr_common = -1
+ if dofadr0 >= 0 and dofadr1 >= 0:
+ # traverse up the tree until common ancestral dof is found
+ while dofadr0 != dofadr1:
+ if dofadr0 < dofadr1:
+ dofadr1 = dof_parentid[dofadr1]
+ else:
+ dofadr0 = dof_parentid[dofadr0]
+
+ if dofadr0 == -1 or dofadr1 == -1:
+ # reached tree root, no common ancestral dof
+ break
+
+ # found common ancestral dof
+ if dofadr0 == dofadr1:
+ dofadr_common = dofadr0
+
+ translational_transmission = not (gear[0] == 0.0 and gear[1] == 0.0 and gear[2] == 0.0)
+ rotational_transmission = not (gear[3] == 0.0 and gear[4] == 0.0 and gear[5] == 0.0)
+
+ site_xpos = site_xpos_in[worldid, siteid]
+ ref_xpos = site_xpos_in[worldid, refid]
+ ref_xmat = site_xmat_in[worldid, refid]
+
+ length = float(0.0)
+
+ if translational_transmission:
+ # vec: site position in reference site frame
+ vec = wp.transpose(ref_xmat) @ (site_xpos - ref_xpos)
+ length += wp.dot(vec, gear_translation)
+
+ wrench_translation = ref_xmat @ gear_translation
+
+ if rotational_transmission:
+ # get site and refsite quats from parent bodies (avoid converting matrix to quat)
+ quat = math.mul_quat(site_quat[worldid, siteid], xquat_in[worldid, bodyid])
+ refquat = math.mul_quat(site_quat[worldid, refid], xquat_in[worldid, bodyrefid])
+
+ # convert difference to expmap (axis-angle)
+ vec = math.quat_sub(quat, refquat)
+ length += wp.dot(vec, gear_rotational)
+
+ wrench_rotation = ref_xmat @ gear_rotational
+
+ actuator_length_out[worldid, actid] = length
+
+ # TODO(team): parallelize
+ for i in range(nv):
+ jacp, jacr = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ site_xpos,
+ site_bodyid[siteid],
+ i,
+ worldid,
+ )
+
+ # jacref: global Jacobian of reference site
+ jacpref, jacrref = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ ref_xpos,
+ site_bodyid[refid],
+ i,
+ worldid,
+ )
+
+ jacpdif = jacp - jacpref
+ jacrdif = jacr - jacrref
+
+ # if common ancestral dof was found, clear the columns of its parental chain
+ da = dofadr_common
+ while da >= 0:
+ if da == i:
+ jacpdif = wp.vec3(0.0)
+ jacrdif = wp.vec3(0.0)
+ break
+ da = dof_parentid[da]
+
+ # moment: global Jacobian projected on wrench
+ moment = float(0.0)
+
+ if translational_transmission:
+ moment += wp.dot(jacpdif, wrench_translation)
+ if rotational_transmission:
+ moment += wp.dot(jacrdif, wrench_rotation)
+
+ actuator_moment_out[worldid, actid, i] = moment
+ else:
+ wp.printf("unhandled transmission type %d\n", trntype)
+
+
+@wp.kernel
+def _transmission_body_moment(
+ # Model:
+ opt_cone: 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),
+ actuator_trnid: wp.array(dtype=wp.vec2i),
+ actuator_trntype_body_adr: wp.array(dtype=int),
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ contact_dist_in: wp.array(dtype=float),
+ contact_pos_in: wp.array(dtype=wp.vec3),
+ contact_frame_in: wp.array(dtype=wp.mat33),
+ contact_includemargin_in: wp.array(dtype=float),
+ contact_dim_in: wp.array(dtype=int),
+ contact_geom_in: wp.array(dtype=wp.vec2i),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ contact_worldid_in: wp.array(dtype=int),
+ efc_J_in: wp.array3d(dtype=float),
+ # Data out:
+ actuator_moment_out: wp.array3d(dtype=float),
+ actuator_trntype_body_ncon_out: wp.array2d(dtype=int),
+):
+ trnbodyid, conid, dofid = wp.tid()
+ actid = actuator_trntype_body_adr[trnbodyid]
+ bodyid = actuator_trnid[actid][0]
+
+ if conid >= ncon_in[0]:
+ return
+
+ worldid = contact_worldid_in[conid]
+
+ # get geom ids
+ geom = contact_geom_in[conid]
+ g1 = geom[0]
+ g2 = geom[1]
+
+ # contact involving flex, continue
+ if g1 < 0 or g2 < 0:
+ return
+
+ # get body ids
+ b1 = geom_bodyid[g1]
+ b2 = geom_bodyid[g2]
+
+ # irrelevant contact, continue
+ if b1 != bodyid and b2 != bodyid:
+ return
+
+ contact_exclude = int(contact_dist_in[conid] >= contact_includemargin_in[conid])
+
+ if dofid == 0:
+ wp.atomic_add(actuator_trntype_body_ncon_out[worldid], trnbodyid, 1)
+
+ # mark contact normals in efc_force
+ if contact_exclude == 0:
+ contact_dim = contact_dim_in[conid]
+ contact_efc_address = contact_efc_address_in[conid]
+
+ if contact_dim == 1 or opt_cone == int(ConeType.ELLIPTIC.value):
+ efc_force = 1.0
+ efcid0 = contact_efc_address[0]
+ wp.atomic_add(actuator_moment_out[worldid, actid], dofid, efc_J_in[worldid, efcid0, dofid] * efc_force)
+
+ else:
+ npyramid = contact_dim - 1 # number of frictional directions
+ efc_force = 0.5 / float(npyramid)
+
+ for j in range(2 * npyramid):
+ efcid = contact_efc_address[j]
+ wp.atomic_add(actuator_moment_out[worldid, actid], dofid, efc_J_in[worldid, efcid, dofid] * efc_force)
+
+ # excluded contact in gap: get Jacobian, accumulate
+ elif contact_exclude == 1:
+ contact_pos = contact_pos_in[conid]
+ contact_frame = contact_frame_in[conid]
+ normal = wp.vec3(contact_frame[0, 0], contact_frame[0, 1], contact_frame[0, 2])
+
+ # get Jacobian difference
+ jacp1, _ = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ contact_pos,
+ b1,
+ dofid,
+ worldid,
+ )
+ jacp2, _ = support.jac(
+ body_parentid,
+ body_rootid,
+ dof_bodyid,
+ subtree_com_in,
+ cdof_in,
+ contact_pos,
+ b2,
+ dofid,
+ worldid,
+ )
+ jacdif = jacp2 - jacp1
+
+ # project Jacobian along the normal of the contact frame
+ wp.atomic_add(actuator_moment_out[worldid, actid], dofid, wp.dot(normal, jacdif))
+
+
+@wp.kernel
+def _transmission_body_moment_scale(
+ # Model:
+ actuator_trntype_body_adr: wp.array(dtype=int),
+ # Data in:
+ actuator_trntype_body_ncon_in: wp.array2d(dtype=int),
+ # Data out:
+ actuator_moment_out: wp.array3d(dtype=float),
+):
+ worldid, trnbodyid, dofid = wp.tid()
+
+ ncon = actuator_trntype_body_ncon_in[worldid, trnbodyid]
+
+ if ncon > 0:
+ actid = actuator_trntype_body_adr[trnbodyid]
+ actuator_moment_out[worldid, actid, dofid] /= -float(ncon)
+
+
+@event_scope
+def transmission(m: Model, d: Data):
+ """
+ Computes actuator/transmission lengths and moments.
+
+ Updates the actuator length and moments for all actuators in the model, including joint
+ and tendon transmissions.
+ """
+ wp.launch(
+ _transmission,
+ dim=[d.nworld, m.nu],
+ inputs=[
+ m.nv,
+ m.body_parentid,
+ m.body_rootid,
+ m.body_weldid,
+ m.body_dofnum,
+ m.body_dofadr,
+ m.jnt_type,
+ m.jnt_qposadr,
+ m.jnt_dofadr,
+ m.dof_bodyid,
+ m.dof_parentid,
+ m.site_bodyid,
+ m.site_quat,
+ m.actuator_trntype,
+ m.actuator_trnid,
+ m.actuator_gear,
+ m.actuator_cranklength,
+ m.tendon_adr,
+ m.tendon_num,
+ m.wrap_objid,
+ m.wrap_type,
+ d.qpos,
+ d.xquat,
+ d.site_xpos,
+ d.site_xmat,
+ d.subtree_com,
+ d.cdof,
+ d.ten_length,
+ d.ten_J,
+ ],
+ outputs=[d.actuator_length, d.actuator_moment],
+ )
+
+ if m.actuator_trntype_body_adr.size > 0:
+ # reset number of active contacts
+ d.actuator_trntype_body_ncon.zero_()
+
+ # compute moments
+ wp.launch(
+ _transmission_body_moment,
+ dim=(
+ m.actuator_trntype_body_adr.size,
+ d.nconmax,
+ m.nv,
+ ),
+ inputs=[
+ m.opt.cone,
+ m.body_parentid,
+ m.body_rootid,
+ m.dof_bodyid,
+ m.geom_bodyid,
+ m.actuator_trnid,
+ m.actuator_trntype_body_adr,
+ d.ncon,
+ d.subtree_com,
+ d.cdof,
+ d.contact.dist,
+ d.contact.pos,
+ d.contact.frame,
+ d.contact.includemargin,
+ d.contact.dim,
+ d.contact.geom,
+ d.contact.efc_address,
+ d.contact.worldid,
+ d.efc.J,
+ ],
+ outputs=[
+ d.actuator_moment,
+ d.actuator_trntype_body_ncon,
+ ],
+ )
+
+ # scale moments
+ wp.launch(
+ _transmission_body_moment_scale,
+ dim=(d.nworld, m.actuator_trntype_body_adr.size, m.nv),
+ inputs=[
+ m.actuator_trntype_body_adr,
+ d.actuator_trntype_body_ncon,
+ ],
+ outputs=[d.actuator_moment],
+ )
+
+
+@wp.kernel
+def _solve_LD_sparse_x_acc_up(
+ # In:
+ L: wp.array3d(dtype=float),
+ qLD_updates_: wp.array(dtype=wp.vec3i),
+ # Out:
+ x: wp.array2d(dtype=float),
+):
+ worldid, nodeid = wp.tid()
+ update = qLD_updates_[nodeid]
+ i, k, Madr_ki = update[0], update[1], update[2]
+ wp.atomic_sub(x[worldid], i, L[worldid, 0, Madr_ki] * x[worldid, k])
+
+
+@wp.kernel
+def _solve_LD_sparse_qLDiag_mul(
+ # In:
+ D: wp.array2d(dtype=float),
+ # Out:
+ out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ out[worldid, dofid] *= D[worldid, dofid]
+
+
+@wp.kernel
+def _solve_LD_sparse_x_acc_down(
+ # In:
+ L: wp.array3d(dtype=float),
+ qLD_updates_: wp.array(dtype=wp.vec3i),
+ # Out:
+ x: wp.array2d(dtype=float),
+):
+ worldid, nodeid = wp.tid()
+ update = qLD_updates_[nodeid]
+ i, k, Madr_ki = update[0], update[1], update[2]
+ wp.atomic_sub(x[worldid], k, L[worldid, 0, Madr_ki] * x[worldid, i])
+
+
+def _solve_LD_sparse(
+ m: Model,
+ d: Data,
+ L: wp.array3d(dtype=float),
+ D: wp.array2d(dtype=float),
+ x: wp.array2d(dtype=float),
+ y: wp.array2d(dtype=float),
+):
+ """Computes sparse backsubstitution: x = inv(L'*D*L)*y"""
+
+ wp.copy(x, y)
+ for qLD_updates in reversed(m.qLD_updates):
+ wp.launch(_solve_LD_sparse_x_acc_up, dim=(d.nworld, qLD_updates.size), inputs=[L, qLD_updates], outputs=[x])
+
+ wp.launch(_solve_LD_sparse_qLDiag_mul, dim=(d.nworld, m.nv), inputs=[D], outputs=[x])
+
+ for qLD_updates in m.qLD_updates:
+ wp.launch(_solve_LD_sparse_x_acc_down, dim=(d.nworld, qLD_updates.size), inputs=[L, qLD_updates], outputs=[x])
+
+
+@cache_kernel
+def _tile_cholesky_solve(tile: TileSet):
+ """Returns a kernel for dense Cholesky backsubstitution of a tile."""
+
+ @nested_kernel
+ def cholesky_solve(
+ # In:
+ L: wp.array3d(dtype=float),
+ y: wp.array2d(dtype=float),
+ adr: wp.array(dtype=int),
+ # Out:
+ x: wp.array2d(dtype=float),
+ ):
+ worldid, nodeid = wp.tid()
+ TILE_SIZE = wp.static(tile.size)
+
+ dofid = adr[nodeid]
+ y_slice = wp.tile_load(y[worldid], shape=(TILE_SIZE,), offset=(dofid,))
+ L_tile = wp.tile_load(L[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid))
+ x_slice = wp.tile_cholesky_solve(L_tile, y_slice)
+ wp.tile_store(x[worldid], x_slice, offset=(dofid,))
+
+ return cholesky_solve
+
+
+def _solve_LD_dense(m: Model, d: Data, L: wp.array3d(dtype=float), x: wp.array2d(dtype=float), y: wp.array2d(dtype=float)):
+ """Computes dense backsubstitution: x = inv(L'*L)*y"""
+ for tile in m.qM_tiles:
+ wp.launch_tiled(
+ _tile_cholesky_solve(tile),
+ dim=(d.nworld, tile.adr.size),
+ inputs=[L, y, tile.adr],
+ outputs=[x],
+ block_dim=m.block_dim.cholesky_solve,
+ )
+
+
+def solve_LD(
+ m: Model,
+ d: Data,
+ L: wp.array3d(dtype=float),
+ D: wp.array2d(dtype=float),
+ x: wp.array2d(dtype=float),
+ y: wp.array2d(dtype=float),
+):
+ """
+ Computes backsubstitution to solve a linear system of the form x = inv(L'*D*L) * y,
+ where L and D are the factors from the Cholesky factorization of the inertia matrix.
+
+ This function dispatches to either a sparse or dense solver depending on Model options.
+
+ Args:
+ m (Model): The model containing factorization and sparsity information.
+ d (Data): The data object containing workspace and factorization results.
+ L (array3d): Lower-triangular factor from the factorization (sparse or dense).
+ D (array2d): Diagonal factor from the factorization (only used for sparse).
+ x (array2d): Output array for the solution.
+ y (array2d): Input right-hand side array.
+ """
+ if m.opt.is_sparse:
+ _solve_LD_sparse(m, d, L, D, x, y)
+ else:
+ _solve_LD_dense(m, d, L, x, y)
+
+
+@event_scope
+def solve_m(m: Model, d: Data, x: wp.array2d(dtype=float), y: wp.array2d(dtype=float)):
+ """
+ Computes backsubstitution: x = qLD * y.
+
+ Args:
+ m (Model): The model containing inertia and factorization information.
+ d (Data): The data object containing factorization results.
+ x (array2d): Output array for the solution.
+ y (array2d): Input right-hand side array.
+ """
+ solve_LD(m, d, d.qLD, d.qLDiagInv, x, y)
+
+
+@cache_kernel
+def _tile_cholesky_factorize_solve(tile: TileSet):
+ """Returns a kernel for dense Cholesky factorization and backsubstitution of a tile."""
+
+ @nested_kernel
+ def cholesky_factorize_solve(
+ # In:
+ M: wp.array3d(dtype=float),
+ y: wp.array2d(dtype=float),
+ adr: wp.array(dtype=int),
+ # Out:
+ x: wp.array2d(dtype=float),
+ ):
+ worldid, nodeid = wp.tid()
+ TILE_SIZE = wp.static(tile.size)
+
+ dofid = adr[nodeid]
+ M_tile = wp.tile_load(M[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid))
+ y_slice = wp.tile_load(y[worldid], shape=(TILE_SIZE,), offset=(dofid,))
+
+ L_tile = wp.tile_cholesky(M_tile)
+ x_slice = wp.tile_cholesky_solve(L_tile, y_slice)
+ wp.tile_store(x[worldid], x_slice, offset=(dofid,))
+
+ return cholesky_factorize_solve
+
+
+def _factor_solve_i_dense(
+ m: Model, d: Data, M: wp.array3d(dtype=float), x: wp.array2d(dtype=float), y: wp.array2d(dtype=float)
+):
+ for tile in m.qM_tiles:
+ wp.launch_tiled(
+ _tile_cholesky_factorize_solve(tile),
+ dim=(d.nworld, tile.adr.size),
+ inputs=[M, y, tile.adr],
+ outputs=[x],
+ block_dim=m.block_dim.cholesky_factorize_solve,
+ )
+
+
+def factor_solve_i(m, d, M, L, D, x, y):
+ """
+ Factorizes and solves the linear system: x = inv(L'*D*L) * y or x = inv(L'*L) * y,
+ where M is an inertia-like matrix and L, D are its Cholesky-like factors.
+
+ This function first factorizes the matrix M (sparse or dense depending on model options),
+ then solves the system for x given right-hand side y.
+
+ Args:
+ m (Model): The model containing factorization and sparsity information.
+ d (Data): The data object containing workspace and factorization results.
+ M (array3d): The inertia-like matrix to factorize.
+ L (array3d): Output lower-triangular factor from the factorization (sparse or dense).
+ D (array2d): Output diagonal factor from the factorization (only used for sparse).
+ x (array2d): Output array for the solution.
+ y (array2d): Input right-hand side array.
+ """
+ if m.opt.is_sparse:
+ _factor_i_sparse(m, d, M, L, D)
+ _solve_LD_sparse(m, d, L, D, x, y)
+ else:
+ _factor_solve_i_dense(m, d, M, x, y)
+
+
+@wp.kernel
+def _subtree_vel_forward(
+ # Model:
+ body_rootid: wp.array(dtype=int),
+ body_mass: wp.array2d(dtype=float),
+ body_inertia: wp.array2d(dtype=wp.vec3),
+ # Data in:
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ ximat_in: wp.array2d(dtype=wp.mat33),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ # Data out:
+ subtree_linvel_out: wp.array2d(dtype=wp.vec3),
+ subtree_angmom_out: wp.array2d(dtype=wp.vec3),
+ subtree_bodyvel_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ worldid, bodyid = wp.tid()
+
+ cvel = cvel_in[worldid, bodyid]
+ ang = wp.spatial_top(cvel)
+ lin = wp.spatial_bottom(cvel)
+ xipos = xipos_in[worldid, bodyid]
+ ximat = ximat_in[worldid, bodyid]
+ subtree_com_root = subtree_com_in[worldid, body_rootid[bodyid]]
+
+ # update linear velocity
+ lin -= wp.cross(xipos - subtree_com_root, ang)
+
+ subtree_linvel_out[worldid, bodyid] = body_mass[worldid, bodyid] * lin
+ dv = wp.transpose(ximat) @ ang
+ dv[0] *= body_inertia[worldid, bodyid][0]
+ dv[1] *= body_inertia[worldid, bodyid][1]
+ dv[2] *= body_inertia[worldid, bodyid][2]
+ subtree_angmom_out[worldid, bodyid] = ximat @ dv
+ subtree_bodyvel_out[worldid, bodyid] = wp.spatial_vector(ang, lin)
+
+
+@wp.kernel
+def _linear_momentum(
+ # Model:
+ body_parentid: wp.array(dtype=int),
+ body_subtreemass: wp.array2d(dtype=float),
+ # Data in:
+ subtree_linvel_in: wp.array2d(dtype=wp.vec3),
+ # In:
+ body_tree_: wp.array(dtype=int),
+ # Data out:
+ subtree_linvel_out: wp.array2d(dtype=wp.vec3),
+):
+ worldid, nodeid = wp.tid()
+ bodyid = body_tree_[nodeid]
+ if bodyid:
+ pid = body_parentid[bodyid]
+ wp.atomic_add(subtree_linvel_out[worldid], pid, subtree_linvel_in[worldid, bodyid])
+ subtree_linvel_out[worldid, bodyid] /= wp.max(MJ_MINVAL, body_subtreemass[worldid, bodyid])
+
+
+@wp.kernel
+def _angular_momentum(
+ # Model:
+ body_parentid: wp.array(dtype=int),
+ body_mass: wp.array2d(dtype=float),
+ body_subtreemass: wp.array2d(dtype=float),
+ # Data in:
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ subtree_linvel_in: wp.array2d(dtype=wp.vec3),
+ subtree_bodyvel_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ body_tree_: wp.array(dtype=int),
+ # Data out:
+ subtree_angmom_out: wp.array2d(dtype=wp.vec3),
+):
+ worldid, nodeid = wp.tid()
+ bodyid = body_tree_[nodeid]
+
+ if bodyid == 0:
+ return
+
+ pid = body_parentid[bodyid]
+
+ xipos = xipos_in[worldid, bodyid]
+ com = subtree_com_in[worldid, bodyid]
+ com_parent = subtree_com_in[worldid, pid]
+ vel = subtree_bodyvel_in[worldid, bodyid]
+ linvel = subtree_linvel_in[worldid, bodyid]
+ linvel_parent = subtree_linvel_in[worldid, pid]
+ mass = body_mass[worldid, bodyid]
+ subtreemass = body_subtreemass[worldid, bodyid]
+
+ # momentum wrt body i
+ dx = xipos - com
+ dv = wp.spatial_bottom(vel) - linvel
+ dp = dv * mass
+ dL = wp.cross(dx, dp)
+
+ # add to subtree i
+ subtree_angmom_out[worldid, bodyid] += dL
+
+ # add to parent
+ wp.atomic_add(subtree_angmom_out[worldid], pid, subtree_angmom_out[worldid, bodyid])
+
+ # momentum wrt parent
+ dx = com - com_parent
+ dv = linvel - linvel_parent
+ dv *= subtreemass
+ dL = wp.cross(dx, dv)
+ wp.atomic_add(subtree_angmom_out[worldid], pid, dL)
+
+
+def subtree_vel(m: Model, d: Data):
+ """
+ Computes subtree linear velocity and angular momentum.
+
+ Computes the linear momentum and angular momentum for each subtree, accumulating
+ contributions up the kinematic tree.
+ """
+
+ # bodywise quantities
+ wp.launch(
+ _subtree_vel_forward,
+ dim=(d.nworld, m.nbody),
+ inputs=[m.body_rootid, m.body_mass, m.body_inertia, d.xipos, d.ximat, d.subtree_com, d.cvel],
+ outputs=[d.subtree_linvel, d.subtree_angmom, d.subtree_bodyvel],
+ )
+
+ # sum body linear momentum recursively up the kinematic tree
+ for body_tree in reversed(m.body_tree):
+ wp.launch(
+ _linear_momentum,
+ dim=[d.nworld, body_tree.size],
+ inputs=[m.body_parentid, m.body_subtreemass, d.subtree_linvel, body_tree],
+ outputs=[d.subtree_linvel],
+ )
+
+ for body_tree in reversed(m.body_tree):
+ wp.launch(
+ _angular_momentum,
+ dim=[d.nworld, body_tree.size],
+ inputs=[
+ m.body_parentid,
+ m.body_mass,
+ m.body_subtreemass,
+ d.xipos,
+ d.subtree_com,
+ d.subtree_linvel,
+ d.subtree_bodyvel,
+ body_tree,
+ ],
+ outputs=[d.subtree_angmom],
+ )
+
+
+@wp.kernel
+def _joint_tendon(
+ # Model:
+ jnt_qposadr: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ wrap_objid: wp.array(dtype=int),
+ wrap_prm: wp.array(dtype=float),
+ tendon_jnt_adr: wp.array(dtype=int),
+ wrap_jnt_adr: wp.array(dtype=int),
+ # Data in:
+ qpos_in: wp.array2d(dtype=float),
+ # Data out:
+ ten_length_out: wp.array2d(dtype=float),
+ ten_J_out: wp.array3d(dtype=float),
+):
+ worldid, wrapid = wp.tid()
+
+ tendon_jnt_adr_ = tendon_jnt_adr[wrapid]
+ wrap_jnt_adr_ = wrap_jnt_adr[wrapid]
+
+ wrap_objid_ = wrap_objid[wrap_jnt_adr_]
+ prm = wrap_prm[wrap_jnt_adr_]
+
+ # add to length
+ L = prm * qpos_in[worldid, jnt_qposadr[wrap_objid_]]
+ # TODO(team): compare atomic_add and for loop
+ wp.atomic_add(ten_length_out[worldid], tendon_jnt_adr_, L)
+
+ # add to moment
+ ten_J_out[worldid, tendon_jnt_adr_, jnt_dofadr[wrap_objid_]] = prm
+
+
+@wp.kernel
+def _spatial_site_tendon(
+ # Model:
+ nv: int,
+ body_parentid: wp.array(dtype=int),
+ body_rootid: wp.array(dtype=int),
+ dof_bodyid: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ wrap_objid: wp.array(dtype=int),
+ tendon_site_pair_adr: wp.array(dtype=int),
+ wrap_site_pair_adr: wp.array(dtype=int),
+ wrap_pulley_scale: wp.array(dtype=float),
+ # Data in:
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ # Data out:
+ ten_length_out: wp.array2d(dtype=float),
+ ten_J_out: wp.array3d(dtype=float),
+):
+ worldid, elementid = wp.tid()
+
+ # site pairs
+ site_pair_adr = wrap_site_pair_adr[elementid]
+ ten_adr = tendon_site_pair_adr[elementid]
+
+ # pulley scaling
+ pulley_scale = wrap_pulley_scale[site_pair_adr]
+
+ id0 = wrap_objid[site_pair_adr + 0]
+ id1 = wrap_objid[site_pair_adr + 1]
+
+ pnt0 = site_xpos_in[worldid, id0]
+ pnt1 = site_xpos_in[worldid, id1]
+ dif = pnt1 - pnt0
+ vec, length = math.normalize_with_norm(dif)
+ wp.atomic_add(ten_length_out[worldid], ten_adr, length * pulley_scale)
+
+ if length < MJ_MINVAL:
+ vec = wp.vec3(1.0, 0.0, 0.0)
+
+ body0 = site_bodyid[id0]
+ body1 = site_bodyid[id1]
+ if body0 != body1:
+ # TODO(team): parallelize
+ for i in range(nv):
+ jacp1, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pnt0, body0, i, worldid)
+ jacp2, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pnt1, body1, i, worldid)
+
+ J = wp.dot(jacp2 - jacp1, vec)
+ if J:
+ wp.atomic_add(ten_J_out[worldid, ten_adr], i, J * pulley_scale)
+
+
+@wp.kernel
+def _spatial_geom_tendon(
+ # Model:
+ nv: 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),
+ site_bodyid: wp.array(dtype=int),
+ wrap_objid: wp.array(dtype=int),
+ wrap_prm: wp.array(dtype=float),
+ wrap_type: wp.array(dtype=int),
+ tendon_geom_adr: wp.array(dtype=int),
+ wrap_geom_adr: wp.array(dtype=int),
+ wrap_pulley_scale: wp.array(dtype=float),
+ # Data in:
+ geom_xpos_in: wp.array2d(dtype=wp.vec3),
+ geom_xmat_in: wp.array2d(dtype=wp.mat33),
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ # Data out:
+ ten_length_out: wp.array2d(dtype=float),
+ ten_J_out: wp.array3d(dtype=float),
+ wrap_geom_xpos_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ worldid, elementid = wp.tid()
+ wrap_adr = wrap_geom_adr[elementid]
+ ten_adr = tendon_geom_adr[elementid]
+
+ # pulley scaling
+ pulley_scale = wrap_pulley_scale[wrap_adr]
+
+ # site-geom-site
+ wrap_objid_site0 = wrap_objid[wrap_adr - 1]
+ wrap_objid_geom = wrap_objid[wrap_adr + 0]
+ wrap_objid_site1 = wrap_objid[wrap_adr + 1]
+
+ # get site positions before and after geom
+ site_pnt0 = site_xpos_in[worldid, wrap_objid_site0]
+ site_pnt1 = site_xpos_in[worldid, wrap_objid_site1]
+
+ # get geom information
+ geom_xpos = geom_xpos_in[worldid, wrap_objid_geom]
+ geom_xmat = geom_xmat_in[worldid, wrap_objid_geom]
+ geomsize = geom_size[worldid, wrap_objid_geom][0]
+ geom_type = wrap_type[wrap_adr]
+
+ # get body ids for site-geom-site instances
+ bodyid_site0 = site_bodyid[wrap_objid_site0]
+ bodyid_geom = geom_bodyid[wrap_objid_geom]
+ bodyid_site1 = site_bodyid[wrap_objid_site1]
+
+ # find wrap object sidesite (if it exists)
+ sideid = int(wp.round(wrap_prm[wrap_adr]))
+ if sideid >= 0:
+ side = site_xpos_in[worldid, sideid]
+ else:
+ side = wp.vec3(wp.inf)
+
+ # compute geom wrap length and connect points (if wrap occurs)
+ length_geomgeom, geom_pnt0, geom_pnt1 = util_misc.wrap(site_pnt0, site_pnt1, geom_xpos, geom_xmat, geomsize, geom_type, side)
+
+ # store geom points
+ wrap_geom_xpos_out[worldid, elementid] = wp.spatial_vector(geom_pnt0, geom_pnt1)
+
+ if length_geomgeom >= 0.0:
+ dif_sitegeom = geom_pnt0 - site_pnt0
+ dif_geomsite = site_pnt1 - geom_pnt1
+ vec_sitegeom, length_sitegeom = math.normalize_with_norm(dif_sitegeom)
+ vec_geomsite, length_geomsite = math.normalize_with_norm(dif_geomsite)
+
+ # length
+ length_sitegeomsite = length_sitegeom + length_geomgeom + length_geomsite
+
+ if length_sitegeomsite:
+ wp.atomic_add(ten_length_out[worldid], ten_adr, length_sitegeomsite * pulley_scale)
+
+ # moment
+ if length_sitegeom < MJ_MINVAL:
+ vec_sitegeom = wp.vec3(1.0, 0.0, 0.0)
+
+ if length_geomsite < MJ_MINVAL:
+ vec_geomsite = wp.vec3(1.0, 0.0, 0.0)
+
+ dif_body_sitegeom = bodyid_site0 != bodyid_geom
+ dif_body_geomsite = bodyid_geom != bodyid_site1
+
+ # TODO(team): parallelize
+ for i in range(nv):
+ J = float(0.0)
+ # site-geom
+ if dif_body_sitegeom:
+ jacp_site0, _ = support.jac(
+ body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt0, bodyid_site0, i, worldid
+ )
+
+ jacp_geom0, _ = support.jac(
+ body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, geom_pnt0, bodyid_geom, i, worldid
+ )
+
+ J += wp.dot(jacp_geom0 - jacp_site0, vec_sitegeom)
+
+ # geom-site
+ if dif_body_geomsite:
+ jacp_geom1, _ = support.jac(
+ body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, geom_pnt1, bodyid_geom, i, worldid
+ )
+
+ jacp_site1, _ = support.jac(
+ body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt1, bodyid_site1, i, worldid
+ )
+
+ J += wp.dot(jacp_site1 - jacp_geom1, vec_geomsite)
+
+ if J:
+ wp.atomic_add(ten_J_out[worldid, ten_adr], i, J * pulley_scale)
+ else:
+ dif_sitesite = site_pnt1 - site_pnt0
+ vec_sitesite, length_sitesite = math.normalize_with_norm(dif_sitesite)
+
+ # length
+ if length_sitesite:
+ wp.atomic_add(ten_length_out[worldid], ten_adr, length_sitesite * pulley_scale)
+
+ # moment
+ if length_sitesite < MJ_MINVAL:
+ vec_sitesite = wp.vec3(1.0, 0.0, 0.0)
+
+ if bodyid_site0 != bodyid_site1:
+ # TODO(team): parallelize
+ for i in range(nv):
+ jacp1, _ = support.jac(
+ body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt0, bodyid_site0, i, worldid
+ )
+ jacp2, _ = support.jac(
+ body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt1, bodyid_site1, i, worldid
+ )
+
+ J = wp.dot(jacp2 - jacp1, vec_sitesite)
+
+ if J:
+ wp.atomic_add(ten_J_out[worldid, ten_adr], i, J * pulley_scale)
+
+
+@wp.kernel
+def _spatial_tendon_wrap(
+ # Model:
+ ntendon: int,
+ tendon_adr: wp.array(dtype=int),
+ tendon_num: wp.array(dtype=int),
+ wrap_objid: wp.array(dtype=int),
+ wrap_type: wp.array(dtype=int),
+ # Data in:
+ site_xpos_in: wp.array2d(dtype=wp.vec3),
+ wrap_geom_xpos_in: wp.array2d(dtype=wp.spatial_vector),
+ # Data out:
+ ten_wrapadr_out: wp.array2d(dtype=int),
+ ten_wrapnum_out: wp.array2d(dtype=int),
+ wrap_obj_out: wp.array2d(dtype=wp.vec2i),
+ wrap_xpos_out: wp.array2d(dtype=wp.spatial_vector),
+):
+ worldid = wp.tid()
+
+ wrapcount = int(0)
+ wrapgeomid = int(0)
+ for i in range(ntendon):
+ adr = tendon_adr[i]
+ ten_wrapadr_out[worldid, i] = wrapcount
+ wrapnum = int(0)
+ tendonnum = tendon_num[i]
+
+ # process fixed tendon
+ if wrap_type[adr] == int(WrapType.JOINT.value):
+ continue
+
+ # process spatial tendon
+ j = int(0)
+ while j < tendonnum - 1:
+ # get 1st and 2nd object
+ type0 = wrap_type[adr + j + 0]
+ type1 = wrap_type[adr + j + 1]
+ id0 = wrap_objid[adr + j + 0]
+ id1 = wrap_objid[adr + j + 1]
+
+ # pulley
+ pulley0 = type0 == int(WrapType.PULLEY.value)
+ if pulley0 or type1 == int(WrapType.PULLEY.value):
+ if pulley0:
+ row = wrapcount // 2
+ col = wrapcount % 2
+ wrap_xpos_out[worldid, row][3 * col + 0] = 0.0
+ wrap_xpos_out[worldid, row][3 * col + 1] = 0.0
+ wrap_xpos_out[worldid, row][3 * col + 2] = 0.0
+
+ wrap_obj_out[worldid, row][col] = -2
+
+ wrapnum += 1
+ wrapcount += 1
+
+ # move to next
+ j += 1
+ continue
+
+ # init sequence; assume it starts with site
+ wpnt_site0 = site_xpos_in[worldid, id0]
+
+ # second object is geom: process site-geom-site
+ if type1 == int(WrapType.SPHERE.value) or type1 == int(WrapType.CYLINDER.value):
+ wrap_geom_xpos = wrap_geom_xpos_in[worldid, wrapgeomid]
+ wpnt_geom0 = wp.spatial_top(wrap_geom_xpos)
+ wrapgeomid += 1
+
+ wrapid = id1
+ id1 = wrap_objid[adr + j + 2]
+ if wp.norm_l2(wpnt_geom0) < wp.inf:
+ wpnt_geom1 = wp.spatial_bottom(wrap_geom_xpos)
+ wpnt_site1 = site_xpos_in[worldid, id1]
+
+ # assign to wrap
+ row0 = (wrapcount + 0) // 2
+ col0 = (wrapcount + 0) % 2
+ row1 = (wrapcount + 1) // 2
+ col1 = (wrapcount + 1) % 2
+ row2 = (wrapcount + 2) // 2
+ col2 = (wrapcount + 2) % 2
+ row3 = (wrapcount + 3) // 2
+ col3 = (wrapcount + 3) % 2
+
+ wrap_xpos_out[worldid, row0][3 * col0 + 0] = wpnt_site0[0]
+ wrap_xpos_out[worldid, row0][3 * col0 + 1] = wpnt_site0[1]
+ wrap_xpos_out[worldid, row0][3 * col0 + 2] = wpnt_site0[2]
+
+ wrap_xpos_out[worldid, row1][3 * col1 + 0] = wpnt_geom0[0]
+ wrap_xpos_out[worldid, row1][3 * col1 + 1] = wpnt_geom0[1]
+ wrap_xpos_out[worldid, row1][3 * col1 + 2] = wpnt_geom0[2]
+
+ wrap_xpos_out[worldid, row2][3 * col2 + 0] = wpnt_geom1[0]
+ wrap_xpos_out[worldid, row2][3 * col2 + 1] = wpnt_geom1[1]
+ wrap_xpos_out[worldid, row2][3 * col2 + 2] = wpnt_geom1[2]
+
+ wrap_xpos_out[worldid, row3][3 * col3 + 0] = wpnt_site1[0]
+ wrap_xpos_out[worldid, row3][3 * col3 + 1] = wpnt_site1[1]
+ wrap_xpos_out[worldid, row3][3 * col3 + 2] = wpnt_site1[2]
+
+ wrap_obj_out[worldid, row0][col0] = -1
+ wrap_obj_out[worldid, row1][col1] = wrapid
+ wrap_obj_out[worldid, row2][col2] = wrapid
+
+ wrapnum += 3
+ wrapcount += 3
+ j += 2
+
+ else:
+ row0 = (wrapcount + 0) // 2
+ col0 = (wrapcount + 0) % 2
+
+ wrap_xpos_out[worldid, row0][3 * col0 + 0] = wpnt_site0[0]
+ wrap_xpos_out[worldid, row0][3 * col0 + 1] = wpnt_site0[1]
+ wrap_xpos_out[worldid, row0][3 * col0 + 2] = wpnt_site0[2]
+
+ wrap_obj_out[worldid, row0][col0] = -1
+
+ wrapnum += 1
+ wrapcount += 1
+ j += 2
+
+ else:
+ row0 = (wrapcount + 0) // 2
+ col0 = (wrapcount + 0) % 2
+
+ wrap_xpos_out[worldid, row0][3 * col0 + 0] = wpnt_site0[0]
+ wrap_xpos_out[worldid, row0][3 * col0 + 1] = wpnt_site0[1]
+ wrap_xpos_out[worldid, row0][3 * col0 + 2] = wpnt_site0[2]
+
+ wrap_obj_out[worldid, row0][col0] = -1
+
+ wrapnum += 1
+ wrapcount += 1
+ j += 1
+
+ # assign last site before pulley or tendon end
+ if adr + j + 1 < wrap_type.shape[0]:
+ last_before_pulley = wrap_type[adr + j + 1] == int(WrapType.PULLEY.value)
+ else:
+ last_before_pulley = False
+
+ if j == tendonnum - 1 or last_before_pulley:
+ row0 = (wrapcount + 0) // 2
+ col0 = (wrapcount + 0) % 2
+
+ wpnt_site1 = site_xpos_in[worldid, id1]
+ wrap_xpos_out[worldid, row0][3 * col0 + 0] = wpnt_site1[0]
+ wrap_xpos_out[worldid, row0][3 * col0 + 1] = wpnt_site1[1]
+ wrap_xpos_out[worldid, row0][3 * col0 + 2] = wpnt_site1[2]
+
+ wrap_obj_out[worldid, row0][col0] = -1
+ wrapnum += 1
+ wrapcount += 1
+
+ ten_wrapnum_out[worldid, i] = wrapnum
+
+
+def tendon(m: Model, d: Data):
+ """
+ Computes tendon lengths and moments.
+
+ Updates the tendon length and moment arrays for all tendons in the model, including joint,
+ site, and geom tendons.
+ """
+ if not m.ntendon:
+ return
+
+ d.ten_length.zero_()
+ d.ten_J.zero_()
+
+ # process joint tendons
+ wp.launch(
+ _joint_tendon,
+ dim=(d.nworld, m.wrap_jnt_adr.size),
+ inputs=[m.jnt_qposadr, m.jnt_dofadr, m.wrap_objid, m.wrap_prm, m.tendon_jnt_adr, m.wrap_jnt_adr, d.qpos],
+ outputs=[d.ten_length, d.ten_J],
+ )
+
+ spatial_site = m.wrap_site_pair_adr.size > 0
+ spatial_geom = m.wrap_geom_adr.size > 0
+
+ if spatial_site or spatial_geom:
+ d.wrap_xpos.zero_()
+ d.wrap_obj.zero_()
+
+ # process spatial site tendons
+ wp.launch(
+ _spatial_site_tendon,
+ dim=(d.nworld, m.wrap_site_pair_adr.size),
+ inputs=[
+ m.nv,
+ m.body_parentid,
+ m.body_rootid,
+ m.dof_bodyid,
+ m.site_bodyid,
+ m.wrap_objid,
+ m.tendon_site_pair_adr,
+ m.wrap_site_pair_adr,
+ m.wrap_pulley_scale,
+ d.site_xpos,
+ d.subtree_com,
+ d.cdof,
+ ],
+ outputs=[d.ten_length, d.ten_J],
+ )
+
+ # process spatial geom tendons
+ wp.launch(
+ _spatial_geom_tendon,
+ dim=(d.nworld, m.wrap_geom_adr.size),
+ inputs=[
+ m.nv,
+ m.body_parentid,
+ m.body_rootid,
+ m.dof_bodyid,
+ m.geom_bodyid,
+ m.geom_size,
+ m.site_bodyid,
+ m.wrap_objid,
+ m.wrap_prm,
+ m.wrap_type,
+ m.tendon_geom_adr,
+ m.wrap_geom_adr,
+ m.wrap_pulley_scale,
+ d.geom_xpos,
+ d.geom_xmat,
+ d.site_xpos,
+ d.subtree_com,
+ d.cdof,
+ ],
+ outputs=[d.ten_length, d.ten_J, d.wrap_geom_xpos],
+ )
+
+ if spatial_site or spatial_geom:
+ wp.launch(
+ _spatial_tendon_wrap,
+ dim=(d.nworld,),
+ inputs=[m.ntendon, m.tendon_adr, m.tendon_num, m.wrap_objid, m.wrap_type, d.site_xpos, d.wrap_geom_xpos],
+ outputs=[d.ten_wrapadr, d.ten_wrapnum, d.wrap_obj, d.wrap_xpos],
+ )
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth_test.py
new file mode 100644
index 00000000..9ceed777
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth_test.py
@@ -0,0 +1,398 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 smooth dynamics functions."""
+
+import mujoco
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+from absl.testing import parameterized
+
+import mujoco_warp as mjwarp
+
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+from mujoco.mjx.third_party.mujoco_warp._src import types
+
+# tolerance for difference between MuJoCo and MJWarp smooth calculations - mostly
+# due to float precision
+_TOLERANCE = 5e-5
+
+
+def _assert_eq(a, b, name):
+ tol = _TOLERANCE * 10 # avoid test noise
+ err_msg = f"mismatch: {name}"
+ np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+
+
+class SmoothTest(parameterized.TestCase):
+ def test_kinematics(self):
+ """Tests kinematics."""
+ _, mjd, m, d = test_util.fixture("pendula.xml")
+
+ for arr in (
+ d.xanchor,
+ d.xaxis,
+ d.xpos,
+ d.xquat,
+ d.xmat,
+ d.xipos,
+ d.ximat,
+ d.geom_xpos,
+ d.geom_xmat,
+ d.site_xpos,
+ d.site_xmat,
+ ):
+ arr.zero_()
+
+ mjwarp.kinematics(m, d)
+
+ _assert_eq(d.xanchor.numpy()[0], mjd.xanchor, "xanchor")
+ _assert_eq(d.xaxis.numpy()[0], mjd.xaxis, "xaxis")
+ _assert_eq(d.xpos.numpy()[0], mjd.xpos, "xpos")
+ _assert_eq(d.xquat.numpy()[0], mjd.xquat, "xquat")
+ _assert_eq(d.xmat.numpy()[0], mjd.xmat.reshape((-1, 3, 3)), "xmat")
+ _assert_eq(d.xipos.numpy()[0], mjd.xipos, "xipos")
+ _assert_eq(d.ximat.numpy()[0], mjd.ximat.reshape((-1, 3, 3)), "ximat")
+ _assert_eq(d.geom_xpos.numpy()[0], mjd.geom_xpos, "geom_xpos")
+ _assert_eq(d.geom_xmat.numpy()[0], mjd.geom_xmat.reshape((-1, 3, 3)), "geom_xmat")
+ _assert_eq(d.site_xpos.numpy()[0], mjd.site_xpos, "site_xpos")
+ _assert_eq(d.site_xmat.numpy()[0], mjd.site_xmat.reshape((-1, 3, 3)), "site_xmat")
+
+ def test_com_pos(self):
+ """Tests com_pos."""
+ _, mjd, m, d = test_util.fixture("pendula.xml")
+
+ for arr in (d.subtree_com, d.cinert, d.cdof):
+ arr.zero_()
+
+ mjwarp.com_pos(m, d)
+ _assert_eq(d.subtree_com.numpy()[0], mjd.subtree_com, "subtree_com")
+ _assert_eq(d.cinert.numpy()[0], mjd.cinert, "cinert")
+ _assert_eq(d.cdof.numpy()[0], mjd.cdof, "cdof")
+
+ def test_camlight(self):
+ """Tests camlight."""
+ _, mjd, m, d = test_util.fixture("pendula.xml")
+
+ d.cam_xpos.zero_()
+ d.cam_xmat.zero_()
+ d.light_xpos.zero_()
+ d.light_xdir.zero_()
+
+ mjwarp.camlight(m, d)
+ _assert_eq(d.cam_xpos.numpy()[0], mjd.cam_xpos, "cam_xpos")
+ _assert_eq(d.cam_xmat.numpy()[0], mjd.cam_xmat.reshape((-1, 3, 3)), "cam_xmat")
+ _assert_eq(d.light_xpos.numpy()[0], mjd.light_xpos, "light_xpos")
+ _assert_eq(d.light_xdir.numpy()[0], mjd.light_xdir, "light_xdir")
+
+ @parameterized.parameters(True, False)
+ def test_crb(self, sparse: bool):
+ """Tests crb."""
+ mjm, mjd, m, d = test_util.fixture("pendula.xml", sparse=sparse)
+
+ d.crb.zero_()
+
+ mjwarp.crb(m, d)
+ _assert_eq(d.crb.numpy()[0], mjd.crb, "crb")
+
+ if sparse:
+ _assert_eq(d.qM.numpy()[0, 0], mjd.qM, "qM")
+ else:
+ qM = np.zeros((mjm.nv, mjm.nv))
+ mujoco.mj_fullM(mjm, qM, mjd.qM)
+ _assert_eq(d.qM.numpy()[0], qM, "qM")
+
+ @parameterized.parameters(True, False)
+ def test_factor_m(self, sparse: bool):
+ """Tests factor_m."""
+ _, mjd, m, d = test_util.fixture("pendula.xml", sparse=sparse)
+
+ qLD = d.qLD.numpy()[0].copy()
+ for arr in (d.qLD, d.qLDiagInv):
+ arr.zero_()
+
+ mjwarp.factor_m(m, d)
+
+ if sparse:
+ _assert_eq(d.qLD.numpy()[0, 0], mjd.qLD, "qLD (sparse)")
+ _assert_eq(d.qLDiagInv.numpy()[0], mjd.qLDiagInv, "qLDiagInv")
+ else:
+ _assert_eq(d.qLD.numpy()[0], qLD, "qLD (dense)")
+
+ @parameterized.parameters(True, False)
+ def test_solve_m(self, sparse: bool):
+ """Tests solve_m."""
+ mjm, mjd, m, d = test_util.fixture("pendula.xml", sparse=sparse)
+
+ qfrc_smooth = np.tile(mjd.qfrc_smooth, (1, 1))
+ qacc_smooth = np.zeros(
+ shape=(
+ 1,
+ mjm.nv,
+ ),
+ dtype=float,
+ )
+ mujoco.mj_solveM(mjm, mjd, qacc_smooth, qfrc_smooth)
+
+ d.qacc_smooth.zero_()
+
+ mjwarp.solve_m(m, d, d.qacc_smooth, d.qfrc_smooth)
+ _assert_eq(d.qacc_smooth.numpy()[0], qacc_smooth[0], "qacc_smooth")
+
+ @parameterized.parameters(True, False)
+ def test_rne(self, gravity):
+ """Tests rne."""
+ _, mjd, m, d = test_util.fixture("pendula.xml", gravity=gravity)
+
+ d.qfrc_bias.zero_()
+
+ mjwarp.rne(m, d)
+ _assert_eq(d.qfrc_bias.numpy()[0], mjd.qfrc_bias, "qfrc_bias")
+
+ @parameterized.parameters(True, False)
+ def test_rne_postconstraint(self, gravity):
+ """Tests rne_postconstraint."""
+ mjm, mjd, m, d = test_util.fixture("pendula.xml", gravity=gravity)
+
+ mjd.xfrc_applied = np.random.uniform(low=-0.01, high=0.01, size=mjd.xfrc_applied.shape)
+ d.xfrc_applied = wp.array(np.expand_dims(mjd.xfrc_applied, axis=0), dtype=wp.spatial_vector)
+
+ mujoco.mj_rnePostConstraint(mjm, mjd)
+
+ for arr in (d.cacc, d.cfrc_int, d.cfrc_ext):
+ arr.zero_()
+
+ mjwarp.rne_postconstraint(m, d)
+
+ _assert_eq(d.cacc.numpy()[0], mjd.cacc, "cacc")
+ _assert_eq(d.cfrc_int.numpy()[0], mjd.cfrc_int, "cfrc_int")
+ _assert_eq(d.cfrc_ext.numpy()[0], mjd.cfrc_ext, "cfrc_ext")
+
+ _EQUALITY = """
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """
+ mjm, mjd, m, d = test_util.fixture(xml=_EQUALITY, kick=True, keyframe=0)
+
+ mujoco.mj_rnePostConstraint(mjm, mjd)
+
+ d.cfrc_ext.zero_()
+ mjwarp.rne_postconstraint(m, d)
+
+ _assert_eq(d.cfrc_ext.numpy()[0], mjd.cfrc_ext, "cfrc_ext (equality)")
+
+ mjm, mjd, m, d = test_util.fixture("constraints.xml", keyframe=1, equality=False)
+
+ mujoco.mj_rnePostConstraint(mjm, mjd)
+
+ d.cfrc_ext.zero_()
+
+ # clear equality constraint counts
+ d.ne_connect.zero_()
+ d.ne_weld.zero_()
+ d.ne_jnt.zero_()
+
+ mjwarp.rne_postconstraint(m, d)
+
+ _assert_eq(d.cfrc_ext.numpy()[0], mjd.cfrc_ext, "cfrc_ext (contact)")
+
+ def test_com_vel(self):
+ """Tests com_vel."""
+ _, mjd, m, d = test_util.fixture("pendula.xml")
+
+ for arr in (d.cvel, d.cdof_dot):
+ arr.zero_()
+
+ mjwarp.com_vel(m, d)
+ _assert_eq(d.cvel.numpy()[0], mjd.cvel, "cvel")
+ _assert_eq(d.cdof_dot.numpy()[0], mjd.cdof_dot, "cdof_dot")
+
+ @parameterized.parameters("pendula.xml", "actuation/site.xml", "actuation/slidercrank.xml")
+ def test_transmission(self, xml):
+ """Tests transmission."""
+ mjm, mjd, m, d = test_util.fixture(xml)
+
+ for arr in (d.actuator_length, d.actuator_moment):
+ arr.zero_()
+
+ actuator_moment = np.zeros((mjm.nu, mjm.nv))
+ mujoco.mju_sparse2dense(
+ actuator_moment,
+ mjd.actuator_moment,
+ mjd.moment_rownnz,
+ mjd.moment_rowadr,
+ mjd.moment_colind,
+ )
+
+ mjwarp._src.smooth.transmission(m, d)
+ _assert_eq(d.actuator_length.numpy()[0], mjd.actuator_length, "actuator_length")
+ _assert_eq(d.actuator_moment.numpy()[0], actuator_moment, "actuator_moment")
+
+ @parameterized.product(keyframe=list(range(4)), cone=list(types.ConeType))
+ def test_actuator_adhesion(self, keyframe, cone):
+ """Tests adhesion actuator."""
+ mjm, mjd, m, d = test_util.fixture("actuation/adhesion.xml", keyframe=keyframe, cone=cone)
+
+ d.actuator_length.zero_()
+ d.actuator_moment.zero_()
+ mjwarp._src.collision_driver.collision(m, d) # compute contact.includemargin
+ mjwarp._src.constraint.make_constraint(m, d) # compute contact.efc_address
+ mjwarp._src.smooth.transmission(m, d)
+
+ actuator_moment = np.zeros((mjm.nu, mjm.nv))
+ mujoco.mju_sparse2dense(actuator_moment, mjd.actuator_moment, mjd.moment_rownnz, mjd.moment_rowadr, mjd.moment_colind)
+
+ _assert_eq(d.actuator_length.numpy()[0], mjd.actuator_length, "actuator_length")
+ _assert_eq(d.actuator_moment.numpy()[0], actuator_moment, "acutator_moment")
+
+ def test_subtree_vel(self):
+ """Tests subtree_vel."""
+ mjm, mjd, m, d = test_util.fixture("pendula.xml")
+
+ for arr in (d.subtree_linvel, d.subtree_angmom):
+ arr.zero_()
+
+ mujoco.mj_subtreeVel(mjm, mjd)
+ mjwarp.subtree_vel(m, d)
+
+ _assert_eq(d.subtree_linvel.numpy()[0], mjd.subtree_linvel, "subtree_linvel")
+ _assert_eq(d.subtree_angmom.numpy()[0], mjd.subtree_angmom, "subtree_angmom")
+
+ @parameterized.parameters(
+ "tendon/fixed.xml",
+ "tendon/site.xml",
+ "tendon/pulley_site.xml",
+ "tendon/fixed_site.xml",
+ "tendon/pulley_fixed_site.xml",
+ "tendon/site_fixed.xml",
+ "tendon/pulley_site_fixed.xml",
+ "tendon/wrap.xml",
+ "tendon/pulley_wrap.xml",
+ )
+ def test_tendon(self, xml):
+ """Tests tendon."""
+ mjm, mjd, m, d = test_util.fixture(xml, keyframe=0)
+
+ for arr in (d.ten_length, d.ten_J, d.actuator_length, d.actuator_moment):
+ arr.zero_()
+
+ mjwarp.tendon(m, d)
+ mjwarp.transmission(m, d)
+
+ _assert_eq(d.ten_length.numpy()[0], mjd.ten_length, "ten_length")
+ _assert_eq(d.ten_J.numpy()[0], mjd.ten_J.reshape((mjm.ntendon, mjm.nv)), "ten_J")
+ _assert_eq(d.wrap_xpos.numpy()[0], mjd.wrap_xpos, "wrap_xpos")
+ _assert_eq(d.wrap_obj.numpy()[0], mjd.wrap_obj, "wrap_obj")
+ _assert_eq(d.ten_wrapnum.numpy()[0], mjd.ten_wrapnum, "ten_wrapnum")
+ _assert_eq(d.ten_wrapadr.numpy()[0], mjd.ten_wrapadr, "ten_wrapadr")
+ _assert_eq(d.actuator_length.numpy()[0], mjd.actuator_length, "actuator_length")
+ actuator_moment = np.zeros((mjm.nu, mjm.nv))
+ mujoco.mju_sparse2dense(
+ actuator_moment,
+ mjd.actuator_moment,
+ mjd.moment_rownnz,
+ mjd.moment_rowadr,
+ mjd.moment_colind,
+ )
+ _assert_eq(d.actuator_moment.numpy()[0], actuator_moment, "actuator_moment")
+
+ @parameterized.parameters(True, False)
+ def test_factor_solve_i(self, sparse):
+ mjm, mjd, m, d = test_util.fixture(
+ xml="""
+
+
+
+
+
+
+
+
+ """,
+ sparse=sparse,
+ )
+
+ qM = np.zeros((mjm.nv, mjm.nv))
+ mujoco.mj_fullM(mjm, qM, mjd.qM)
+
+ d.qLD.zero_()
+ if sparse:
+ d.qLDiagInv.zero_()
+
+ res = wp.zeros((1, mjm.nv), dtype=float)
+ vec = wp.ones((1, mjm.nv), dtype=float)
+
+ mjwarp._src.smooth.factor_solve_i(m, d, d.qM, d.qLD, d.qLDiagInv, res, vec)
+
+ _assert_eq(res.numpy()[0], np.linalg.solve(qM, vec.numpy()[0]), "qM \\ 1")
+
+ def test_tendon_armature(self):
+ mjm, mjd, m, d = test_util.fixture("tendon/armature.xml", keyframe=0)
+
+ # qM
+ d.qM.zero_()
+
+ mjwarp._src.smooth.crb(m, d)
+ mjwarp._src.smooth.tendon_armature(m, d)
+
+ qM = np.zeros((mjm.nv, mjm.nv))
+ mujoco.mj_fullM(mjm, qM, mjd.qM)
+ _assert_eq(d.qM.numpy()[0], qM, "qM")
+
+ # qfrc_bias
+ d.qfrc_bias.zero_()
+
+ mjwarp._src.smooth.rne(m, d)
+ mjwarp._src.smooth.tendon_bias(m, d, d.qfrc_bias)
+ _assert_eq(d.qfrc_bias.numpy()[0], mjd.qfrc_bias, "qfrc_bias")
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py
new file mode 100644
index 00000000..31de24df
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py
@@ -0,0 +1,2578 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+from math import ceil
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src import math
+from mujoco.mjx.third_party.mujoco_warp._src import smooth
+from mujoco.mjx.third_party.mujoco_warp._src import support
+from mujoco.mjx.third_party.mujoco_warp._src import types
+from mujoco.mjx.third_party.mujoco_warp._src.block_cholesky import create_blocked_cholesky_func
+from mujoco.mjx.third_party.mujoco_warp._src.block_cholesky import create_blocked_cholesky_solve_func
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel
+
+
+@wp.func
+def _rescale(nv: int, stat_meaninertia: float, value: float) -> float:
+ return value / (stat_meaninertia * float(wp.max(1, nv)))
+
+
+@wp.func
+def _in_bracket(x: wp.vec3, y: wp.vec3) -> bool:
+ return (x[1] < y[1] and y[1] < 0.0) or (x[1] > y[1] and y[1] > 0.0)
+
+
+@wp.func
+def _eval_pt(quad: wp.vec3, alpha: float) -> wp.vec3:
+ return wp.vec3(
+ alpha * alpha * quad[2] + alpha * quad[1] + quad[0],
+ 2.0 * alpha * quad[2] + quad[1],
+ 2.0 * quad[2],
+ )
+
+
+@wp.func
+def _eval_pt_elliptic(
+ # In:
+ impratio: float,
+ friction: types.vec5,
+ u0: float,
+ uu: float,
+ uv: float,
+ vv: float,
+ jv: float,
+ D: float,
+ quad: wp.vec3,
+ alpha: float,
+) -> wp.vec3:
+ mu = friction[0] / wp.sqrt(impratio)
+ v0 = jv * mu
+ n = u0 + alpha * v0
+ tsqr = uu + alpha * (2.0 * uv + alpha * vv)
+ t = wp.sqrt(tsqr) # tangential force
+
+ bottom_zone = ((tsqr <= 0.0) and (n < 0)) or ((tsqr > 0.0) and ((mu * n + t) <= 0.0))
+ middle_zone = (tsqr > 0) and (n < (mu * t)) and ((mu * n + t) > 0.0)
+
+ # elliptic bottom zone: quadratic cose
+ if bottom_zone:
+ pt = _eval_pt(quad, alpha)
+ else:
+ pt = wp.vec3(0.0)
+
+ # elliptic middle zone
+ if t == 0.0:
+ t += types.MJ_MINVAL
+
+ if tsqr == 0.0:
+ tsqr += types.MJ_MINVAL
+
+ n1 = v0
+ t1 = (uv + alpha * vv) / t
+ t2 = vv / t - (uv + alpha * vv) * t1 / tsqr
+
+ if middle_zone:
+ mu2 = mu * mu
+ dm = D / wp.max(mu2 * (1.0 + mu2), types.MJ_MINVAL)
+ nmt = n - mu * t
+ n1mut1 = n1 - mu * t1
+
+ pt += wp.vec3(
+ 0.5 * dm * nmt * nmt,
+ dm * nmt * n1mut1,
+ dm * (n1mut1 * n1mut1 - nmt * mu * t2),
+ )
+
+ return pt
+
+
+@wp.kernel
+def linesearch_iterative_init_gtol_p0_gauss(
+ # Model:
+ nv: int,
+ opt_tolerance: wp.array(dtype=float),
+ opt_ls_tolerance: wp.array(dtype=float),
+ stat_meaninertia: float,
+ # Data in:
+ efc_search_dot_in: wp.array(dtype=float),
+ efc_quad_gauss_in: wp.array(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_gtol_out: wp.array(dtype=float),
+ efc_p0_out: wp.array(dtype=wp.vec3),
+):
+ worldid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ tolerance = opt_tolerance[worldid]
+ ls_tolerance = opt_ls_tolerance[worldid]
+ snorm = wp.math.sqrt(efc_search_dot_in[worldid])
+ scale = stat_meaninertia * wp.float(wp.max(1, nv))
+ efc_gtol_out[worldid] = tolerance * ls_tolerance * snorm * scale
+
+ quad = efc_quad_gauss_in[worldid]
+ efc_p0_out[worldid] = wp.vec3(quad[0], quad[1], 2.0 * quad[2])
+
+
+@wp.kernel
+def linesearch_iterative_init_p0_elliptic0(
+ # Data in:
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nl_in: wp.array(dtype=int),
+ nefc_in: wp.array(dtype=int),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ efc_quad_in: wp.array2d(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ efc_condim_in: wp.array2d(dtype=int),
+ # Data out:
+ efc_p0_out: wp.array(dtype=wp.vec3),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ active = efc_Jaref_in[worldid, efcid] < 0.0
+
+ nef = ne_in[worldid] + nf_in[worldid]
+ nefl = nef + nl_in[worldid]
+ if efcid < nef:
+ active = True
+ elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1:
+ active = False
+
+ if active:
+ quad = efc_quad_in[worldid, efcid]
+ wp.atomic_add(efc_p0_out, worldid, wp.vec3(quad[0], quad[1], 2.0 * quad[2]))
+
+
+@wp.kernel
+def linesearch_iterative_init_p0_elliptic1(
+ # Model:
+ opt_impratio: wp.array(dtype=float),
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ contact_worldid_in: wp.array(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_jv_in: wp.array2d(dtype=float),
+ efc_quad_in: wp.array2d(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ efc_u_in: wp.array(dtype=types.vec6),
+ efc_uu_in: wp.array(dtype=float),
+ efc_uv_in: wp.array(dtype=float),
+ efc_vv_in: wp.array(dtype=float),
+ # Data out:
+ efc_p0_out: wp.array(dtype=wp.vec3),
+):
+ conid = wp.tid()
+
+ if conid >= ncon_in[0]:
+ return
+
+ worldid = contact_worldid_in[conid]
+ if efc_done_in[worldid]:
+ return
+
+ if contact_dim_in[conid] < 2:
+ return
+
+ efcid = contact_efc_address_in[conid, 0]
+
+ pt = _eval_pt_elliptic(
+ opt_impratio[worldid],
+ contact_friction_in[conid],
+ efc_u_in[conid][0],
+ efc_uu_in[conid],
+ efc_uv_in[conid],
+ efc_vv_in[conid],
+ efc_jv_in[worldid, efcid],
+ efc_D_in[worldid, efcid],
+ efc_quad_in[worldid, efcid],
+ 0.0,
+ )
+
+ wp.atomic_add(efc_p0_out, worldid, pt)
+
+
+@wp.kernel
+def linesearch_iterative_init_p0_pyramidal(
+ # Data in:
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nefc_in: wp.array(dtype=int),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ efc_quad_in: wp.array2d(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_p0_out: wp.array(dtype=wp.vec3),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ if efc_Jaref_in[worldid, efcid] >= 0.0 and efcid >= ne_in[worldid] + nf_in[worldid]:
+ return
+
+ quad = efc_quad_in[worldid, efcid]
+
+ wp.atomic_add(efc_p0_out, worldid, wp.vec3(quad[0], quad[1], 2.0 * quad[2]))
+
+
+@wp.kernel
+def linesearch_iterative_init_lo_gauss(
+ # Data in:
+ efc_quad_gauss_in: wp.array(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ efc_p0_in: wp.array(dtype=wp.vec3),
+ # Data out:
+ efc_lo_out: wp.array(dtype=wp.vec3),
+ efc_lo_alpha_out: wp.array(dtype=float),
+):
+ worldid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ p0 = efc_p0_in[worldid]
+ alpha = -math.safe_div(p0[1], p0[2])
+ efc_lo_out[worldid] = _eval_pt(efc_quad_gauss_in[worldid], alpha)
+ efc_lo_alpha_out[worldid] = alpha
+
+
+@wp.kernel
+def linesearch_iterative_init_lo_elliptic0(
+ # Data in:
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nl_in: wp.array(dtype=int),
+ nefc_in: wp.array(dtype=int),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ efc_jv_in: wp.array2d(dtype=float),
+ efc_quad_in: wp.array2d(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ efc_lo_alpha_in: wp.array(dtype=float),
+ efc_condim_in: wp.array2d(dtype=int),
+ # Data out:
+ efc_lo_out: wp.array(dtype=wp.vec3),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ alpha = efc_lo_alpha_in[worldid]
+
+ active = efc_Jaref_in[worldid, efcid] + alpha * efc_jv_in[worldid, efcid] < 0.0
+
+ nef = ne_in[worldid] + nf_in[worldid]
+ nefl = nef + nl_in[worldid]
+ if efcid < nef:
+ active = True
+ elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1:
+ active = False
+
+ if active:
+ wp.atomic_add(efc_lo_out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha))
+
+
+@wp.kernel
+def linesearch_iterative_init_lo_elliptic1(
+ # Model:
+ opt_impratio: wp.array(dtype=float),
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ contact_worldid_in: wp.array(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_jv_in: wp.array2d(dtype=float),
+ efc_quad_in: wp.array2d(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ efc_lo_alpha_in: wp.array(dtype=float),
+ efc_u_in: wp.array(dtype=types.vec6),
+ efc_uu_in: wp.array(dtype=float),
+ efc_uv_in: wp.array(dtype=float),
+ efc_vv_in: wp.array(dtype=float),
+ # Data out:
+ efc_lo_out: wp.array(dtype=wp.vec3),
+):
+ conid = wp.tid()
+
+ if conid >= ncon_in[0]:
+ return
+
+ worldid = contact_worldid_in[conid]
+ if efc_done_in[worldid]:
+ return
+
+ if contact_dim_in[conid] < 2:
+ return
+
+ efcid = contact_efc_address_in[conid, 0]
+ alpha = efc_lo_alpha_in[worldid]
+ pt = _eval_pt_elliptic(
+ opt_impratio[worldid],
+ contact_friction_in[conid],
+ efc_u_in[conid][0],
+ efc_uu_in[conid],
+ efc_uv_in[conid],
+ efc_vv_in[conid],
+ efc_jv_in[worldid, efcid],
+ efc_D_in[worldid, efcid],
+ efc_quad_in[worldid, efcid],
+ alpha,
+ )
+ wp.atomic_add(efc_lo_out, worldid, pt)
+
+
+@wp.kernel
+def linesearch_iterative_init_lo_pyramidal(
+ # Data in:
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nefc_in: wp.array(dtype=int),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ efc_jv_in: wp.array2d(dtype=float),
+ efc_quad_in: wp.array2d(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ efc_lo_alpha_in: wp.array(dtype=float),
+ # Data out:
+ efc_lo_out: wp.array(dtype=wp.vec3),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ alpha = efc_lo_alpha_in[worldid]
+
+ if efc_Jaref_in[worldid, efcid] + alpha * efc_jv_in[worldid, efcid] < 0.0 or (efcid < ne_in[worldid] + nf_in[worldid]):
+ wp.atomic_add(efc_lo_out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha))
+
+
+@wp.kernel
+def linesearch_iterative_init_bounds(
+ # Data in:
+ efc_done_in: wp.array(dtype=bool),
+ efc_p0_in: wp.array(dtype=wp.vec3),
+ efc_lo_in: wp.array(dtype=wp.vec3),
+ efc_lo_alpha_in: wp.array(dtype=float),
+ # Data out:
+ efc_lo_out: wp.array(dtype=wp.vec3),
+ efc_lo_alpha_out: wp.array(dtype=float),
+ efc_hi_out: wp.array(dtype=wp.vec3),
+ efc_hi_alpha_out: wp.array(dtype=float),
+):
+ worldid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ p0 = efc_p0_in[worldid]
+ lo = efc_lo_in[worldid]
+ lo_alpha = efc_lo_alpha_in[worldid]
+ lo_less = lo[1] < p0[1]
+
+ efc_lo_out[worldid] = wp.where(lo_less, lo, p0)
+ efc_lo_alpha_out[worldid] = wp.where(lo_less, lo_alpha, 0.0)
+ efc_hi_out[worldid] = wp.where(lo_less, p0, lo)
+ efc_hi_alpha_out[worldid] = wp.where(lo_less, 0.0, lo_alpha)
+
+
+@wp.kernel
+def linesearch_iterative_next_alpha_gauss(
+ # Data in:
+ efc_quad_gauss_in: wp.array(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ efc_ls_done_in: wp.array(dtype=bool),
+ efc_lo_in: wp.array(dtype=wp.vec3),
+ efc_lo_alpha_in: wp.array(dtype=float),
+ efc_hi_in: wp.array(dtype=wp.vec3),
+ efc_hi_alpha_in: wp.array(dtype=float),
+ # Data out:
+ efc_lo_next_out: wp.array(dtype=wp.vec3),
+ efc_lo_next_alpha_out: wp.array(dtype=float),
+ efc_hi_next_out: wp.array(dtype=wp.vec3),
+ efc_hi_next_alpha_out: wp.array(dtype=float),
+ efc_mid_out: wp.array(dtype=wp.vec3),
+ efc_mid_alpha_out: wp.array(dtype=float),
+):
+ worldid = wp.tid()
+
+ if efc_ls_done_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ quad = efc_quad_gauss_in[worldid]
+
+ lo = efc_lo_in[worldid]
+ lo_alpha = efc_lo_alpha_in[worldid]
+ lo_next_alpha = lo_alpha - math.safe_div(lo[1], lo[2])
+ efc_lo_next_out[worldid] = _eval_pt(quad, lo_next_alpha)
+ efc_lo_next_alpha_out[worldid] = lo_next_alpha
+
+ hi = efc_hi_in[worldid]
+ hi_alpha = efc_hi_alpha_in[worldid]
+ hi_next_alpha = hi_alpha - math.safe_div(hi[1], hi[2])
+ efc_hi_next_out[worldid] = _eval_pt(quad, hi_next_alpha)
+ efc_hi_next_alpha_out[worldid] = hi_next_alpha
+
+ mid_alpha = 0.5 * (lo_alpha + hi_alpha)
+ efc_mid_out[worldid] = _eval_pt(quad, mid_alpha)
+ efc_mid_alpha_out[worldid] = mid_alpha
+
+
+@wp.kernel
+def linesearch_iterative_next_quad_elliptic0(
+ # Data in:
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nl_in: wp.array(dtype=int),
+ nefc_in: wp.array(dtype=int),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ efc_jv_in: wp.array2d(dtype=float),
+ efc_quad_in: wp.array2d(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ efc_ls_done_in: wp.array(dtype=bool),
+ efc_lo_next_alpha_in: wp.array(dtype=float),
+ efc_hi_next_alpha_in: wp.array(dtype=float),
+ efc_mid_alpha_in: wp.array(dtype=float),
+ efc_condim_in: wp.array2d(dtype=int),
+ # Data out:
+ efc_lo_next_out: wp.array(dtype=wp.vec3),
+ efc_hi_next_out: wp.array(dtype=wp.vec3),
+ efc_mid_out: wp.array(dtype=wp.vec3),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ if efc_ls_done_in[worldid]:
+ return
+
+ nef = ne_in[worldid] + nf_in[worldid]
+ nefl = nef + nl_in[worldid]
+
+ quad = efc_quad_in[worldid, efcid]
+ jaref = efc_Jaref_in[worldid, efcid]
+ jv = efc_jv_in[worldid, efcid]
+
+ alpha = efc_lo_next_alpha_in[worldid]
+
+ active = jaref + alpha * jv < 0.0
+ if efcid < nef:
+ active = True
+ elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1:
+ active = False
+
+ if active:
+ wp.atomic_add(efc_lo_next_out, worldid, _eval_pt(quad, alpha))
+
+ alpha = efc_hi_next_alpha_in[worldid]
+
+ active = jaref + alpha * jv < 0.0
+ if efcid < nef:
+ active = True
+ elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1:
+ active = False
+
+ if active:
+ wp.atomic_add(efc_hi_next_out, worldid, _eval_pt(quad, alpha))
+
+ alpha = efc_mid_alpha_in[worldid]
+
+ active = jaref + alpha * jv < 0.0
+ if efcid < nef:
+ active = True
+ elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1:
+ active = False
+
+ if active:
+ wp.atomic_add(efc_mid_out, worldid, _eval_pt(quad, alpha))
+
+
+@wp.kernel
+def linesearch_iterative_next_quad_elliptic1(
+ # Model:
+ opt_impratio: wp.array(dtype=float),
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ contact_worldid_in: wp.array(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_jv_in: wp.array2d(dtype=float),
+ efc_quad_in: wp.array2d(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ efc_lo_next_alpha_in: wp.array(dtype=float),
+ efc_hi_next_alpha_in: wp.array(dtype=float),
+ efc_mid_alpha_in: wp.array(dtype=float),
+ efc_u_in: wp.array(dtype=types.vec6),
+ efc_uu_in: wp.array(dtype=float),
+ efc_uv_in: wp.array(dtype=float),
+ efc_vv_in: wp.array(dtype=float),
+ # Data out:
+ efc_lo_next_out: wp.array(dtype=wp.vec3),
+ efc_hi_next_out: wp.array(dtype=wp.vec3),
+ efc_mid_out: wp.array(dtype=wp.vec3),
+):
+ conid = wp.tid()
+
+ if conid >= ncon_in[0]:
+ return
+
+ worldid = contact_worldid_in[conid]
+
+ if efc_done_in[worldid]:
+ return
+
+ if contact_dim_in[conid] < 2:
+ return
+
+ efcid = contact_efc_address_in[conid, 0]
+ impratio = opt_impratio[worldid]
+ friction = contact_friction_in[conid]
+ u = efc_u_in[conid][0]
+ uu = efc_uu_in[conid]
+ uv = efc_uv_in[conid]
+ vv = efc_vv_in[conid]
+ jv = efc_jv_in[worldid, efcid]
+ d = efc_D_in[worldid, efcid]
+ quad = efc_quad_in[worldid, efcid]
+
+ alpha = efc_lo_next_alpha_in[worldid]
+ pt = _eval_pt_elliptic(impratio, friction, u, uu, uv, vv, jv, d, quad, alpha)
+ wp.atomic_add(efc_lo_next_out, worldid, pt)
+
+ alpha = efc_hi_next_alpha_in[worldid]
+ pt = _eval_pt_elliptic(impratio, friction, u, uu, uv, vv, jv, d, quad, alpha)
+ wp.atomic_add(efc_hi_next_out, worldid, pt)
+
+ alpha = efc_mid_alpha_in[worldid]
+ pt = _eval_pt_elliptic(impratio, friction, u, uu, uv, vv, jv, d, quad, alpha)
+ wp.atomic_add(efc_mid_out, worldid, pt)
+
+
+@wp.kernel
+def linesearch_iterative_next_quad_pyramidal(
+ # Data in:
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nefc_in: wp.array(dtype=int),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ efc_jv_in: wp.array2d(dtype=float),
+ efc_quad_in: wp.array2d(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ efc_ls_done_in: wp.array(dtype=bool),
+ efc_lo_next_alpha_in: wp.array(dtype=float),
+ efc_hi_next_alpha_in: wp.array(dtype=float),
+ efc_mid_alpha_in: wp.array(dtype=float),
+ # Data out:
+ efc_lo_next_out: wp.array(dtype=wp.vec3),
+ efc_hi_next_out: wp.array(dtype=wp.vec3),
+ efc_mid_out: wp.array(dtype=wp.vec3),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ if efc_ls_done_in[worldid]:
+ return
+
+ nef_active = efcid < ne_in[worldid] + nf_in[worldid]
+
+ quad = efc_quad_in[worldid, efcid]
+ jaref = efc_Jaref_in[worldid, efcid]
+ jv = efc_jv_in[worldid, efcid]
+
+ alpha = efc_lo_next_alpha_in[worldid]
+ if jaref + alpha * jv < 0.0 or nef_active:
+ wp.atomic_add(efc_lo_next_out, worldid, _eval_pt(quad, alpha))
+
+ alpha = efc_hi_next_alpha_in[worldid]
+ if jaref + alpha * jv < 0.0 or nef_active:
+ wp.atomic_add(efc_hi_next_out, worldid, _eval_pt(quad, alpha))
+
+ alpha = efc_mid_alpha_in[worldid]
+ if jaref + alpha * jv < 0.0 or nef_active:
+ wp.atomic_add(efc_mid_out, worldid, _eval_pt(quad, alpha))
+
+
+@wp.kernel
+def linesearch_iterative_swap(
+ # Data in:
+ efc_gtol_in: wp.array(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ efc_ls_done_in: wp.array(dtype=bool),
+ efc_p0_in: wp.array(dtype=wp.vec3),
+ efc_lo_in: wp.array(dtype=wp.vec3),
+ efc_lo_alpha_in: wp.array(dtype=float),
+ efc_hi_in: wp.array(dtype=wp.vec3),
+ efc_hi_alpha_in: wp.array(dtype=float),
+ efc_lo_next_in: wp.array(dtype=wp.vec3),
+ efc_lo_next_alpha_in: wp.array(dtype=float),
+ efc_hi_next_in: wp.array(dtype=wp.vec3),
+ efc_hi_next_alpha_in: wp.array(dtype=float),
+ efc_mid_in: wp.array(dtype=wp.vec3),
+ efc_mid_alpha_in: wp.array(dtype=float),
+ # Data out:
+ efc_alpha_out: wp.array(dtype=float),
+ efc_ls_done_out: wp.array(dtype=bool),
+ efc_lo_out: wp.array(dtype=wp.vec3),
+ efc_lo_alpha_out: wp.array(dtype=float),
+ efc_hi_out: wp.array(dtype=wp.vec3),
+ efc_hi_alpha_out: wp.array(dtype=float),
+):
+ worldid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ if efc_ls_done_in[worldid]:
+ return
+
+ lo = efc_lo_in[worldid]
+ lo_alpha = efc_lo_alpha_in[worldid]
+ hi = efc_hi_in[worldid]
+ hi_alpha = efc_hi_alpha_in[worldid]
+ lo_next = efc_lo_next_in[worldid]
+ lo_next_alpha = efc_lo_next_alpha_in[worldid]
+ hi_next = efc_hi_next_in[worldid]
+ hi_next_alpha = efc_hi_next_alpha_in[worldid]
+ mid = efc_mid_in[worldid]
+ mid_alpha = efc_mid_alpha_in[worldid]
+
+ # swap lo:
+ swap_lo_lo_next = _in_bracket(lo, lo_next)
+ lo = wp.where(swap_lo_lo_next, lo_next, lo)
+ lo_alpha = wp.where(swap_lo_lo_next, lo_next_alpha, lo_alpha)
+ swap_lo_mid = _in_bracket(lo, mid)
+ lo = wp.where(swap_lo_mid, mid, lo)
+ lo_alpha = wp.where(swap_lo_mid, mid_alpha, lo_alpha)
+ swap_lo_hi_next = _in_bracket(lo, hi_next)
+ lo = wp.where(swap_lo_hi_next, hi_next, lo)
+ lo_alpha = wp.where(swap_lo_hi_next, hi_next_alpha, lo_alpha)
+ efc_lo_out[worldid] = lo
+ efc_lo_alpha_out[worldid] = lo_alpha
+ swap_lo = swap_lo_lo_next or swap_lo_mid or swap_lo_hi_next
+
+ # swap hi:
+ swap_hi_hi_next = _in_bracket(hi, hi_next)
+ hi = wp.where(swap_hi_hi_next, hi_next, hi)
+ hi_alpha = wp.where(swap_hi_hi_next, hi_next_alpha, hi_alpha)
+ swap_hi_mid = _in_bracket(hi, mid)
+ hi = wp.where(swap_hi_mid, mid, hi)
+ hi_alpha = wp.where(swap_hi_mid, mid_alpha, hi_alpha)
+ swap_hi_lo_next = _in_bracket(hi, lo_next)
+ hi = wp.where(swap_hi_lo_next, lo_next, hi)
+ hi_alpha = wp.where(swap_hi_lo_next, lo_next_alpha, hi_alpha)
+ efc_hi_out[worldid] = hi
+ efc_hi_alpha_out[worldid] = hi_alpha
+ swap_hi = swap_hi_hi_next or swap_hi_mid or swap_hi_lo_next
+
+ # if we did not adjust the interval, we are done
+ # also done if either low or hi slope is nearly flat
+ gtol = efc_gtol_in[worldid]
+ efc_ls_done_out[worldid] = (not swap_lo and not swap_hi) or (lo[1] < 0 and lo[1] > -gtol) or (hi[1] > 0 and hi[1] < gtol)
+
+ # update alpha if we have an improvement
+ p0 = efc_p0_in[worldid]
+ alpha = 0.0
+ improved = lo[0] < p0[0] or hi[0] < p0[0]
+ lo_better = lo[0] < hi[0]
+ alpha = wp.where(improved and lo_better, lo_alpha, alpha)
+ alpha = wp.where(improved and not lo_better, hi_alpha, alpha)
+ efc_alpha_out[worldid] = alpha
+
+
+def _linesearch_iterative(m: types.Model, d: types.Data):
+ """Iterative linesearch."""
+ d.efc.ls_done.zero_()
+
+ wp.launch(
+ linesearch_iterative_init_gtol_p0_gauss,
+ dim=(d.nworld,),
+ inputs=[
+ m.nv, m.opt.tolerance, m.opt.ls_tolerance, m.stat.meaninertia, d.efc.search_dot,
+ d.efc.quad_gauss, d.efc.done
+ ],
+ outputs=[d.efc.gtol, d.efc.p0]) # fmt: skip
+
+ if m.opt.cone == types.ConeType.ELLIPTIC:
+ wp.launch(
+ linesearch_iterative_init_p0_elliptic0,
+ dim=(d.nworld, d.njmax,),
+ inputs=[
+ d.ne, d.nf, d.nl, d.nefc, d.efc.Jaref, d.efc.quad, d.efc.done,
+ d.efc.condim
+ ],
+ outputs=[d.efc.p0]) # fmt: skip
+ wp.launch(
+ linesearch_iterative_init_p0_elliptic1,
+ dim=(d.nconmax),
+ inputs=[
+ m.opt.impratio, d.ncon, d.contact.friction, d.contact.dim,
+ d.contact.efc_address, d.contact.worldid, d.efc.D, d.efc.jv, d.efc.quad,
+ d.efc.done, d.efc.u, d.efc.uu, d.efc.uv, d.efc.vv
+ ],
+ outputs=[d.efc.p0]) # fmt: skip
+ else:
+ wp.launch(
+ linesearch_iterative_init_p0_pyramidal,
+ dim=(d.nworld, d.njmax,),
+ inputs=[
+ d.ne, d.nf, d.nefc, d.efc.Jaref, d.efc.quad, d.efc.done
+ ], outputs=[d.efc.p0]) # fmt: skip
+
+ wp.launch(
+ linesearch_iterative_init_lo_gauss,
+ dim=(d.nworld,),
+ inputs=[
+ d.efc.quad_gauss, d.efc.done, d.efc.p0
+ ],
+ outputs=[d.efc.lo, d.efc.lo_alpha]) # fmt: skip
+
+ if m.opt.cone == types.ConeType.ELLIPTIC:
+ wp.launch(
+ linesearch_iterative_init_lo_elliptic0,
+ dim=(d.nworld, d.njmax,),
+ inputs=[
+ d.ne, d.nf, d.nl, d.nefc, d.efc.Jaref, d.efc.jv, d.efc.quad,
+ d.efc.done, d.efc.lo_alpha, d.efc.condim
+ ],
+ outputs=[d.efc.lo]) # fmt: skip
+ wp.launch(
+ linesearch_iterative_init_lo_elliptic1,
+ dim=(d.nconmax),
+ inputs=[
+ m.opt.impratio, d.ncon, d.contact.friction, d.contact.dim,
+ d.contact.efc_address, d.contact.worldid, d.efc.D, d.efc.jv, d.efc.quad,
+ d.efc.done, d.efc.lo_alpha, d.efc.u, d.efc.uu, d.efc.uv, d.efc.vv
+ ],
+ outputs=[d.efc.lo]) # fmt: skip
+ else:
+ wp.launch(
+ linesearch_iterative_init_lo_pyramidal,
+ dim=(d.nworld, d.njmax,),
+ inputs=[
+ d.ne, d.nf, d.nefc, d.efc.Jaref, d.efc.jv,
+ d.efc.quad, d.efc.done, d.efc.lo_alpha
+ ],
+ outputs=[d.efc.lo]) # fmt: skip
+
+ # set the lo/hi interval bounds
+
+ wp.launch(
+ linesearch_iterative_init_bounds,
+ dim=(d.nworld,),
+ inputs=[d.efc.done, d.efc.p0, d.efc.lo, d.efc.lo_alpha],
+ outputs=[d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha]) # fmt: skip
+
+ for _ in range(m.opt.ls_iterations):
+ # NOTE: we always launch ls_iterations kernels, but the kernels may early exit if done
+ # is true. this preserves cudagraph requirements (no dynamic kernel launching) at the
+ # expense of extra launches
+ wp.launch(
+ linesearch_iterative_next_alpha_gauss,
+ dim=(d.nworld,),
+ inputs=[
+ d.efc.quad_gauss, d.efc.done, d.efc.ls_done, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha
+ ],
+ outputs=[
+ d.efc.lo_next, d.efc.lo_next_alpha, d.efc.hi_next, d.efc.hi_next_alpha, d.efc.mid, d.efc.mid_alpha
+ ]) # fmt: skip
+
+ if m.opt.cone == types.ConeType.ELLIPTIC:
+ wp.launch(
+ linesearch_iterative_next_quad_elliptic0,
+ dim=(d.nworld, d.njmax,),
+ inputs=[
+ d.ne, d.nf, d.nl, d.nefc, d.efc.Jaref, d.efc.jv, d.efc.quad, d.efc.done, d.efc.ls_done,
+ d.efc.lo_next_alpha, d.efc.hi_next_alpha, d.efc.mid_alpha, d.efc.condim
+ ],
+ outputs=[d.efc.lo_next, d.efc.hi_next, d.efc.mid]) # fmt: skip
+ wp.launch(
+ linesearch_iterative_next_quad_elliptic1,
+ dim=(d.nconmax),
+ inputs=[
+ m.opt.impratio, d.ncon, d.contact.friction, d.contact.dim, d.contact.efc_address, d.contact.worldid, d.efc.D,
+ d.efc.jv, d.efc.quad, d.efc.done, d.efc.lo_next_alpha, d.efc.hi_next_alpha, d.efc.mid_alpha, d.efc.u, d.efc.uu,
+ d.efc.uv, d.efc.vv
+ ],
+ outputs=[d.efc.lo_next, d.efc.hi_next, d.efc.mid]) # fmt: skip
+ else:
+ wp.launch(
+ linesearch_iterative_next_quad_pyramidal,
+ dim=(d.nworld, d.njmax,),
+ inputs=[
+ d.ne, d.nf, d.nefc, d.efc.Jaref, d.efc.jv, d.efc.quad, d.efc.done, d.efc.ls_done,
+ d.efc.lo_next_alpha, d.efc.hi_next_alpha, d.efc.mid_alpha
+ ],
+ outputs=[d.efc.lo_next, d.efc.hi_next, d.efc.mid]) # fmt: skip
+
+ wp.launch(
+ linesearch_iterative_swap,
+ dim=(d.nworld,),
+ inputs=[
+ d.efc.gtol, d.efc.done, d.efc.ls_done, d.efc.p0, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha,
+ d.efc.lo_next, d.efc.lo_next_alpha, d.efc.hi_next, d.efc.hi_next_alpha, d.efc.mid, d.efc.mid_alpha
+ ],
+ outputs=[
+ d.efc.alpha, d.efc.ls_done, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha
+ ]) # fmt: skip
+
+
+@wp.kernel
+def linesearch_parallel_fused(
+ # Model:
+ nlsp: int,
+ # Data in:
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nefc_in: wp.array(dtype=int),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ efc_jv_in: wp.array2d(dtype=float),
+ efc_quad_in: wp.array2d(dtype=wp.vec3),
+ efc_quad_gauss_in: wp.array(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_cost_candidate_out: wp.array2d(dtype=float),
+):
+ worldid, alphaid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ efc_quad_total_candidate = efc_quad_gauss_in[worldid]
+
+ alpha = float(alphaid) / float(nlsp - 1)
+ ne = ne_in[worldid]
+ nf = nf_in[worldid]
+ for efcid in range(nefc_in[worldid]):
+ Jaref = efc_Jaref_in[worldid, efcid]
+ jv = efc_jv_in[worldid, efcid]
+ quad = efc_quad_in[worldid, efcid]
+
+ if (Jaref + alpha * jv) < 0.0 or (efcid < ne + nf):
+ efc_quad_total_candidate += quad
+
+ alpha_sq = alpha * alpha
+ quad_total0 = efc_quad_total_candidate[0]
+ quad_total1 = efc_quad_total_candidate[1]
+ quad_total2 = efc_quad_total_candidate[2]
+
+ efc_cost_candidate_out[worldid, alphaid] = alpha_sq * quad_total2 + alpha * quad_total1 + quad_total0
+
+
+@wp.kernel
+def linesearch_parallel_best_alpha(
+ # Model:
+ nlsp: int,
+ # Data in:
+ efc_done_in: wp.array(dtype=bool),
+ efc_cost_candidate_in: wp.array2d(dtype=float),
+ # Data out:
+ efc_alpha_out: wp.array(dtype=float),
+):
+ worldid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ # TODO(team): investigate alternatives to wp.argmin
+ # TODO(thowell): how did this use to work?
+ bestid = int(0)
+ best_cost = float(wp.inf)
+ for i in range(nlsp):
+ cost = efc_cost_candidate_in[worldid, i]
+ if cost < best_cost:
+ best_cost = cost
+ bestid = i
+
+ efc_alpha_out[worldid] = float(bestid) / float(nlsp - 1)
+
+
+def _linesearch_parallel(m: types.Model, d: types.Data):
+ wp.launch(
+ linesearch_parallel_fused,
+ dim=(d.nworld, m.nlsp),
+ inputs=[
+ m.nlsp,
+ d.ne,
+ d.nf,
+ d.nefc,
+ d.efc.Jaref,
+ d.efc.jv,
+ d.efc.quad,
+ d.efc.quad_gauss,
+ d.efc.done,
+ ],
+ outputs=[d.efc.cost_candidate],
+ )
+
+ wp.launch(
+ linesearch_parallel_best_alpha,
+ dim=(d.nworld),
+ inputs=[m.nlsp, d.efc.done, d.efc.cost_candidate],
+ outputs=[d.efc.alpha],
+ )
+
+
+@wp.kernel
+def linesearch_zero_jv(
+ # Data in:
+ nefc_in: wp.array(dtype=int),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_jv_out: wp.array2d(dtype=float),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ efc_jv_out[worldid, efcid] = 0.0
+
+
+@cache_kernel
+def linesearch_jv_fused(nv: int, dofs_per_thread: int):
+ @nested_kernel
+ def kernel(
+ # Data in:
+ nefc_in: wp.array(dtype=int),
+ efc_J_in: wp.array3d(dtype=float),
+ efc_search_in: wp.array2d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_jv_out: wp.array2d(dtype=float),
+ ):
+ worldid, efcid, dofstart = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ jv_out = float(0.0)
+
+ if wp.static(dofs_per_thread >= nv):
+ for i in range(wp.static(min(dofs_per_thread, nv))):
+ jv_out += efc_J_in[worldid, efcid, i] * efc_search_in[worldid, i]
+ efc_jv_out[worldid, efcid] = jv_out
+
+ else:
+ for i in range(wp.static(dofs_per_thread)):
+ ii = dofstart * wp.static(dofs_per_thread) + i
+ if ii < nv:
+ jv_out += efc_J_in[worldid, efcid, ii] * efc_search_in[worldid, ii]
+ wp.atomic_add(efc_jv_out, worldid, efcid, jv_out)
+
+ return kernel
+
+
+@wp.kernel
+def linesearch_init_quad_gauss(
+ # Model:
+ nv: int,
+ # Data in:
+ qfrc_smooth_in: wp.array2d(dtype=float),
+ efc_Ma_in: wp.array2d(dtype=float),
+ efc_search_in: wp.array2d(dtype=float),
+ efc_gauss_in: wp.array(dtype=float),
+ efc_mv_in: wp.array2d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_quad_gauss_out: wp.array(dtype=wp.vec3),
+):
+ worldid = wp.tid()
+ if efc_done_in[worldid]:
+ return
+
+ quad_gauss_0 = efc_gauss_in[worldid]
+ quad_gauss_1 = float(0.0)
+ quad_gauss_2 = float(0.0)
+ for i in range(nv):
+ search = efc_search_in[worldid, i]
+ quad_gauss_1 += search * (efc_Ma_in[worldid, i] - qfrc_smooth_in[worldid, i])
+ quad_gauss_2 += 0.5 * search * efc_mv_in[worldid, i]
+
+ efc_quad_gauss_out[worldid] = wp.vec3(quad_gauss_0, quad_gauss_1, quad_gauss_2)
+
+
+@wp.kernel
+def linesearch_init_quad(
+ # Data in:
+ nefc_in: wp.array(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_frictionloss_in: wp.array2d(dtype=float),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ efc_jv_in: wp.array2d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # In:
+ disable_floss: bool,
+ # Data out:
+ efc_quad_out: wp.array2d(dtype=wp.vec3),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ Jaref = efc_Jaref_in[worldid, efcid]
+ jv = efc_jv_in[worldid, efcid]
+ efc_D = efc_D_in[worldid, efcid]
+ floss = efc_frictionloss_in[worldid, efcid]
+
+ if floss > 0.0 and not disable_floss:
+ rf = math.safe_div(floss, efc_D)
+ if Jaref <= -rf:
+ efc_quad_out[worldid, efcid] = wp.vec3(floss * (-0.5 * rf - Jaref), -floss * jv, 0.0)
+ return
+ elif Jaref >= rf:
+ efc_quad_out[worldid, efcid] = wp.vec3(floss * (-0.5 * rf + Jaref), floss * jv, 0.0)
+ return
+
+ efc_quad_out[worldid, efcid] = wp.vec3(0.5 * Jaref * Jaref * efc_D, jv * Jaref * efc_D, 0.5 * jv * jv * efc_D)
+
+
+@wp.kernel
+def linesearch_quad_elliptic(
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ contact_worldid_in: wp.array(dtype=int),
+ efc_jv_in: wp.array2d(dtype=float),
+ efc_quad_in: wp.array2d(dtype=wp.vec3),
+ efc_done_in: wp.array(dtype=bool),
+ efc_u_in: wp.array(dtype=types.vec6),
+ # Data out:
+ efc_quad_out: wp.array2d(dtype=wp.vec3),
+ efc_uv_out: wp.array(dtype=float),
+ efc_vv_out: wp.array(dtype=float),
+):
+ conid, dimid = wp.tid()
+ dimid += 1
+
+ if conid >= ncon_in[0]:
+ return
+
+ worldid = contact_worldid_in[conid]
+ if efc_done_in[worldid]:
+ return
+
+ condim = contact_dim_in[conid]
+
+ if condim == 1 or (dimid >= condim):
+ return
+
+ efcid0 = contact_efc_address_in[conid, 0]
+ efcid = contact_efc_address_in[conid, dimid]
+
+ # complete vector quadratic (for bottom zone)
+ wp.atomic_add(efc_quad_out, worldid, efcid0, efc_quad_in[worldid, efcid])
+
+ # rescale to make primal cone circular
+ u = efc_u_in[conid][dimid]
+ v = efc_jv_in[worldid, efcid] * contact_friction_in[conid][dimid - 1]
+ wp.atomic_add(efc_uv_out, conid, u * v)
+ wp.atomic_add(efc_vv_out, conid, v * v)
+
+
+@wp.kernel
+def linesearch_qacc_ma(
+ # Data in:
+ efc_search_in: wp.array2d(dtype=float),
+ efc_mv_in: wp.array2d(dtype=float),
+ efc_alpha_in: wp.array(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ qacc_out: wp.array2d(dtype=float),
+ efc_Ma_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ alpha = efc_alpha_in[worldid]
+ qacc_out[worldid, dofid] += alpha * efc_search_in[worldid, dofid]
+ efc_Ma_out[worldid, dofid] += alpha * efc_mv_in[worldid, dofid]
+
+
+@wp.kernel
+def linesearch_jaref(
+ # Data in:
+ nefc_in: wp.array(dtype=int),
+ efc_jv_in: wp.array2d(dtype=float),
+ efc_alpha_in: wp.array(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_Jaref_out: wp.array2d(dtype=float),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ efc_Jaref_out[worldid, efcid] += efc_alpha_in[worldid] * efc_jv_in[worldid, efcid]
+
+
+@event_scope
+def _linesearch(m: types.Model, d: types.Data):
+ # mv = qM @ search
+ support.mul_m(m, d, d.efc.mv, d.efc.search, d.efc.done)
+
+ # jv = efc_J @ search
+ # TODO(team): is there a better way of doing batched matmuls with dynamic array sizes?
+
+ # if we are only using 1 thread, it makes sense to do more dofs as we can also skip the
+ # init kernel. For more than 1 thread, dofs_per_thread is lower for better load balancing.
+
+ if m.nv > 50:
+ dofs_per_thread = 20
+ else:
+ dofs_per_thread = 50
+
+ threads_per_efc = ceil(m.nv / dofs_per_thread)
+ # we need to clear the jv array if we're doing atomic adds.
+ if threads_per_efc > 1:
+ wp.launch(
+ linesearch_zero_jv,
+ dim=(d.nworld, d.njmax),
+ inputs=[d.nefc, d.efc.done],
+ outputs=[d.efc.jv],
+ )
+
+ wp.launch(
+ linesearch_jv_fused(m.nv, dofs_per_thread),
+ dim=(d.nworld, d.njmax, threads_per_efc),
+ inputs=[d.nefc, d.efc.J, d.efc.search, d.efc.done],
+ outputs=[d.efc.jv],
+ )
+
+ # prepare quadratics
+ # quad_gauss = [gauss, search.T @ Ma - search.T @ qfrc_smooth, 0.5 * search.T @ mv]
+ wp.launch(
+ linesearch_init_quad_gauss,
+ dim=(d.nworld),
+ inputs=[m.nv, d.qfrc_smooth, d.efc.Ma, d.efc.search, d.efc.gauss, d.efc.mv, d.efc.done],
+ outputs=[d.efc.quad_gauss],
+ )
+
+ # quad = [0.5 * Jaref * Jaref * efc_D, jv * Jaref * efc_D, 0.5 * jv * jv * efc_D]
+
+ disable_floss = m.opt.disableflags & types.DisableBit.FRICTIONLOSS
+ wp.launch(
+ linesearch_init_quad,
+ dim=(d.nworld, d.njmax),
+ inputs=[
+ d.nefc,
+ d.efc.D,
+ d.efc.frictionloss,
+ d.efc.Jaref,
+ d.efc.jv,
+ d.efc.done,
+ disable_floss,
+ ],
+ outputs=[d.efc.quad],
+ )
+
+ if m.opt.cone == types.ConeType.ELLIPTIC:
+ d.efc.uv.zero_()
+ d.efc.vv.zero_()
+ wp.launch(
+ linesearch_quad_elliptic,
+ dim=(d.nconmax, m.condim_max - 1),
+ inputs=[
+ d.ncon,
+ d.contact.friction,
+ d.contact.dim,
+ d.contact.efc_address,
+ d.contact.worldid,
+ d.efc.jv,
+ d.efc.quad,
+ d.efc.done,
+ d.efc.u,
+ ],
+ outputs=[d.efc.quad, d.efc.uv, d.efc.vv],
+ )
+
+ if m.opt.ls_parallel:
+ _linesearch_parallel(m, d)
+ else:
+ _linesearch_iterative(m, d)
+
+ wp.launch(
+ linesearch_qacc_ma,
+ dim=(d.nworld, m.nv),
+ inputs=[d.efc.search, d.efc.mv, d.efc.alpha, d.efc.done],
+ outputs=[d.qacc, d.efc.Ma],
+ )
+
+ wp.launch(
+ linesearch_jaref,
+ dim=(d.nworld, d.njmax),
+ inputs=[d.nefc, d.efc.jv, d.efc.alpha, d.efc.done],
+ outputs=[d.efc.Jaref],
+ )
+
+
+@wp.kernel
+def solve_init_efc(
+ # Data out:
+ solver_niter_out: wp.array(dtype=int),
+ efc_search_dot_out: wp.array(dtype=float),
+ efc_cost_out: wp.array(dtype=float),
+ efc_done_out: wp.array(dtype=bool),
+):
+ worldid = wp.tid()
+ efc_cost_out[worldid] = wp.inf
+ solver_niter_out[worldid] = 0
+ efc_done_out[worldid] = False
+ efc_search_dot_out[worldid] = 0.0
+
+
+@wp.kernel
+def solve_init_jaref(
+ # Model:
+ nv: int,
+ # Data in:
+ nefc_in: wp.array(dtype=int),
+ qacc_in: wp.array2d(dtype=float),
+ efc_J_in: wp.array3d(dtype=float),
+ efc_aref_in: wp.array2d(dtype=float),
+ # Data out:
+ efc_Jaref_out: wp.array2d(dtype=float),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ jaref = float(0.0)
+ for i in range(nv):
+ jaref += efc_J_in[worldid, efcid, i] * qacc_in[worldid, i]
+
+ efc_Jaref_out[worldid, efcid] = jaref - efc_aref_in[worldid, efcid]
+
+
+@wp.kernel
+def solve_init_search(
+ # Data in:
+ efc_Mgrad_in: wp.array2d(dtype=float),
+ # Data out:
+ efc_search_out: wp.array2d(dtype=float),
+ efc_search_dot_out: wp.array(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ search = -1.0 * efc_Mgrad_in[worldid, dofid]
+ efc_search_out[worldid, dofid] = search
+ wp.atomic_add(efc_search_dot_out, worldid, search * search)
+
+
+@wp.kernel
+def update_constraint_init_cost(
+ # Data in:
+ efc_cost_in: wp.array(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_gauss_out: wp.array(dtype=float),
+ efc_cost_out: wp.array(dtype=float),
+ efc_prev_cost_out: wp.array(dtype=float),
+):
+ worldid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ efc_gauss_out[worldid] = 0.0
+ efc_prev_cost_out[worldid] = efc_cost_in[worldid]
+ efc_cost_out[worldid] = 0.0
+
+
+@wp.kernel
+def update_constraint_efc_pyramidal(
+ # Data in:
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nefc_in: wp.array(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_frictionloss_in: wp.array2d(dtype=float),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ # In:
+ disable_floss: int,
+ # Data out:
+ efc_force_out: wp.array2d(dtype=float),
+ efc_cost_out: wp.array(dtype=float),
+ efc_active_out: wp.array2d(dtype=bool),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ efc_D = efc_D_in[worldid, efcid]
+ Jaref = efc_Jaref_in[worldid, efcid]
+
+ cost = 0.5 * efc_D * Jaref * Jaref
+ efc_force = -efc_D * Jaref
+
+ ne = ne_in[worldid]
+ nf = nf_in[worldid]
+
+ if efcid < ne:
+ # equality
+ pass
+ elif efcid < ne + nf and not disable_floss:
+ # friction
+ f = efc_frictionloss_in[worldid, efcid]
+ if f > 0.0:
+ rf = math.safe_div(f, efc_D)
+ if Jaref <= -rf:
+ efc_force_out[worldid, efcid] = f
+ efc_active_out[worldid, efcid] = False
+ wp.atomic_add(efc_cost_out, worldid, -0.5 * rf - Jaref)
+ return
+ elif Jaref >= rf:
+ efc_force_out[worldid, efcid] = -f
+ efc_active_out[worldid, efcid] = False
+ wp.atomic_add(efc_cost_out, worldid, -0.5 * rf + Jaref)
+ return
+ else:
+ # limit, contact
+ if Jaref >= 0.0:
+ efc_force_out[worldid, efcid] = 0.0
+ efc_active_out[worldid, efcid] = False
+ return
+
+ efc_force_out[worldid, efcid] = efc_force
+ efc_active_out[worldid, efcid] = True
+ wp.atomic_add(efc_cost_out, worldid, cost)
+
+
+@wp.kernel
+def update_constraint_u_elliptic(
+ # Model:
+ opt_impratio: wp.array(dtype=float),
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ contact_worldid_in: wp.array(dtype=int),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_u_out: wp.array(dtype=types.vec6),
+ efc_uu_out: wp.array(dtype=float),
+ efc_condim_out: wp.array2d(dtype=int),
+):
+ conid, dimid = wp.tid()
+
+ if conid >= ncon_in[0]:
+ return
+
+ worldid = contact_worldid_in[conid]
+ if efc_done_in[worldid]:
+ return
+
+ efcid = contact_efc_address_in[conid, dimid]
+
+ condim = contact_dim_in[conid]
+ efc_condim_out[worldid, efcid] = condim
+
+ if condim == 1:
+ return
+
+ if dimid < condim:
+ if dimid == 0:
+ fri = contact_friction_in[conid][0] / wp.sqrt(opt_impratio[worldid])
+ else:
+ fri = contact_friction_in[conid][dimid - 1]
+ u = efc_Jaref_in[worldid, efcid] * fri
+ efc_u_out[conid][dimid] = u
+ if dimid > 0:
+ wp.atomic_add(efc_uu_out, conid, u * u)
+
+
+@wp.kernel
+def update_constraint_active_elliptic_bottom_zone(
+ # Model:
+ opt_impratio: wp.array(dtype=float),
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ contact_worldid_in: wp.array(dtype=int),
+ efc_done_in: wp.array(dtype=bool),
+ efc_u_in: wp.array(dtype=types.vec6),
+ efc_uu_in: wp.array(dtype=float),
+ # Data out:
+ efc_active_out: wp.array2d(dtype=bool),
+):
+ conid, dimid = wp.tid()
+
+ if conid >= ncon_in[0]:
+ return
+
+ worldid = contact_worldid_in[conid]
+ if efc_done_in[worldid]:
+ return
+
+ condim = contact_dim_in[conid]
+ if condim == 1:
+ return
+
+ mu = contact_friction_in[conid][0] / wp.sqrt(opt_impratio[worldid])
+ n = efc_u_in[conid][0]
+ tt = efc_uu_in[conid]
+ if tt <= 0.0:
+ t = 0.0
+ else:
+ t = wp.sqrt(tt)
+
+ # bottom zone: quadratic
+ bottom_zone = ((t <= 0.0) and (n < 0.0)) or ((t > 0.0) and ((mu * n + t) <= 0.0))
+
+ # update active
+ efcid = contact_efc_address_in[conid, dimid]
+ efc_active_out[worldid, efcid] = bottom_zone
+
+
+@wp.kernel
+def update_constraint_efc_elliptic0(
+ # Data in:
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nl_in: wp.array(dtype=int),
+ nefc_in: wp.array(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_frictionloss_in: wp.array2d(dtype=float),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ efc_active_in: wp.array2d(dtype=bool),
+ efc_done_in: wp.array(dtype=bool),
+ # In:
+ disable_floss: int,
+ # Data out:
+ efc_force_out: wp.array2d(dtype=float),
+ efc_cost_out: wp.array(dtype=float),
+ efc_active_out: wp.array2d(dtype=bool),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ efc_D = efc_D_in[worldid, efcid]
+ Jaref = efc_Jaref_in[worldid, efcid]
+
+ ne = ne_in[worldid]
+ nf = nf_in[worldid]
+ nl = nl_in[worldid]
+
+ if efcid < ne:
+ # equality
+ efc_active_out[worldid, efcid] = True
+ elif efcid < ne + nf and not disable_floss:
+ # friction
+ f = efc_frictionloss_in[worldid, efcid]
+ if f > 0.0:
+ rf = math.safe_div(f, efc_D)
+ if Jaref <= -rf:
+ efc_force_out[worldid, efcid] = f
+ efc_active_out[worldid, efcid] = False
+ wp.atomic_add(efc_cost_out, worldid, -0.5 * rf - Jaref)
+ return
+ elif Jaref >= rf:
+ efc_force_out[worldid, efcid] = -f
+ efc_active_out[worldid, efcid] = False
+ wp.atomic_add(efc_cost_out, worldid, -0.5 * rf + Jaref)
+ return
+ elif efcid < ne + nf + nl:
+ # limits
+ if Jaref < 0.0:
+ efc_active_out[worldid, efcid] = True
+ else:
+ efc_force_out[worldid, efcid] = 0.0
+ efc_active_out[worldid, efcid] = False
+ return
+ else:
+ # contact
+ if not efc_active_in[worldid, efcid]: # calculated by solve_active_elliptic_bottom_zone
+ efc_force_out[worldid, efcid] = 0.0
+ return
+
+ efc_force_out[worldid, efcid] = -efc_D * Jaref
+ wp.atomic_add(efc_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref)
+
+
+@wp.kernel
+def update_constraint_efc_elliptic1(
+ # Model:
+ opt_impratio: wp.array(dtype=float),
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ contact_worldid_in: wp.array(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ efc_u_in: wp.array(dtype=types.vec6),
+ efc_uu_in: wp.array(dtype=float),
+ # Data out:
+ efc_force_out: wp.array2d(dtype=float),
+ efc_cost_out: wp.array(dtype=float),
+):
+ conid, dimid = wp.tid()
+
+ if conid >= ncon_in[0]:
+ return
+
+ worldid = contact_worldid_in[conid]
+ if efc_done_in[worldid]:
+ return
+
+ condim = contact_dim_in[conid]
+
+ if condim == 1 or dimid >= condim:
+ return
+
+ friction = contact_friction_in[conid]
+ efcid = contact_efc_address_in[conid, dimid]
+
+ mu = friction[0] / wp.sqrt(opt_impratio[worldid])
+ n = efc_u_in[conid][0]
+ tt = efc_uu_in[conid]
+ if tt <= 0.0:
+ t = 0.0
+ else:
+ t = wp.sqrt(tt)
+
+ # middle zone: cone
+ middle_zone = (t > 0.0) and (n < (mu * t)) and ((mu * n + t) > 0.0)
+
+ # tangent and friction for middle zone:
+ if middle_zone:
+ efcid0 = contact_efc_address_in[conid, 0]
+ mu2 = mu * mu
+ dm = efc_D_in[worldid, efcid0] / wp.max(mu2 * float(1.0 + mu2), types.MJ_MINVAL)
+
+ nmt = n - mu * t
+
+ force = -dm * nmt * mu
+ if dimid > 0:
+ force_fri = -force / t
+ force_fri *= efc_u_in[conid][dimid] * friction[dimid - 1]
+ efc_force_out[worldid, efcid] += force_fri
+ else:
+ efc_force_out[worldid, efcid] += force
+ worldid = contact_worldid_in[conid]
+ wp.atomic_add(efc_cost_out, worldid, 0.5 * dm * nmt * nmt)
+
+
+@wp.kernel
+def update_constraint_zero_qfrc_constraint(
+ # Data in:
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ qfrc_constraint_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ qfrc_constraint_out[worldid, dofid] = 0.0
+
+
+@wp.kernel
+def update_constraint_init_qfrc_constraint(
+ # Model:
+ nv: int,
+ # Data in:
+ nefc_in: wp.array(dtype=int),
+ efc_J_in: wp.array3d(dtype=float),
+ efc_force_in: wp.array2d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ qfrc_constraint_out: wp.array2d(dtype=float),
+):
+ worldid, efcid = wp.tid()
+
+ if efcid >= nefc_in[worldid]:
+ return
+
+ if efc_done_in[worldid]:
+ return
+
+ force = efc_force_in[worldid, efcid]
+ for i in range(nv):
+ wp.atomic_add(
+ qfrc_constraint_out[worldid],
+ i,
+ efc_J_in[worldid, efcid, i] * force,
+ )
+
+
+@wp.kernel
+def update_constraint_gauss_cost(
+ # Data in:
+ qacc_in: wp.array2d(dtype=float),
+ qfrc_smooth_in: wp.array2d(dtype=float),
+ qacc_smooth_in: wp.array2d(dtype=float),
+ efc_Ma_in: wp.array2d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_gauss_out: wp.array(dtype=float),
+ efc_cost_out: wp.array(dtype=float),
+):
+ worldid, dofid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ gauss_cost = (
+ 0.5
+ * (efc_Ma_in[worldid, dofid] - qfrc_smooth_in[worldid, dofid])
+ * (qacc_in[worldid, dofid] - qacc_smooth_in[worldid, dofid])
+ )
+ wp.atomic_add(efc_gauss_out, worldid, gauss_cost)
+ wp.atomic_add(efc_cost_out, worldid, gauss_cost)
+
+
+def _update_constraint(m: types.Model, d: types.Data):
+ """Update constraint arrays after each solve iteration."""
+
+ disable_floss = m.opt.disableflags & types.DisableBit.FRICTIONLOSS
+
+ wp.launch(
+ update_constraint_init_cost,
+ dim=(d.nworld),
+ inputs=[d.efc.cost, d.efc.done],
+ outputs=[d.efc.gauss, d.efc.cost, d.efc.prev_cost],
+ )
+
+ if m.opt.cone == types.ConeType.PYRAMIDAL:
+ wp.launch(
+ update_constraint_efc_pyramidal,
+ dim=(d.nworld, d.njmax),
+ inputs=[
+ d.ne,
+ d.nf,
+ d.nefc,
+ d.efc.D,
+ d.efc.frictionloss,
+ d.efc.Jaref,
+ disable_floss,
+ ],
+ outputs=[d.efc.force, d.efc.cost, d.efc.active],
+ )
+ elif m.opt.cone == types.ConeType.ELLIPTIC:
+ d.efc.uu.zero_()
+ d.efc.active.zero_()
+ d.efc.condim.fill_(-1)
+ wp.launch(
+ update_constraint_u_elliptic,
+ dim=(d.nconmax, m.condim_max),
+ inputs=[
+ m.opt.impratio,
+ d.ncon,
+ d.contact.friction,
+ d.contact.dim,
+ d.contact.efc_address,
+ d.contact.worldid,
+ d.efc.Jaref,
+ d.efc.done,
+ ],
+ outputs=[d.efc.u, d.efc.uu, d.efc.condim],
+ )
+ wp.launch(
+ update_constraint_active_elliptic_bottom_zone,
+ dim=(d.nconmax, m.condim_max),
+ inputs=[
+ m.opt.impratio,
+ d.ncon,
+ d.contact.friction,
+ d.contact.dim,
+ d.contact.efc_address,
+ d.contact.worldid,
+ d.efc.done,
+ d.efc.u,
+ d.efc.uu,
+ ],
+ outputs=[d.efc.active],
+ )
+ wp.launch(
+ update_constraint_efc_elliptic0,
+ dim=(d.nworld, d.njmax),
+ inputs=[
+ d.ne,
+ d.nf,
+ d.nl,
+ d.nefc,
+ d.efc.D,
+ d.efc.frictionloss,
+ d.efc.Jaref,
+ d.efc.active,
+ d.efc.done,
+ disable_floss,
+ ],
+ outputs=[d.efc.force, d.efc.cost, d.efc.active],
+ )
+ wp.launch(
+ update_constraint_efc_elliptic1,
+ dim=(d.nconmax, m.condim_max),
+ inputs=[
+ m.opt.impratio,
+ d.ncon,
+ d.contact.friction,
+ d.contact.dim,
+ d.contact.efc_address,
+ d.contact.worldid,
+ d.efc.D,
+ d.efc.done,
+ d.efc.u,
+ d.efc.uu,
+ ],
+ outputs=[d.efc.force, d.efc.cost],
+ )
+ else:
+ raise ValueError(f"Unknown cone type: {m.opt.cone}")
+
+ # qfrc_constraint = efc_J.T @ efc_force
+ wp.launch(
+ update_constraint_zero_qfrc_constraint,
+ dim=(d.nworld, m.nv),
+ inputs=[d.efc.done],
+ outputs=[d.qfrc_constraint],
+ )
+
+ wp.launch(
+ update_constraint_init_qfrc_constraint,
+ dim=(d.nworld, d.njmax),
+ inputs=[m.nv, d.nefc, d.efc.J, d.efc.force, d.efc.done],
+ outputs=[d.qfrc_constraint],
+ )
+
+ # gauss = 0.5 * (Ma - qfrc_smooth).T @ (qacc - qacc_smooth)
+ wp.launch(
+ update_constraint_gauss_cost,
+ dim=(d.nworld, m.nv),
+ inputs=[d.qacc, d.qfrc_smooth, d.qacc_smooth, d.efc.Ma, d.efc.done],
+ outputs=[d.efc.gauss, d.efc.cost],
+ )
+
+
+@wp.kernel
+def update_gradient_zero_grad_dot(
+ # Data in:
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_grad_dot_out: wp.array(dtype=float),
+):
+ worldid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ efc_grad_dot_out[worldid] = 0.0
+
+
+@wp.kernel
+def update_gradient_grad(
+ # Data in:
+ qfrc_smooth_in: wp.array2d(dtype=float),
+ qfrc_constraint_in: wp.array2d(dtype=float),
+ efc_Ma_in: wp.array2d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_grad_out: wp.array2d(dtype=float),
+ efc_grad_dot_out: wp.array(dtype=float),
+):
+ worldid, dofid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ grad = efc_Ma_in[worldid, dofid] - qfrc_smooth_in[worldid, dofid] - qfrc_constraint_in[worldid, dofid]
+ efc_grad_out[worldid, dofid] = grad
+ wp.atomic_add(efc_grad_dot_out, worldid, grad * grad)
+
+
+@wp.kernel
+def update_gradient_zero_h_lower(
+ # Model:
+ dof_tri_row: wp.array(dtype=int),
+ dof_tri_col: wp.array(dtype=int),
+ # Data in:
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_h_out: wp.array3d(dtype=float),
+):
+ worldid, elementid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ rowid = dof_tri_row[elementid]
+ colid = dof_tri_col[elementid]
+ efc_h_out[worldid, rowid, colid] = 0.0
+
+
+@wp.kernel
+def update_gradient_set_h_qM_lower_sparse(
+ # Model:
+ qM_fullm_i: wp.array(dtype=int),
+ qM_fullm_j: wp.array(dtype=int),
+ # Data in:
+ qM_in: wp.array3d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_h_out: wp.array3d(dtype=float),
+):
+ worldid, elementid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ i = qM_fullm_i[elementid]
+ j = qM_fullm_j[elementid]
+ efc_h_out[worldid, i, j] = qM_in[worldid, 0, elementid]
+
+
+@wp.kernel
+def update_gradient_copy_lower_triangle(
+ # Model:
+ dof_tri_row: wp.array(dtype=int),
+ dof_tri_col: wp.array(dtype=int),
+ # Data in:
+ qM_in: wp.array3d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_h_out: wp.array3d(dtype=float),
+):
+ worldid, elementid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ rowid = dof_tri_row[elementid]
+ colid = dof_tri_col[elementid]
+ efc_h_out[worldid, rowid, colid] = qM_in[worldid, rowid, colid]
+
+
+@wp.kernel
+def update_gradient_JTDAJ(
+ # Model:
+ dof_tri_row: wp.array(dtype=int),
+ dof_tri_col: wp.array(dtype=int),
+ # Data in:
+ nefc_in: wp.array(dtype=int),
+ efc_J_in: wp.array3d(dtype=float),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_active_in: wp.array2d(dtype=bool),
+ efc_done_in: wp.array(dtype=bool),
+ # In:
+ # Data out:
+ efc_h_out: wp.array3d(dtype=float),
+):
+ worldid, elementid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ nefc = nefc_in[worldid]
+
+ dofi = dof_tri_row[elementid]
+ dofj = dof_tri_col[elementid]
+
+ for efcid in range(nefc):
+ efc_D = efc_D_in[worldid, efcid]
+ active = efc_active_in[worldid, efcid]
+
+ if efc_D == 0.0 or not active:
+ continue
+
+ # TODO(team): sparse efc_J
+ value = efc_J_in[worldid, efcid, dofi] * efc_J_in[worldid, efcid, dofj] * efc_D
+ if value != 0.0:
+ wp.atomic_add(efc_h_out[worldid, dofi], dofj, value)
+
+
+@wp.kernel
+def update_gradient_JTCJ(
+ # Model:
+ opt_impratio: wp.array(dtype=float),
+ dof_tri_row: wp.array(dtype=int),
+ dof_tri_col: wp.array(dtype=int),
+ # Data in:
+ nconmax_in: int,
+ ncon_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ contact_worldid_in: wp.array(dtype=int),
+ efc_J_in: wp.array3d(dtype=float),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ efc_u_in: wp.array(dtype=types.vec6),
+ efc_uu_in: wp.array(dtype=float),
+ # In:
+ nblocks_perblock: int,
+ dim_block: int,
+ # Data out:
+ efc_h_out: wp.array3d(dtype=float),
+):
+ conid_start, elementid = wp.tid()
+
+ dof1id = dof_tri_row[elementid]
+ dof2id = dof_tri_col[elementid]
+
+ for i in range(nblocks_perblock):
+ conid = conid_start + i * dim_block
+
+ if conid >= min(ncon_in[0], nconmax_in):
+ return
+
+ worldid = contact_worldid_in[conid]
+ if efc_done_in[worldid]:
+ continue
+
+ condim = contact_dim_in[conid]
+
+ if condim == 1:
+ continue
+
+ fri = contact_friction_in[conid]
+ mu = fri[0] / wp.sqrt(opt_impratio[worldid])
+ n = efc_u_in[conid][0]
+ tt = efc_uu_in[conid]
+ if tt <= 0.0:
+ t = 0.0
+ else:
+ t = wp.sqrt(tt)
+
+ middle_zone = (t > 0) and (n < (mu * t)) and ((mu * n + t) > 0.0)
+
+ if not middle_zone:
+ continue
+
+ t = wp.max(t, types.MJ_MINVAL)
+ ttt = wp.max(t * t * t, types.MJ_MINVAL)
+
+ mu2 = mu * mu
+ efc0 = contact_efc_address_in[conid, 0]
+ dm = efc_D_in[worldid, efc0] / wp.max(mu2 * (1.0 + mu2), types.MJ_MINVAL)
+
+ if dm == 0.0:
+ continue
+
+ u = efc_u_in[conid]
+
+ efc_h = float(0.0)
+
+ for dim1id in range(condim):
+ if dim1id == 0:
+ efcid1 = efc0
+ else:
+ efcid1 = contact_efc_address_in[conid, dim1id]
+
+ efc_J11 = efc_J_in[worldid, efcid1, dof1id]
+ efc_J12 = efc_J_in[worldid, efcid1, dof2id]
+
+ ui = u[dim1id]
+
+ for dim2id in range(0, dim1id + 1):
+ if dim2id == 0:
+ efcid2 = efc0
+ else:
+ efcid2 = contact_efc_address_in[conid, dim2id]
+
+ efc_J21 = efc_J_in[worldid, efcid2, dof1id]
+ efc_J22 = efc_J_in[worldid, efcid2, dof2id]
+
+ uj = u[dim2id]
+
+ # set first row/column: (1, -mu/t * u)
+ if dim1id == 0 and dim2id == 0:
+ hcone = 1.0
+ elif dim1id == 0:
+ hcone = -mu / t * uj
+ elif dim2id == 0:
+ hcone = -mu / t * ui
+ else:
+ hcone = mu * n / ttt * ui * uj
+
+ # add to diagonal: mu^2 - mu * n / t
+ if dim1id == dim2id:
+ hcone += mu2 - mu * n / t
+
+ # pre and post multiply by diag(mu, friction) scale by dm
+ if dim1id == 0:
+ fri1 = mu
+ else:
+ fri1 = fri[dim1id - 1]
+
+ if dim2id == 0:
+ fri2 = mu
+ else:
+ fri2 = fri[dim2id - 1]
+
+ hcone *= dm * fri1 * fri2
+
+ if hcone != 0.0:
+ efc_h += hcone * efc_J11 * efc_J22
+
+ if dim1id != dim2id:
+ efc_h += hcone * efc_J12 * efc_J21
+
+ worldid = contact_worldid_in[conid]
+ efc_h_out[worldid, dof1id, dof2id] += efc_h
+
+
+@cache_kernel
+def update_gradient_cholesky(tile_size: int):
+ @nested_kernel
+ def kernel(
+ # Data in:
+ efc_grad_in: wp.array2d(dtype=float),
+ efc_h_in: wp.array3d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_Mgrad_out: wp.array2d(dtype=float),
+ ):
+ worldid = wp.tid()
+ TILE_SIZE = wp.static(tile_size)
+
+ if efc_done_in[worldid]:
+ return
+
+ mat_tile = wp.tile_load(efc_h_in[worldid], shape=(TILE_SIZE, TILE_SIZE))
+ fact_tile = wp.tile_cholesky(mat_tile)
+ input_tile = wp.tile_load(efc_grad_in[worldid], shape=TILE_SIZE)
+ output_tile = wp.tile_cholesky_solve(fact_tile, input_tile)
+ wp.tile_store(efc_Mgrad_out[worldid], output_tile)
+
+ return kernel
+
+
+@cache_kernel
+def update_gradient_cholesky_blocked(tile_size: int):
+ @nested_kernel
+ def kernel(
+ # Data in:
+ efc_grad_in: wp.array3d(dtype=float),
+ efc_h_in: wp.array3d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ matrix_size: int,
+ cholesky_L_tmp: wp.array3d(dtype=float),
+ cholesky_y_tmp: wp.array3d(dtype=float),
+ # Data out:
+ efc_Mgrad_out: wp.array3d(dtype=float),
+ ):
+ worldid, tid_block = wp.tid()
+ TILE_SIZE = wp.static(tile_size)
+
+ if efc_done_in[worldid]:
+ return
+
+ wp.static(create_blocked_cholesky_func(TILE_SIZE))(tid_block, efc_h_in[worldid], matrix_size, cholesky_L_tmp[worldid])
+ wp.static(create_blocked_cholesky_solve_func(TILE_SIZE))(
+ tid_block, cholesky_L_tmp[worldid], efc_grad_in[worldid], cholesky_y_tmp[worldid], matrix_size, efc_Mgrad_out[worldid]
+ )
+
+ return kernel
+
+
+def _update_gradient(m: types.Model, d: types.Data):
+ # grad = Ma - qfrc_smooth - qfrc_constraint
+ wp.launch(update_gradient_zero_grad_dot, dim=(d.nworld), inputs=[d.efc.done], outputs=[d.efc.grad_dot])
+
+ wp.launch(
+ update_gradient_grad,
+ dim=(d.nworld, m.nv),
+ inputs=[d.qfrc_smooth, d.qfrc_constraint, d.efc.Ma, d.efc.done],
+ outputs=[d.efc.grad, d.efc.grad_dot],
+ )
+
+ if m.opt.solver == types.SolverType.CG:
+ smooth.solve_m(m, d, d.efc.Mgrad, d.efc.grad)
+ elif m.opt.solver == types.SolverType.NEWTON:
+ # h = qM + (efc_J.T * efc_D * active) @ efc_J
+ if m.opt.is_sparse:
+ wp.launch(
+ update_gradient_zero_h_lower,
+ dim=(d.nworld, m.dof_tri_row.size),
+ inputs=[m.dof_tri_row, m.dof_tri_col, d.efc.done],
+ outputs=[d.efc.h],
+ )
+ wp.launch(
+ update_gradient_set_h_qM_lower_sparse,
+ dim=(d.nworld, m.qM_fullm_i.size),
+ inputs=[m.qM_fullm_i, m.qM_fullm_j, d.qM, d.efc.done],
+ outputs=[d.efc.h],
+ )
+ else:
+ wp.launch(
+ update_gradient_copy_lower_triangle,
+ dim=(d.nworld, m.dof_tri_row.size),
+ inputs=[m.dof_tri_row, m.dof_tri_col, d.qM, d.efc.done],
+ outputs=[d.efc.h],
+ )
+
+ wp.launch(
+ update_gradient_JTDAJ,
+ dim=(d.nworld, m.dof_tri_row.size),
+ inputs=[
+ m.dof_tri_row,
+ m.dof_tri_col,
+ d.nefc,
+ d.efc.J,
+ d.efc.D,
+ d.efc.active,
+ d.efc.done,
+ ],
+ outputs=[d.efc.h],
+ )
+
+ if m.opt.cone == types.ConeType.ELLIPTIC:
+ # Optimization: launching update_gradient_JTCJ with limited number of blocks on a GPU.
+ # Profiling suggests that only a fraction of blocks out of the original
+ # d.njmax blocks do the actual work. It aims to minimize #CTAs with no
+ # effective work. It launches with #blocks that's proportional to the number
+ # of SMs on the GPU. We can now query the SM count:
+ # https://github.com/NVIDIA/warp/commit/f3814e7e5459e5fd13032cf0fddb3daddd510f30
+
+ # make dim_block and nblocks_perblock static for update_gradient_JTCJ to allow
+ # loop unrolling
+ if wp.get_device().is_cuda:
+ sm_count = wp.get_device().sm_count
+
+ # Here we assume one block has 256 threads. We use a factor of 6, which
+ # can be changed in the future to fine-tune the perf. The optimal factor will
+ # depend on the kernel's occupancy, which determines how many blocks can
+ # simultaneously run on the SM. TODO: This factor can be tuned further.
+ dim_block = ceil((sm_count * 6 * 256) / m.dof_tri_row.size)
+ else:
+ # fall back for CPU
+ dim_block = d.nconmax
+
+ nblocks_perblock = int((d.nconmax + dim_block - 1) / dim_block)
+
+ wp.launch(
+ update_gradient_JTCJ,
+ dim=(dim_block, m.dof_tri_row.size),
+ inputs=[
+ m.opt.impratio,
+ m.dof_tri_row,
+ m.dof_tri_col,
+ d.nconmax,
+ d.ncon,
+ d.contact.friction,
+ d.contact.dim,
+ d.contact.efc_address,
+ d.contact.worldid,
+ d.efc.J,
+ d.efc.D,
+ d.efc.done,
+ d.efc.u,
+ d.efc.uu,
+ nblocks_perblock,
+ dim_block,
+ ],
+ outputs=[d.efc.h],
+ )
+
+ # TODO(team): Define good threshold for blocked vs non-blocked cholesky
+ if m.nv < 32:
+ wp.launch_tiled(
+ update_gradient_cholesky(m.nv),
+ dim=(d.nworld,),
+ inputs=[d.efc.grad, d.efc.h, d.efc.done],
+ outputs=[d.efc.Mgrad],
+ block_dim=m.block_dim.update_gradient_cholesky,
+ )
+ else:
+ wp.launch_tiled(
+ update_gradient_cholesky_blocked(32),
+ dim=(d.nworld,),
+ inputs=[
+ d.efc.grad.reshape(shape=(d.nworld, m.nv, 1)),
+ d.efc.h,
+ d.efc.done,
+ m.nv,
+ d.efc.cholesky_L_tmp,
+ d.efc.cholesky_y_tmp.reshape(shape=(d.nworld, m.nv, 1)),
+ ],
+ outputs=[d.efc.Mgrad.reshape(shape=(d.nworld, m.nv, 1))],
+ block_dim=m.block_dim.update_gradient_cholesky,
+ )
+ else:
+ raise ValueError(f"Unknown solver type: {m.opt.solver}")
+
+
+@wp.kernel
+def solve_prev_grad_Mgrad(
+ # Data in:
+ efc_grad_in: wp.array2d(dtype=float),
+ efc_Mgrad_in: wp.array2d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_prev_grad_out: wp.array2d(dtype=float),
+ efc_prev_Mgrad_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ efc_prev_grad_out[worldid, dofid] = efc_grad_in[worldid, dofid]
+ efc_prev_Mgrad_out[worldid, dofid] = efc_Mgrad_in[worldid, dofid]
+
+
+@wp.kernel
+def solve_zero_beta_num_den(
+ # Data in:
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_beta_num_out: wp.array(dtype=float),
+ efc_beta_den_out: wp.array(dtype=float),
+):
+ worldid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ efc_beta_num_out[worldid] = 0.0
+ efc_beta_den_out[worldid] = 0.0
+
+
+@wp.kernel
+def solve_beta_num_den(
+ # Data in:
+ efc_grad_in: wp.array2d(dtype=float),
+ efc_Mgrad_in: wp.array2d(dtype=float),
+ efc_prev_grad_in: wp.array2d(dtype=float),
+ efc_prev_Mgrad_in: wp.array2d(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_beta_num_out: wp.array(dtype=float),
+ efc_beta_den_out: wp.array(dtype=float),
+):
+ worldid, dofid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ prev_Mgrad = efc_prev_Mgrad_in[worldid][dofid]
+ wp.atomic_add(
+ efc_beta_num_out,
+ worldid,
+ efc_grad_in[worldid, dofid] * (efc_Mgrad_in[worldid, dofid] - prev_Mgrad),
+ )
+ wp.atomic_add(efc_beta_den_out, worldid, efc_prev_grad_in[worldid, dofid] * prev_Mgrad)
+
+
+@wp.kernel
+def solve_beta(
+ # Data in:
+ efc_beta_num_in: wp.array(dtype=float),
+ efc_beta_den_in: wp.array(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_beta_out: wp.array(dtype=float),
+):
+ worldid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ efc_beta_out[worldid] = wp.max(0.0, efc_beta_num_in[worldid] / wp.max(types.MJ_MINVAL, efc_beta_den_in[worldid]))
+
+
+@wp.kernel
+def solve_zero_search_dot(
+ # Data in:
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_search_dot_out: wp.array(dtype=float),
+):
+ worldid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ efc_search_dot_out[worldid] = 0.0
+
+
+@wp.kernel
+def solve_search_update(
+ # Model:
+ opt_solver: int,
+ # Data in:
+ efc_Mgrad_in: wp.array2d(dtype=float),
+ efc_search_in: wp.array2d(dtype=float),
+ efc_beta_in: wp.array(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ efc_search_out: wp.array2d(dtype=float),
+ efc_search_dot_out: wp.array(dtype=float),
+):
+ worldid, dofid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ search = -1.0 * efc_Mgrad_in[worldid, dofid]
+
+ if opt_solver == wp.static(types.SolverType.CG.value):
+ search += efc_beta_in[worldid] * efc_search_in[worldid, dofid]
+
+ efc_search_out[worldid, dofid] = search
+ wp.atomic_add(efc_search_dot_out, worldid, search * search)
+
+
+@wp.kernel
+def solve_done(
+ # Model:
+ nv: int,
+ opt_tolerance: wp.array(dtype=float),
+ opt_iterations: int,
+ stat_meaninertia: float,
+ # Data in:
+ efc_grad_dot_in: wp.array(dtype=float),
+ efc_cost_in: wp.array(dtype=float),
+ efc_prev_cost_in: wp.array(dtype=float),
+ efc_done_in: wp.array(dtype=bool),
+ # Data out:
+ solver_niter_out: wp.array(dtype=int),
+ nsolving_out: wp.array(dtype=int),
+ efc_done_out: wp.array(dtype=bool),
+):
+ worldid = wp.tid()
+
+ if efc_done_in[worldid]:
+ return
+
+ solver_niter_out[worldid] += 1
+ tolerance = opt_tolerance[worldid]
+
+ improvement = _rescale(nv, stat_meaninertia, efc_prev_cost_in[worldid] - efc_cost_in[worldid])
+ gradient = _rescale(nv, stat_meaninertia, wp.math.sqrt(efc_grad_dot_in[worldid]))
+ done = (improvement < tolerance) or (gradient < tolerance)
+ if done or solver_niter_out[worldid] == opt_iterations:
+ # if the solver has converged or the maximum number of iterations has been reached then
+ # mark this world as done and remove it from the number of unconverged worlds
+ efc_done_out[worldid] = True
+ wp.atomic_add(nsolving_out, 0, -1)
+
+
+@event_scope
+def _solver_iteration(
+ m: types.Model,
+ d: types.Data,
+):
+ _linesearch(m, d)
+
+ if m.opt.solver == types.SolverType.CG:
+ wp.launch(
+ solve_prev_grad_Mgrad,
+ dim=(d.nworld, m.nv),
+ inputs=[d.efc.grad, d.efc.Mgrad, d.efc.done],
+ outputs=[d.efc.prev_grad, d.efc.prev_Mgrad],
+ )
+
+ _update_constraint(m, d)
+ _update_gradient(m, d)
+
+ # polak-ribiere
+ if m.opt.solver == types.SolverType.CG:
+ wp.launch(
+ solve_zero_beta_num_den,
+ dim=(d.nworld),
+ inputs=[d.efc.done],
+ outputs=[d.efc.beta_num, d.efc.beta_den],
+ )
+
+ wp.launch(
+ solve_beta_num_den,
+ dim=(d.nworld, m.nv),
+ inputs=[d.efc.grad, d.efc.Mgrad, d.efc.prev_grad, d.efc.prev_Mgrad, d.efc.done],
+ outputs=[d.efc.beta_num, d.efc.beta_den],
+ )
+
+ wp.launch(
+ solve_beta,
+ dim=(d.nworld,),
+ inputs=[d.efc.beta_num, d.efc.beta_den, d.efc.done],
+ outputs=[d.efc.beta],
+ )
+
+ wp.launch(solve_zero_search_dot, dim=(d.nworld), inputs=[d.efc.done], outputs=[d.efc.search_dot])
+
+ wp.launch(
+ solve_search_update,
+ dim=(d.nworld, m.nv),
+ inputs=[m.opt.solver, d.efc.Mgrad, d.efc.search, d.efc.beta, d.efc.done],
+ outputs=[d.efc.search, d.efc.search_dot],
+ )
+
+ wp.launch(
+ solve_done,
+ dim=(d.nworld,),
+ inputs=[
+ m.nv,
+ m.opt.tolerance,
+ m.opt.iterations,
+ m.stat.meaninertia,
+ d.efc.grad_dot,
+ d.efc.cost,
+ d.efc.prev_cost,
+ d.efc.done,
+ ],
+ outputs=[d.solver_niter, d.nsolving, d.efc.done],
+ )
+
+
+def create_context(m: types.Model, d: types.Data, grad: bool = True):
+ # initialize some efc arrays
+ wp.launch(
+ solve_init_efc,
+ dim=(d.nworld),
+ outputs=[d.solver_niter, d.efc.search_dot, d.efc.cost, d.efc.done],
+ )
+
+ # jaref = d.efc_J @ d.qacc - d.efc_aref
+ wp.launch(
+ solve_init_jaref,
+ dim=(d.nworld, d.njmax),
+ inputs=[m.nv, d.nefc, d.qacc, d.efc.J, d.efc.aref],
+ outputs=[d.efc.Jaref],
+ )
+
+ # Ma = qM @ qacc
+ support.mul_m(m, d, d.efc.Ma, d.qacc, d.efc.done)
+
+ _update_constraint(m, d)
+
+ if grad:
+ _update_gradient(m, d)
+
+
+def _copy_acc(m: types.Model, d: types.Data):
+ wp.copy(d.qacc, d.qacc_smooth)
+ wp.copy(d.qacc_warmstart, d.qacc_smooth)
+ d.solver_niter.fill_(0)
+
+
+@event_scope
+def solve(m: types.Model, d: types.Data):
+ if d.njmax == 0:
+ _copy_acc(m, d)
+ else:
+ _solve(m, d)
+
+
+def _solve(m: types.Model, d: types.Data):
+ """Finds forces that satisfy constraints."""
+ # warmstart
+ wp.copy(d.qacc, d.qacc_warmstart)
+
+ # create context
+ create_context(m, d, grad=True)
+
+ # search = -Mgrad
+ wp.launch(
+ solve_init_search,
+ dim=(d.nworld, m.nv),
+ inputs=[d.efc.Mgrad],
+ outputs=[d.efc.search, d.efc.search_dot],
+ )
+
+ if m.opt.iterations != 0 and m.opt.graph_conditional:
+ # Note: the iteration kernel (indicated by while_body) is repeatedly launched
+ # as long as condition_iteration is not zero.
+ # condition_iteration is a warp array of size 1 and type int, it counts the number
+ # of worlds that are not converged, it becomes 0 when all worlds are converged.
+ # When the number of iterations reaches m.opt.iterations, solver_niter
+ # becomes zero and all worlds are marked as converged to avoid an infinite loop.
+ # note: we only launch the iteration kernel if everything is not done
+ d.nsolving.fill_(d.nworld)
+ wp.capture_while(
+ d.nsolving,
+ while_body=_solver_iteration,
+ m=m,
+ d=d,
+ )
+ else:
+ # This branch is mostly for when JAX is used as it is currently not compatible
+ # with CUDA graph conditional.
+ # It should be removed when JAX becomes compatible.
+ for i in range(m.opt.iterations):
+ _solver_iteration(m, d)
+
+ wp.copy(d.qacc_warmstart, d.qacc)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py
new file mode 100644
index 00000000..11d043f6
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py
@@ -0,0 +1,344 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 solver functions."""
+
+import mujoco
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+from absl.testing import parameterized
+
+import mujoco_warp as mjwarp
+
+from mujoco.mjx.third_party.mujoco_warp._src import solver
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType
+from mujoco.mjx.third_party.mujoco_warp._src.types import SolverType
+
+# tolerance for difference between MuJoCo and MJWarp solver calculations - mostly
+# due to float precision
+_TOLERANCE = 5e-3
+
+
+def _assert_eq(a, b, name):
+ tol = _TOLERANCE * 10 # avoid test noise
+ err_msg = f"mismatch: {name}"
+ np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+
+
+class SolverTest(parameterized.TestCase):
+ @parameterized.product(cone=tuple(ConeType), solver_=tuple(SolverType))
+ def test_cost(self, cone, solver_):
+ """Tests cost function is correct."""
+ for keyframe in range(3):
+ mjm, mjd, m, d = test_util.fixture(
+ "constraints.xml",
+ keyframe=keyframe,
+ cone=cone,
+ solver=solver_,
+ iterations=0,
+ )
+
+ def cost(qacc):
+ jaref = np.zeros(mjd.nefc, dtype=float)
+ cost = np.zeros(1)
+ mujoco.mj_mulJacVec(mjm, mjd, jaref, qacc)
+ mujoco.mj_constraintUpdate(mjm, mjd, jaref - mjd.efc_aref, cost, 0)
+ return cost
+
+ mj_cost = cost(mjd.qacc)
+
+ # solve with 0 iterations just initializes constraints and costs and then exits
+ mjwarp.solve(m, d)
+
+ mjwarp_cost = d.efc.cost.numpy()[0] - d.efc.gauss.numpy()[0]
+
+ _assert_eq(mjwarp_cost, mj_cost, name="cost")
+
+ @parameterized.parameters(
+ (ConeType.PYRAMIDAL, SolverType.CG, 5, 5, False, False),
+ (ConeType.ELLIPTIC, SolverType.CG, 5, 5, False, False),
+ (ConeType.PYRAMIDAL, SolverType.NEWTON, 2, 4, False, False),
+ (ConeType.ELLIPTIC, SolverType.NEWTON, 2, 5, False, False),
+ (ConeType.PYRAMIDAL, SolverType.NEWTON, 2, 4, True, True),
+ (ConeType.ELLIPTIC, SolverType.NEWTON, 3, 16, True, True),
+ )
+ def test_solve(self, cone, solver_, iterations, ls_iterations, sparse, ls_parallel):
+ """Tests solve."""
+ for keyframe in range(3):
+ mjm, mjd, m, d = test_util.fixture(
+ "constraints.xml",
+ keyframe=keyframe,
+ sparse=sparse,
+ cone=cone,
+ solver=solver_,
+ iterations=iterations,
+ ls_iterations=ls_iterations,
+ ls_parallel=ls_parallel,
+ )
+
+ qacc_warmstart = mjd.qacc_warmstart.copy()
+ mujoco.mj_forward(mjm, mjd)
+ mjd.qacc_warmstart = qacc_warmstart
+
+ d.qacc.zero_()
+ d.qfrc_constraint.zero_()
+ d.efc.force.zero_()
+
+ if solver_ == mujoco.mjtSolver.mjSOL_CG:
+ mjwarp.factor_m(m, d)
+ mjwarp.solve(m, d)
+
+ def cost(qacc):
+ jaref = np.zeros(mjd.nefc, dtype=float)
+ cost = np.zeros(1)
+ mujoco.mj_mulJacVec(mjm, mjd, jaref, qacc)
+ mujoco.mj_constraintUpdate(mjm, mjd, jaref - mjd.efc_aref, cost, 0)
+ return cost
+
+ mj_cost = cost(mjd.qacc)
+ mjwarp_cost = cost(d.qacc.numpy()[0])
+ self.assertLessEqual(mjwarp_cost, mj_cost * 1.025)
+
+ if m.opt.solver == mujoco.mjtSolver.mjSOL_NEWTON:
+ _assert_eq(d.qacc.numpy()[0], mjd.qacc, "qacc")
+ _assert_eq(d.qfrc_constraint.numpy()[0], mjd.qfrc_constraint, "qfrc_constraint")
+ _assert_eq(d.efc.force.numpy()[0, : mjd.nefc], mjd.efc_force, "efc_force")
+
+ @parameterized.parameters(
+ (ConeType.PYRAMIDAL, SolverType.CG, 25, 5),
+ (ConeType.PYRAMIDAL, SolverType.NEWTON, 2, 4),
+ )
+ def test_solve_batch(self, cone, solver_, iterations, ls_iterations):
+ """Tests solve (batch)."""
+
+ mjm0, mjd0, _, _ = test_util.fixture(
+ "humanoid/humanoid.xml",
+ keyframe=0,
+ sparse=False,
+ cone=cone,
+ solver=solver_,
+ iterations=iterations,
+ ls_iterations=ls_iterations,
+ )
+ qacc_warmstart0 = mjd0.qacc_warmstart.copy()
+ mujoco.mj_forward(mjm0, mjd0)
+ mjd0.qacc_warmstart = qacc_warmstart0
+
+ mjm1, mjd1, _, _ = test_util.fixture(
+ "humanoid/humanoid.xml",
+ keyframe=2,
+ sparse=False,
+ cone=cone,
+ solver=solver_,
+ iterations=iterations,
+ ls_iterations=ls_iterations,
+ )
+ qacc_warmstart1 = mjd1.qacc_warmstart.copy()
+ mujoco.mj_forward(mjm1, mjd1)
+ mjd1.qacc_warmstart = qacc_warmstart1
+
+ mjm2, mjd2, _, _ = test_util.fixture(
+ "humanoid/humanoid.xml",
+ keyframe=1,
+ sparse=False,
+ cone=cone,
+ solver=solver_,
+ iterations=iterations,
+ ls_iterations=ls_iterations,
+ )
+ qacc_warmstart2 = mjd2.qacc_warmstart.copy()
+ mujoco.mj_forward(mjm2, mjd2)
+ mjd2.qacc_warmstart = qacc_warmstart2
+
+ nefc_active = mjd0.nefc + mjd1.nefc + mjd2.nefc
+ ne_active = mjd0.ne + mjd1.ne + mjd2.ne
+
+ mjm, mjd, m, _ = test_util.fixture(
+ "humanoid/humanoid.xml",
+ sparse=False,
+ cone=cone,
+ solver=solver_,
+ iterations=iterations,
+ ls_iterations=ls_iterations,
+ )
+ d = mjwarp.put_data(mjm, mjd, nworld=3, njmax=2 * nefc_active)
+
+ d.nefc = wp.array([nefc_active, nefc_active, nefc_active], dtype=wp.int32, ndim=1)
+ d.ne = wp.array([ne_active, ne_active, ne_active], dtype=wp.int32, ndim=1)
+
+ qacc_warmstart = np.vstack(
+ [
+ np.expand_dims(qacc_warmstart0, axis=0),
+ np.expand_dims(qacc_warmstart1, axis=0),
+ np.expand_dims(qacc_warmstart2, axis=0),
+ ]
+ )
+
+ qM0 = np.zeros((mjm0.nv, mjm0.nv))
+ mujoco.mj_fullM(mjm0, qM0, mjd0.qM)
+ qM1 = np.zeros((mjm1.nv, mjm1.nv))
+ mujoco.mj_fullM(mjm1, qM1, mjd1.qM)
+ qM2 = np.zeros((mjm2.nv, mjm2.nv))
+ mujoco.mj_fullM(mjm2, qM2, mjd2.qM)
+
+ qM = np.vstack(
+ [
+ np.expand_dims(qM0, axis=0),
+ np.expand_dims(qM1, axis=0),
+ np.expand_dims(qM2, axis=0),
+ ]
+ )
+ qacc_smooth = np.vstack(
+ [
+ np.expand_dims(mjd0.qacc_smooth, axis=0),
+ np.expand_dims(mjd1.qacc_smooth, axis=0),
+ np.expand_dims(mjd2.qacc_smooth, axis=0),
+ ]
+ )
+ qfrc_smooth = np.vstack(
+ [
+ np.expand_dims(mjd0.qfrc_smooth, axis=0),
+ np.expand_dims(mjd1.qfrc_smooth, axis=0),
+ np.expand_dims(mjd2.qfrc_smooth, axis=0),
+ ]
+ )
+
+ # Reshape the Jacobians
+ efc_J0 = mjd0.efc_J.reshape((mjd0.nefc, mjm0.nv))
+ efc_J1 = mjd1.efc_J.reshape((mjd1.nefc, mjm1.nv))
+ efc_J2 = mjd2.efc_J.reshape((mjd2.nefc, mjm2.nv))
+
+ efc_J_fill = np.zeros((3, d.njmax, m.nv))
+ efc_J_fill[0, : mjd0.nefc, :] = efc_J0
+ efc_J_fill[1, : mjd1.nefc, :] = efc_J1
+ efc_J_fill[2, : mjd2.nefc, :] = efc_J2
+
+ # Similarly for D and aref values
+ efc_D0 = mjd0.efc_D[: mjd0.nefc]
+ efc_D1 = mjd1.efc_D[: mjd1.nefc]
+ efc_D2 = mjd2.efc_D[: mjd2.nefc]
+
+ efc_D_fill = np.zeros((3, d.njmax))
+ efc_D_fill[0, : mjd0.nefc] = efc_D0
+ efc_D_fill[1, : mjd1.nefc] = efc_D1
+ efc_D_fill[2, : mjd2.nefc] = efc_D2
+
+ efc_aref0 = mjd0.efc_aref[: mjd0.nefc]
+ efc_aref1 = mjd1.efc_aref[: mjd1.nefc]
+ efc_aref2 = mjd2.efc_aref[: mjd2.nefc]
+
+ efc_aref_fill = np.zeros((3, d.njmax))
+ efc_aref_fill[0, : mjd0.nefc] = efc_aref0
+ efc_aref_fill[1, : mjd1.nefc] = efc_aref1
+ efc_aref_fill[2, : mjd2.nefc] = efc_aref2
+
+ d.qacc_warmstart = wp.from_numpy(qacc_warmstart, dtype=wp.float32)
+ d.qM = wp.from_numpy(qM, dtype=wp.float32)
+ d.qacc_smooth = wp.from_numpy(qacc_smooth, dtype=wp.float32)
+ d.qfrc_smooth = wp.from_numpy(qfrc_smooth, dtype=wp.float32)
+ d.efc.J = wp.from_numpy(efc_J_fill, dtype=wp.float32)
+ d.efc.D = wp.from_numpy(efc_D_fill, dtype=wp.float32)
+ d.efc.aref = wp.from_numpy(efc_aref_fill, dtype=wp.float32)
+
+ if solver_ == SolverType.CG:
+ m0 = mjwarp.put_model(mjm0)
+ d0 = mjwarp.put_data(mjm0, mjd0)
+ mjwarp.factor_m(m0, d0)
+ qLD0 = d0.qLD.numpy()
+
+ m1 = mjwarp.put_model(mjm1)
+ d1 = mjwarp.put_data(mjm1, mjd1)
+ mjwarp.factor_m(m1, d1)
+ qLD1 = d1.qLD.numpy()
+
+ m2 = mjwarp.put_model(mjm2)
+ d2 = mjwarp.put_data(mjm2, mjd2)
+ mjwarp.factor_m(m2, d2)
+ qLD2 = d2.qLD.numpy()
+
+ qLD = np.vstack([qLD0, qLD1, qLD2])
+ d.qLD = wp.from_numpy(qLD, dtype=wp.float32)
+
+ d.qacc.zero_()
+ d.qfrc_constraint.zero_()
+ d.efc.force.zero_()
+ solver.solve(m, d)
+
+ def cost(m, d, qacc):
+ jaref = np.zeros(d.nefc, dtype=float)
+ cost = np.zeros(1)
+ mujoco.mj_mulJacVec(m, d, jaref, qacc)
+ mujoco.mj_constraintUpdate(m, d, jaref - d.efc_aref, cost, 0)
+ return cost
+
+ mj_cost0 = cost(mjm0, mjd0, mjd0.qacc)
+ mjwarp_cost0 = cost(mjm0, mjd0, d.qacc.numpy()[0])
+ self.assertLessEqual(mjwarp_cost0, mj_cost0 * 1.025)
+
+ mj_cost1 = cost(mjm1, mjd1, mjd1.qacc)
+ mjwarp_cost1 = cost(mjm1, mjd1, d.qacc.numpy()[1])
+ self.assertLessEqual(mjwarp_cost1, mj_cost1 * 1.025)
+
+ mj_cost2 = cost(mjm2, mjd2, mjd2.qacc)
+ mjwarp_cost2 = cost(mjm2, mjd2, d.qacc.numpy()[2])
+ self.assertLessEqual(mjwarp_cost2, mj_cost2 * 1.025)
+
+ if m.opt.solver == SolverType.NEWTON:
+ _assert_eq(d.qacc.numpy()[0], mjd0.qacc, "qacc0")
+ _assert_eq(d.qacc.numpy()[1], mjd1.qacc, "qacc1")
+ _assert_eq(d.qacc.numpy()[2], mjd2.qacc, "qacc2")
+
+ _assert_eq(d.qfrc_constraint.numpy()[0], mjd0.qfrc_constraint, "qfrc_constraint0")
+ _assert_eq(d.qfrc_constraint.numpy()[1], mjd1.qfrc_constraint, "qfrc_constraint1")
+ _assert_eq(d.qfrc_constraint.numpy()[2], mjd2.qfrc_constraint, "qfrc_constraint2")
+
+ # Get world 0 forces - equality constraints at start, inequality constraints later
+ nieq0 = mjd0.nefc - mjd0.ne
+ nieq1 = mjd1.nefc - mjd1.ne
+ nieq2 = mjd2.nefc - mjd2.ne
+ world0_eq_forces = d.efc.force.numpy()[0, : mjd0.ne]
+ world0_ineq_forces = d.efc.force.numpy()[0, ne_active : ne_active + nieq0]
+ world0_forces = np.concatenate([world0_eq_forces, world0_ineq_forces])
+ _assert_eq(world0_forces, mjd0.efc_force, "efc_force0")
+
+ # Get world 1 forces
+ world1_eq_forces = d.efc.force.numpy()[1, : mjd1.ne]
+ world1_ineq_forces = d.efc.force.numpy()[1, ne_active : ne_active + nieq1]
+ world1_forces = np.concatenate([world1_eq_forces, world1_ineq_forces])
+ _assert_eq(world1_forces, mjd1.efc_force, "efc_force1")
+
+ # Get world 2 forces
+ world2_eq_forces = d.efc.force.numpy()[2, : mjd2.ne]
+ world2_ineq_forces = d.efc.force.numpy()[2, ne_active : ne_active + nieq2]
+ world2_forces = np.concatenate([world2_eq_forces, world2_ineq_forces])
+ _assert_eq(world2_forces, mjd2.efc_force, "efc_force2")
+
+ def test_frictionloss(self):
+ """Tests solver with frictionloss."""
+ for keyframe in range(3):
+ _, mjd, m, d = test_util.fixture("constraints.xml", keyframe=keyframe)
+ mjwarp.solve(m, d)
+
+ _assert_eq(d.nf.numpy()[0], mjd.nf, "nf")
+ _assert_eq(d.qacc.numpy()[0], mjd.qacc, "qacc")
+ _assert_eq(d.qfrc_constraint.numpy()[0], mjd.qfrc_constraint, "qfrc_constraint")
+ _assert_eq(d.efc.force.numpy()[0, : mjd.nefc], mjd.efc_force, "efc_force")
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py
new file mode 100644
index 00000000..b59ffadd
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py
@@ -0,0 +1,554 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+from typing import Tuple
+
+import warp as wp
+
+from mujoco.mjx.third_party.mujoco_warp._src.math import motion_cross
+from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import JointType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+from mujoco.mjx.third_party.mujoco_warp._src.types import TileSet
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec5
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
+from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel
+
+wp.set_module_options({"enable_backward": False})
+
+
+@wp.kernel
+def mul_m_sparse_diag(
+ # Model:
+ dof_Madr: wp.array(dtype=int),
+ # Data in:
+ qM_in: wp.array3d(dtype=float),
+ # In:
+ vec: wp.array2d(dtype=float),
+ skip: wp.array(dtype=bool),
+ # Out:
+ res: wp.array2d(dtype=float),
+):
+ """Diagonal update for sparse matmul."""
+ worldid, dofid = wp.tid()
+
+ if skip[worldid]:
+ return
+
+ res[worldid, dofid] = qM_in[worldid, 0, dof_Madr[dofid]] * vec[worldid, dofid]
+
+
+@wp.kernel
+def mul_m_sparse_ij(
+ # Model:
+ qM_mulm_i: wp.array(dtype=int),
+ qM_mulm_j: wp.array(dtype=int),
+ qM_madr_ij: wp.array(dtype=int),
+ # Data in:
+ qM_in: wp.array3d(dtype=float),
+ # In:
+ vec: wp.array2d(dtype=float),
+ skip: wp.array(dtype=bool),
+ # Out:
+ res: wp.array2d(dtype=float),
+):
+ """Off-diagonal update for sparse matmul."""
+ worldid, elementid = wp.tid()
+
+ if skip[worldid]:
+ return
+
+ i = qM_mulm_i[elementid]
+ j = qM_mulm_j[elementid]
+ madr_ij = qM_madr_ij[elementid]
+
+ qM_ij = qM_in[worldid, 0, madr_ij]
+
+ wp.atomic_add(res[worldid], i, qM_ij * vec[worldid, j])
+ wp.atomic_add(res[worldid], j, qM_ij * vec[worldid, i])
+
+
+@cache_kernel
+def mul_m_dense(tile: TileSet):
+ """Returns a matmul kernel for some tile size"""
+
+ @nested_kernel
+ def kernel(
+ # Data In:
+ qM_in: wp.array3d(dtype=float),
+ # In:
+ adr: wp.array(dtype=int),
+ vec: wp.array3d(dtype=float),
+ skip: wp.array(dtype=bool),
+ # Out:
+ res: wp.array3d(dtype=float),
+ ):
+ worldid, nodeid = wp.tid()
+ TILE_SIZE = wp.static(tile.size)
+
+ if skip[worldid]:
+ return
+
+ dofid = adr[nodeid]
+ qM_tile = wp.tile_load(qM_in[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid))
+ vec_tile = wp.tile_load(vec[worldid], shape=(TILE_SIZE, 1), offset=(dofid, 0))
+ res_tile = wp.tile_matmul(qM_tile, vec_tile)
+ wp.tile_store(res[worldid], res_tile, offset=(dofid, 0))
+
+ return kernel
+
+
+@event_scope
+def mul_m(
+ m: Model,
+ d: Data,
+ res: wp.array2d(dtype=float),
+ vec: wp.array2d(dtype=float),
+ skip: wp.array(dtype=bool),
+ M: wp.array3d(dtype=float) = None,
+):
+ """Multiply vectors by inertia matrix.
+
+ Args:
+ m (Model): The model containing kinematic and dynamic information (device).
+ d (Data): The data object containing the current state and output arrays (device).
+ res (wp.array2d(dtype=float)): Result: qM @ vec.
+ vec (wp.array2d(dtype=float)): Input vector to multiply by qM.
+ skip (wp.array(dtype=flooat)): Skip output.
+ M (wp.array3d(dtype=float), optional): Input matrix: M @ vec.
+ """
+
+ if M is None:
+ M = d.qM
+
+ if m.opt.is_sparse:
+ wp.launch(
+ mul_m_sparse_diag,
+ dim=(d.nworld, m.nv),
+ inputs=[m.dof_Madr, M, vec, skip],
+ outputs=[res],
+ )
+
+ wp.launch(
+ mul_m_sparse_ij,
+ dim=(d.nworld, m.qM_madr_ij.size),
+ inputs=[m.qM_mulm_i, m.qM_mulm_j, m.qM_madr_ij, M, vec, skip],
+ outputs=[res],
+ )
+
+ else:
+ for tile in m.qM_tiles:
+ wp.launch_tiled(
+ mul_m_dense(tile),
+ dim=(d.nworld, tile.adr.size),
+ inputs=[
+ M,
+ tile.adr,
+ # note reshape: tile_matmul expects 2d input
+ vec.reshape(vec.shape + (1,)),
+ skip,
+ ],
+ outputs=[res.reshape(res.shape + (1,))],
+ block_dim=m.block_dim.mul_m_dense,
+ )
+
+
+@wp.kernel
+def xfrc_accumulate_kernel(
+ # Model:
+ nbody: int,
+ body_parentid: wp.array(dtype=int),
+ body_rootid: wp.array(dtype=int),
+ dof_bodyid: wp.array(dtype=int),
+ # Data in:
+ xfrc_applied_in: wp.array2d(dtype=wp.spatial_vector),
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ # Out:
+ out: wp.array2d(dtype=float),
+):
+ """Accumulate applied forces on the subtree of a dof."""
+ worldid, dofid = wp.tid()
+ cdof = cdof_in[worldid, dofid]
+ rotational_cdof = wp.spatial_top(cdof)
+ jac = wp.spatial_vector(cdof[3], cdof[4], cdof[5], cdof[0], cdof[1], cdof[2])
+
+ bodyid = dof_bodyid[dofid]
+ accumul = float(0.0)
+
+ for child in range(bodyid, nbody):
+ # any body that is in the subtree of dof_bodyid is part of the jacobian
+ parentid = child
+ while parentid != 0 and parentid != bodyid:
+ parentid = body_parentid[parentid]
+ if parentid == 0:
+ continue # body is not part of the subtree
+ offset = xipos_in[worldid, child] - subtree_com_in[worldid, body_rootid[child]]
+ cross_term = wp.cross(rotational_cdof, offset)
+ xfrc_applied = xfrc_applied_in[worldid, child]
+ accumul += wp.dot(jac, xfrc_applied) + wp.dot(cross_term, wp.spatial_top(xfrc_applied))
+
+ out[worldid, dofid] += accumul
+
+
+@wp.kernel
+def _apply_ft(
+ # Model:
+ nbody: int,
+ body_parentid: wp.array(dtype=int),
+ body_rootid: wp.array(dtype=int),
+ dof_bodyid: wp.array(dtype=int),
+ # Data in:
+ xipos_in: wp.array2d(dtype=wp.vec3),
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ ft_in: wp.array2d(dtype=wp.spatial_vector),
+ flg_add: bool,
+ # Out:
+ qfrc_out: wp.array2d(dtype=float),
+):
+ worldid, dofid = wp.tid()
+ cdof = cdof_in[worldid, dofid]
+ rotational_cdof = wp.vec3(cdof[0], cdof[1], cdof[2])
+ jac = wp.spatial_vector(cdof[3], cdof[4], cdof[5], cdof[0], cdof[1], cdof[2])
+
+ dofbodyid = dof_bodyid[dofid]
+ accumul = float(0.0)
+
+ for bodyid in range(dofbodyid, nbody):
+ # any body that is in the subtree of dofbodyid is part of the jacobian
+ parentid = bodyid
+ while parentid != 0 and parentid != dofbodyid:
+ parentid = body_parentid[parentid]
+ if parentid == 0:
+ continue # body is not part of the subtree
+ offset = xipos_in[worldid, bodyid] - subtree_com_in[worldid, body_rootid[bodyid]]
+ cross_term = wp.cross(rotational_cdof, offset)
+ ft_body = ft_in[worldid, bodyid]
+ accumul += wp.dot(jac, ft_body) + wp.dot(cross_term, wp.spatial_top(ft_body))
+
+ if flg_add:
+ qfrc_out[worldid, dofid] += accumul
+ else:
+ qfrc_out[worldid, dofid] = accumul
+
+
+def apply_ft(m: Model, d: Data, ft: wp.array2d(dtype=wp.spatial_vector), qfrc: wp.array2d(dtype=float), flg_add: bool):
+ wp.launch(
+ kernel=_apply_ft,
+ dim=(d.nworld, m.nv),
+ inputs=[m.nbody, m.body_parentid, m.body_rootid, m.dof_bodyid, d.xipos, d.subtree_com, d.cdof, ft, flg_add],
+ outputs=[qfrc],
+ )
+
+
+@event_scope
+def xfrc_accumulate(m: Model, d: Data, qfrc: wp.array2d(dtype=float)):
+ """
+ Map applied forces at each body via Jacobians to dof space and accumulate.
+
+ Args:
+ m (Model): The model containing kinematic and dynamic information (device).
+ d (Data): The data object containing the current state and output arrays (device).
+ qfrc (wp.array2d(dtype=float)): Total applied force mapped to dof space.
+ """
+ apply_ft(m, d, d.xfrc_applied, qfrc, True)
+
+
+@wp.func
+def all_same(v0: wp.vec3, v1: wp.vec3) -> wp.bool:
+ dx = abs(v0[0] - v1[0])
+ dy = abs(v0[1] - v1[1])
+ dz = abs(v0[2] - v1[2])
+
+ return (
+ (dx <= 1.0e-9 or dx <= max(abs(v0[0]), abs(v1[0])) * 1.0e-9)
+ and (dy <= 1.0e-9 or dy <= max(abs(v0[1]), abs(v1[1])) * 1.0e-9)
+ and (dz <= 1.0e-9 or dz <= max(abs(v0[2]), abs(v1[2])) * 1.0e-9)
+ )
+
+
+@wp.func
+def any_different(v0: wp.vec3, v1: wp.vec3) -> wp.bool:
+ dx = abs(v0[0] - v1[0])
+ dy = abs(v0[1] - v1[1])
+ dz = abs(v0[2] - v1[2])
+
+ return (
+ (dx > 1.0e-9 and dx > max(abs(v0[0]), abs(v1[0])) * 1.0e-9)
+ or (dy > 1.0e-9 and dy > max(abs(v0[1]), abs(v1[1])) * 1.0e-9)
+ or (dz > 1.0e-9 and dz > max(abs(v0[2]), abs(v1[2])) * 1.0e-9)
+ )
+
+
+@wp.func
+def _decode_pyramid(pyramid: wp.array(dtype=float), efc_address: int, mu: vec5, condim: int) -> wp.spatial_vector:
+ """Converts pyramid representation to contact force."""
+ force = wp.spatial_vector()
+
+ if condim == 1:
+ force[0] = pyramid[efc_address]
+ return force
+
+ force[0] = float(0.0)
+ for i in range(condim - 1):
+ dir1 = pyramid[2 * i + efc_address]
+ dir2 = pyramid[2 * i + efc_address + 1]
+ force[0] += dir1 + dir2
+ force[i + 1] = (dir1 - dir2) * mu[i]
+
+ return force
+
+
+@wp.func
+def contact_force_fn(
+ # Model:
+ opt_cone: int,
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ contact_frame_in: wp.array(dtype=wp.mat33),
+ contact_friction_in: wp.array(dtype=vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ efc_force_in: wp.array2d(dtype=float),
+ # In:
+ worldid: int,
+ contact_id: int,
+ to_world_frame: bool,
+) -> wp.spatial_vector:
+ """Extract 6D force:torque for one contact, in contact frame by default."""
+ force = wp.spatial_vector(0.0, 0.0, 0.0, 0.0, 0.0, 0.0)
+ condim = contact_dim_in[contact_id]
+ efc_address = contact_efc_address_in[contact_id, 0]
+
+ if contact_id >= 0 and contact_id <= ncon_in[0] and efc_address >= 0:
+ if opt_cone == int(ConeType.PYRAMIDAL.value):
+ force = _decode_pyramid(
+ efc_force_in[worldid],
+ efc_address,
+ contact_friction_in[contact_id],
+ condim,
+ )
+ else:
+ for i in range(condim):
+ force[i] = efc_force_in[worldid, contact_efc_address_in[contact_id, i]]
+
+ if to_world_frame:
+ # Transform both top and bottom parts of spatial vector by the full contact frame matrix
+ t = wp.spatial_top(force) @ contact_frame_in[contact_id]
+ b = wp.spatial_bottom(force) @ contact_frame_in[contact_id]
+ force = wp.spatial_vector(t, b)
+
+ return force
+
+
+@wp.kernel
+def contact_force_kernel(
+ # Model:
+ opt_cone: int,
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ contact_frame_in: wp.array(dtype=wp.mat33),
+ contact_friction_in: wp.array(dtype=vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ contact_worldid_in: wp.array(dtype=int),
+ efc_force_in: wp.array2d(dtype=float),
+ # In:
+ contact_ids: wp.array(dtype=int),
+ to_world_frame: bool,
+ # Out:
+ out: wp.array(dtype=wp.spatial_vector),
+):
+ tid = wp.tid()
+
+ contactid = contact_ids[tid]
+
+ if contactid >= ncon_in[0]:
+ return
+
+ worldid = contact_worldid_in[contactid]
+
+ out[tid] = contact_force_fn(
+ opt_cone,
+ ncon_in,
+ contact_frame_in,
+ contact_friction_in,
+ contact_dim_in,
+ contact_efc_address_in,
+ efc_force_in,
+ worldid,
+ contactid,
+ to_world_frame,
+ )
+
+
+def contact_force(
+ m: Model,
+ d: Data,
+ contact_ids: wp.array(dtype=int),
+ to_world_frame: bool,
+ force: wp.array(dtype=wp.spatial_vector),
+):
+ """
+ Compute forces for contacts in Data.
+
+ Args:
+ m (Model): The model containing kinematic and dynamic information (device).
+ d (Data): The data object containing the current state and output arrays (device).
+ contact_ids (wp.array(dtype=int)): IDs for each contact.
+ to_world_frame (bool): If True, map force from contact to world frame.
+ force (wp.array(dtype=wp.spatial_vector)): Contact forces.
+ """
+ wp.launch(
+ contact_force_kernel,
+ dim=(contact_ids.size,),
+ inputs=[
+ m.opt.cone,
+ d.ncon,
+ d.contact.frame,
+ d.contact.friction,
+ d.contact.dim,
+ d.contact.efc_address,
+ d.contact.worldid,
+ d.efc.force,
+ contact_ids,
+ to_world_frame,
+ ],
+ outputs=[force],
+ )
+
+
+@wp.func
+def transform_force(force: wp.vec3, torque: wp.vec3, offset: wp.vec3) -> wp.spatial_vector:
+ return wp.spatial_vector(torque - wp.cross(offset, force), force)
+
+
+@wp.func
+def transform_force(frc: wp.spatial_vector, offset: wp.vec3) -> wp.spatial_vector:
+ force = wp.spatial_top(frc)
+ torque = wp.spatial_bottom(frc)
+ return transform_force(force, torque, offset)
+
+
+@wp.func
+def jac(
+ # Model:
+ body_parentid: wp.array(dtype=int),
+ body_rootid: wp.array(dtype=int),
+ dof_bodyid: wp.array(dtype=int),
+ # Data in:
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ point: wp.vec3,
+ bodyid: int,
+ dofid: int,
+ worldid: int,
+) -> Tuple[wp.vec3, wp.vec3]:
+ dof_bodyid_ = dof_bodyid[dofid]
+ in_tree = int(dof_bodyid_ == 0)
+ parentid = bodyid
+ while parentid != 0:
+ if parentid == dof_bodyid_:
+ in_tree = 1
+ break
+ parentid = body_parentid[parentid]
+
+ if not in_tree:
+ return wp.vec3(0.0), wp.vec3(0.0)
+
+ offset = point - wp.vec3(subtree_com_in[worldid, body_rootid[bodyid]])
+
+ cdof = cdof_in[worldid, dofid]
+ cdof_ang = wp.spatial_top(cdof)
+ cdof_lin = wp.spatial_bottom(cdof)
+
+ jacp = cdof_lin + wp.cross(cdof_ang, offset)
+ jacr = cdof_ang
+
+ return jacp, jacr
+
+
+@wp.func
+def jac_dot(
+ # Model:
+ body_parentid: wp.array(dtype=int),
+ body_rootid: wp.array(dtype=int),
+ jnt_type: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ dof_bodyid: wp.array(dtype=int),
+ dof_jntid: wp.array(dtype=int),
+ # Data in:
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
+ cdof_in: wp.array2d(dtype=wp.spatial_vector),
+ cvel_in: wp.array2d(dtype=wp.spatial_vector),
+ cdof_dot_in: wp.array2d(dtype=wp.spatial_vector),
+ # In:
+ point: wp.vec3,
+ bodyid: int,
+ dofid: int,
+ worldid: int,
+) -> Tuple[wp.vec3, wp.vec3]:
+ dof_bodyid_ = dof_bodyid[dofid]
+ in_tree = int(dof_bodyid_ == 0)
+ parentid = bodyid
+ while parentid != 0:
+ if parentid == dof_bodyid_:
+ in_tree = 1
+ break
+ parentid = body_parentid[parentid]
+
+ if not in_tree:
+ return wp.vec3(0.0), wp.vec3(0.0)
+
+ com = subtree_com_in[worldid, body_rootid[bodyid]]
+ offset = point - com
+
+ # transform spatial
+ cvel = cvel_in[worldid, bodyid]
+ pvel_lin = wp.spatial_bottom(cvel) - wp.cross(offset, wp.spatial_top(cvel))
+
+ cdof = cdof_in[worldid, dofid]
+ cdof_dot = cdof_dot_in[worldid, dofid]
+
+ # check for quaternion
+ dofjntid = dof_jntid[dofid]
+ jnttype = jnt_type[dofjntid]
+ jntdofadr = jnt_dofadr[dofjntid]
+
+ if (jnttype == int(JointType.BALL.value)) or ((jnttype == int(JointType.FREE.value)) and dofid >= jntdofadr + 3):
+ # compute cdof_dot for quaternion (use current body cvel)
+ cvel = cvel_in[worldid, dof_bodyid[dofid]]
+ cdof_dot = motion_cross(cvel, cdof)
+
+ cdof_dot_ang = wp.spatial_top(cdof_dot)
+ cdof_dot_lin = wp.spatial_bottom(cdof_dot)
+
+ # construct translational Jacobian (correct for rotation)
+ # first correction term, account for varying cdof
+ correction1 = wp.cross(cdof_dot_ang, offset)
+
+ # second correction term, account for point translational velocity
+ correction2 = wp.cross(wp.spatial_top(cdof), pvel_lin)
+
+ jacp = cdof_dot_lin + correction1 + correction2
+ jacr = cdof_dot_ang
+
+ return jacp, jacr
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support_test.py
new file mode 100644
index 00000000..7c639e3b
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support_test.py
@@ -0,0 +1,123 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 support functions."""
+
+import mujoco
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+from absl.testing import parameterized
+
+import mujoco_warp as mjwarp
+
+from mujoco.mjx.third_party.mujoco_warp._src import test_util
+from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType
+
+# tolerance for difference between MuJoCo and MJWarp support calculations - mostly
+# due to float precision
+_TOLERANCE = 5e-5
+
+
+def _assert_eq(a, b, name):
+ tol = _TOLERANCE * 10 # avoid test noise
+ err_msg = f"mismatch: {name}"
+ np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+
+
+class SupportTest(parameterized.TestCase):
+ @parameterized.parameters(True, False)
+ def test_mul_m(self, sparse):
+ """Tests mul_m."""
+ mjm, mjd, m, d = test_util.fixture("pendula.xml", sparse=sparse)
+
+ mj_res = np.zeros(mjm.nv)
+ mj_vec = np.random.uniform(low=-1.0, high=1.0, size=mjm.nv)
+ mujoco.mj_mulM(mjm, mjd, mj_res, mj_vec)
+
+ res = wp.zeros((1, mjm.nv), dtype=wp.float32)
+ vec = wp.from_numpy(np.expand_dims(mj_vec, axis=0), dtype=wp.float32)
+ skip = wp.zeros((d.nworld), dtype=bool)
+ mjwarp.mul_m(m, d, res, vec, skip)
+
+ _assert_eq(res.numpy()[0], mj_res, f"mul_m ({'sparse' if sparse else 'dense'})")
+
+ def test_xfrc_accumulated(self):
+ """Tests that xfrc_accumulate output matches mj_xfrcAccumulate."""
+ mjm, mjd, m, d = test_util.fixture("pendula.xml")
+ xfrc = np.random.randn(*d.xfrc_applied.numpy().shape)
+ d.xfrc_applied = wp.from_numpy(xfrc, dtype=wp.spatial_vector)
+ qfrc = wp.zeros((1, mjm.nv), dtype=wp.float32)
+ mjwarp.xfrc_accumulate(m, d, qfrc)
+
+ qfrc_expected = np.zeros(m.nv)
+ xfrc = xfrc[0]
+ for i in range(1, m.nbody):
+ mujoco.mj_applyFT(mjm, mjd, xfrc[i, :3], xfrc[i, 3:], mjd.xipos[i], i, qfrc_expected)
+ np.testing.assert_almost_equal(qfrc.numpy()[0], qfrc_expected, 6)
+
+ @parameterized.parameters(
+ (ConeType.PYRAMIDAL, 1, False),
+ (ConeType.PYRAMIDAL, 3, False),
+ (ConeType.PYRAMIDAL, 4, False),
+ (ConeType.PYRAMIDAL, 6, False),
+ (ConeType.PYRAMIDAL, 1, True),
+ (ConeType.PYRAMIDAL, 3, True),
+ (ConeType.PYRAMIDAL, 4, True),
+ (ConeType.PYRAMIDAL, 6, True),
+ (ConeType.ELLIPTIC, 1, False),
+ (ConeType.ELLIPTIC, 3, False),
+ (ConeType.ELLIPTIC, 4, False),
+ (ConeType.ELLIPTIC, 6, False),
+ (ConeType.ELLIPTIC, 1, True),
+ (ConeType.ELLIPTIC, 3, True),
+ (ConeType.ELLIPTIC, 4, True),
+ (ConeType.ELLIPTIC, 6, True),
+ )
+ def test_contact_force(self, cone, condim, to_world_frame):
+ _CONTACT = f"""
+
+
+
+
+
+
+
+
+
+
+
+
+ """
+ mjm, mjd, m, d = test_util.fixture(xml=_CONTACT, cone=cone, keyframe=0)
+
+ mj_force = np.zeros(6, dtype=float)
+ mujoco.mj_contactForce(mjm, mjd, 0, mj_force)
+
+ contact_ids = wp.zeros(1, dtype=int)
+ force = wp.zeros(1, dtype=wp.spatial_vector)
+
+ mjwarp.contact_force(m, d, contact_ids, to_world_frame, force)
+
+ if to_world_frame:
+ frame = mjd.contact.frame[0].reshape((3, 3))
+ mj_force = np.concatenate([frame.T @ mj_force[:3], frame.T @ mj_force[3:]])
+
+ _assert_eq(force.numpy()[0], mj_force, "contact force")
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/test_util.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/test_util.py
new file mode 100644
index 00000000..4d9ae7b1
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/test_util.py
@@ -0,0 +1,267 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+"""Utilities for testing."""
+
+import time
+from typing import Callable, Optional, Tuple
+
+import mujoco
+import numpy as np
+import warp as wp
+from etils import epath
+
+from mujoco.mjx.third_party.mujoco_warp._src import io
+from mujoco.mjx.third_party.mujoco_warp._src import warp_util
+from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Data
+from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit
+from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType
+from mujoco.mjx.third_party.mujoco_warp._src.types import Model
+from mujoco.mjx.third_party.mujoco_warp._src.types import SolverType
+from mujoco.mjx.third_party.mujoco_warp._src.util_misc import halton
+
+
+def fixture(
+ fname: Optional[str] = None,
+ xml: Optional[str] = None,
+ keyframe: int = -1,
+ actuation: bool = True,
+ contact: bool = True,
+ constraint: bool = True,
+ equality: bool = True,
+ passive: bool = True,
+ gravity: bool = True,
+ clampctrl: bool = True,
+ filterparent: bool = True,
+ qpos0: bool = False,
+ kick: bool = False,
+ energy: bool = False,
+ eulerdamp: Optional[bool] = None,
+ cone: Optional[ConeType] = None,
+ integrator: Optional[IntegratorType] = None,
+ solver: Optional[SolverType] = None,
+ iterations: Optional[int] = None,
+ ls_iterations: Optional[int] = None,
+ ls_parallel: Optional[bool] = None,
+ sparse: Optional[bool] = None,
+ disableflags: Optional[int] = None,
+ enableflags: Optional[int] = None,
+ applied: bool = False,
+ nstep: int = 3,
+ seed: int = 42,
+ nworld: int = None,
+ nconmax: int = None,
+ njmax: int = None,
+):
+ np.random.seed(seed)
+ if fname is not None:
+ path = epath.resource_path("mjx") / "third_party/mujoco_warp" / "test_data" / fname
+ mjm = mujoco.MjModel.from_xml_path(path.as_posix())
+ elif xml is not None:
+ mjm = mujoco.MjModel.from_xml_string(xml)
+ else:
+ raise ValueError("either fname or xml must be provided")
+
+ if not actuation:
+ mjm.opt.disableflags |= DisableBit.ACTUATION
+ if not contact:
+ mjm.opt.disableflags |= DisableBit.CONTACT
+ if not constraint:
+ mjm.opt.disableflags |= DisableBit.CONSTRAINT
+ if not equality:
+ mjm.opt.disableflags |= DisableBit.EQUALITY
+ if not passive:
+ mjm.opt.disableflags |= DisableBit.PASSIVE
+ if not gravity:
+ mjm.opt.disableflags |= DisableBit.GRAVITY
+ if not clampctrl:
+ mjm.opt.disableflags |= DisableBit.CLAMPCTRL
+ if not eulerdamp:
+ mjm.opt.disableflags |= DisableBit.EULERDAMP
+ if not filterparent:
+ mjm.opt.disableflags |= DisableBit.FILTERPARENT
+
+ if energy:
+ mjm.opt.enableflags |= EnableBit.ENERGY
+
+ if cone is not None:
+ mjm.opt.cone = cone
+ if integrator is not None:
+ mjm.opt.integrator = integrator
+ if disableflags is not None:
+ mjm.opt.disableflags |= disableflags
+ if enableflags is not None:
+ mjm.opt.enableflags |= enableflags
+ if solver is not None:
+ mjm.opt.solver = solver
+ if iterations is not None:
+ mjm.opt.iterations = iterations
+ if ls_iterations is not None:
+ mjm.opt.ls_iterations = ls_iterations
+ if sparse is not None:
+ if sparse:
+ mjm.opt.jacobian = mujoco.mjtJacobian.mjJAC_SPARSE
+ else:
+ mjm.opt.jacobian = mujoco.mjtJacobian.mjJAC_DENSE
+
+ mjd = mujoco.MjData(mjm)
+ if keyframe > -1:
+ mujoco.mj_resetDataKeyframe(mjm, mjd, keyframe)
+ elif qpos0:
+ mjd.qpos[:] = mjm.qpos0
+ else:
+ # set random qpos, underlying code should gracefully handle un-normalized quats
+ mjd.qpos[:] = np.random.random(mjm.nq)
+
+ if kick:
+ # give the system a little kick to ensure we have non-identity rotations
+ mjd.qvel = np.random.uniform(-0.01, 0.01, mjm.nv)
+ mjd.ctrl = np.random.uniform(-0.1, 0.1, size=mjm.nu)
+ if applied:
+ mjd.qfrc_applied = np.random.uniform(-0.1, 0.1, size=mjm.nv)
+ mjd.xfrc_applied = np.random.uniform(-0.1, 0.1, size=mjd.xfrc_applied.shape)
+ if kick or applied:
+ mujoco.mj_step(mjm, mjd, nstep) # let dynamics get state significantly non-zero
+
+ if mjm.nmocap:
+ mjd.mocap_pos = np.random.random(mjd.mocap_pos.shape)
+ mocap_quat = np.random.random(mjd.mocap_quat.shape)
+ mjd.mocap_quat = mocap_quat
+
+ mujoco.mj_forward(mjm, mjd)
+ m = io.put_model(mjm)
+ if ls_parallel is not None:
+ m.opt.ls_parallel = ls_parallel
+
+ d = io.put_data(mjm, mjd, nworld=nworld, nconmax=nconmax, njmax=njmax)
+ return mjm, mjd, m, d
+
+
+def _sum(stack1, stack2):
+ ret = {}
+ for k in stack1:
+ times1, sub_stack1 = stack1[k]
+ times2, sub_stack2 = stack2[k]
+ times = [t1 + t2 for t1, t2 in zip(times1, times2)]
+ ret[k] = (times, _sum(sub_stack1, sub_stack2))
+ return ret
+
+
+@wp.kernel
+def ctrl_noise(
+ # Model:
+ actuator_ctrllimited: wp.array(dtype=bool),
+ actuator_ctrlrange: wp.array2d(dtype=wp.vec2),
+ # In:
+ step: int,
+ ctrlnoise: float,
+ # Data out:
+ ctrl_out: wp.array2d(dtype=float),
+):
+ worldid, actid = wp.tid()
+
+ center = 0.0
+ radius = 1.0
+ ctrlrange = actuator_ctrlrange[0, actid]
+ if actuator_ctrllimited[actid]:
+ center = (ctrlrange[1] + ctrlrange[0]) / 2.0
+ radius = (ctrlrange[1] - ctrlrange[0]) / 2.0
+ radius *= ctrlnoise
+ noise = 2.0 * halton((step + 1) * (worldid + 1), actid + 2) - 1.0
+ ctrl_out[worldid, actid] = center + radius * noise
+
+
+def benchmark(
+ fn: Callable[[Model, Data], None],
+ m: Model,
+ d: Data,
+ nstep: int,
+ event_trace: bool = False,
+ measure_alloc: bool = False,
+ measure_solver_niter: bool = False,
+) -> Tuple[float, float, dict, list, list, list]:
+ """Benchmark a function of Model and Data.
+
+ Args:
+ fn (Callable[[Model, Data], None]): Function to benchmark.
+ m (Model): The model containing kinematic and dynamic information (device).
+ d (Data): The data object containing the current state and output information (device).
+ nstep (int): Number of timesteps.
+ event_trace (bool, optional): If True, time routines decorated with @event_scope.
+ Default is False.
+ measure_alloc (bool, optional): If True, record number of contacts and constraints.
+ Default is False.
+ measure_solver_niter (bool, False): If True, record the number of solver iterations.
+ Default is False.
+ Returns:
+ float: Time to JIT fn.
+ float: Total time to run the benchmark.
+ dict: Trace.
+ list: Number of contacts.
+ list: Number of constraints.
+ list: Number of solver iterations.
+ """
+ jit_beg = time.perf_counter()
+
+ fn(m, d)
+
+ jit_end = time.perf_counter()
+ jit_duration = jit_end - jit_beg
+ wp.synchronize()
+
+ trace = {}
+ ncon, nefc, solver_niter = [], [], []
+
+ with warp_util.EventTracer(enabled=event_trace) as tracer:
+ # capture the whole function as a CUDA graph
+ with wp.ScopedCapture() as capture:
+ fn(m, d)
+ graph = capture.graph
+
+ time_vec = np.zeros(nstep)
+ for i in range(nstep):
+ with wp.ScopedStream(wp.get_stream()):
+ wp.launch(
+ ctrl_noise,
+ dim=(d.nworld, m.nu),
+ inputs=[
+ m.actuator_ctrllimited, m.actuator_ctrlrange, i, 0.01
+ ],
+ outputs=[d.ctrl]) # fmt: skip
+
+ run_beg = time.perf_counter()
+ wp.capture_launch(graph)
+ wp.synchronize()
+
+ run_end = time.perf_counter()
+ time_vec[i] = run_end - run_beg
+ if trace:
+ trace = _sum(trace, tracer.trace())
+ else:
+ trace = tracer.trace()
+ if measure_alloc or measure_solver_niter:
+ wp.synchronize()
+ if measure_alloc:
+ ncon.append(d.ncon.numpy()[0])
+ nefc.append(np.sum(d.nefc.numpy()))
+ if measure_solver_niter:
+ solver_niter.append(d.solver_niter.numpy())
+
+ wp.synchronize()
+ run_duration = np.sum(time_vec)
+
+ return jit_duration, run_duration, trace, ncon, nefc, solver_niter
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py
new file mode 100644
index 00000000..f1e08dca
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py
@@ -0,0 +1,1667 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+import dataclasses
+import enum
+
+import mujoco
+import warp as wp
+
+MJ_MINVAL = mujoco.mjMINVAL
+MJ_MAXVAL = mujoco.mjMAXVAL
+MJ_MINIMP = mujoco.mjMINIMP # minimum constraint impedance
+MJ_MAXIMP = mujoco.mjMAXIMP # maximum constraint impedance
+MJ_MAXCONPAIR = mujoco.mjMAXCONPAIR
+MJ_MINMU = mujoco.mjMINMU # minimum friction
+
+
+# TODO(team): add check that all wp.launch_tiled 'block_dim' settings are configurable
+@dataclasses.dataclass
+class BlockDim:
+ """
+ Block dimension 'block_dim' settings for wp.launch_tiled.
+
+ TODO(team): experimental and may be removed
+ """
+
+ # collision_driver
+ segmented_sort: int = 128
+ # derivative
+ qderiv_actuator_passive_actuation: int = 64
+ qderiv_actuator_passive_no_actuation: int = 256
+ # forward
+ euler_dense: int = 256
+ actuator_velocity: int = 32
+ tendon_velocity: int = 256
+ qfrc_actuator: int = 256
+ # ray
+ ray: int = 64
+ # sensor
+ energy_vel_kinetic: int = 256
+ # smooth
+ cholesky_factorize: int = 256
+ cholesky_solve: int = 256
+ cholesky_factorize_solve: int = 256
+ # solver
+ update_gradient_cholesky: int = 256
+ # support
+ mul_m_dense: int = 256
+
+
+class BroadphaseFilter(enum.IntFlag):
+ """Bitmask specifying which collision functions to run during broadphase.
+
+ Attributes:
+ PLANE: collision between bounding sphere and plane.
+ SPHERE: collision between bounding spheres.
+ AABB: collision between axis-aligned bounding boxes.
+ OBB: collision between oriented bounding boxes.
+ """
+
+ PLANE = 1
+ SPHERE = 2
+ AABB = 4
+ OBB = 8
+
+
+class CamLightType(enum.IntEnum):
+ """Type of camera light.
+
+ Attributes:
+ FIXED: pos and rot fixed in body
+ TRACK: pos tracks body, rot fixed in global
+ TRACKCOM: pos tracks subtree com, rot fixed in body
+ TARGETBODY: pos fixed in body, rot tracks target body
+ TARGETBODYCOM: pos fixed in body, rot tracks target subtree com
+ """
+
+ FIXED = mujoco.mjtCamLight.mjCAMLIGHT_FIXED
+ TRACK = mujoco.mjtCamLight.mjCAMLIGHT_TRACK
+ TRACKCOM = mujoco.mjtCamLight.mjCAMLIGHT_TRACKCOM
+ TARGETBODY = mujoco.mjtCamLight.mjCAMLIGHT_TARGETBODY
+ TARGETBODYCOM = mujoco.mjtCamLight.mjCAMLIGHT_TARGETBODYCOM
+
+
+class DataType(enum.IntFlag):
+ """Sensor data types.
+
+ Attributes:
+ REAL: real values, no constraints
+ POSITIVE: positive values, 0 or negative: inactive
+ """
+
+ REAL = mujoco.mjtDataType.mjDATATYPE_REAL
+ POSITIVE = mujoco.mjtDataType.mjDATATYPE_POSITIVE
+ # unsupported: AXIS, QUATERNION
+
+
+class DisableBit(enum.IntFlag):
+ """Disable default feature bitflags.
+
+ Attributes:
+ CONSTRAINT: entire constraint solver
+ EQUALITY: equality constraints
+ FRICTIONLOSS: joint and tendon frictionloss constraints
+ LIMIT: joint and tendon limit constraints
+ CONTACT: contact constraints
+ PASSIVE: passive forces
+ GRAVITY: gravitational forces
+ CLAMPCTRL: clamp control to specified range
+ ACTUATION: apply actuation forces
+ REFSAFE: integrator safety: make ref[0]>=2*timestep
+ EULERDAMP: implicit damping for Euler integration
+ FILTERPARENT: disable collisions between parent and child bodies
+ SENSOR: sensors
+ """
+
+ CONSTRAINT = mujoco.mjtDisableBit.mjDSBL_CONSTRAINT
+ EQUALITY = mujoco.mjtDisableBit.mjDSBL_EQUALITY
+ FRICTIONLOSS = mujoco.mjtDisableBit.mjDSBL_FRICTIONLOSS
+ LIMIT = mujoco.mjtDisableBit.mjDSBL_LIMIT
+ CONTACT = mujoco.mjtDisableBit.mjDSBL_CONTACT
+ PASSIVE = mujoco.mjtDisableBit.mjDSBL_PASSIVE
+ GRAVITY = mujoco.mjtDisableBit.mjDSBL_GRAVITY
+ CLAMPCTRL = mujoco.mjtDisableBit.mjDSBL_CLAMPCTRL
+ ACTUATION = mujoco.mjtDisableBit.mjDSBL_ACTUATION
+ REFSAFE = mujoco.mjtDisableBit.mjDSBL_REFSAFE
+ EULERDAMP = mujoco.mjtDisableBit.mjDSBL_EULERDAMP
+ FILTERPARENT = mujoco.mjtDisableBit.mjDSBL_FILTERPARENT
+ SENSOR = mujoco.mjtDisableBit.mjDSBL_SENSOR
+ # unsupported: MIDPHASE, WARMSTART
+
+
+class EnableBit(enum.IntFlag):
+ """Enable optional feature bitflags.
+
+ Attributes:
+ ENERGY: energy computation
+ INVDISCRETE: discrete-time inverse dynamics
+ """
+
+ ENERGY = mujoco.mjtEnableBit.mjENBL_ENERGY
+ INVDISCRETE = mujoco.mjtEnableBit.mjENBL_INVDISCRETE
+ # unsupported: OVERRIDE, FWDINV, MULTICCD, ISLAND
+
+
+class TrnType(enum.IntEnum):
+ """Type of actuator transmission.
+
+ Attributes:
+ JOINT: force on joint
+ JOINTINPARENT: force on joint, expressed in parent frame
+ SLIDERCRANK: force via slider-crank linkage
+ TENDON: force on tendon
+ BODY: adhesion force on body's geoms
+ SITE: force on site
+ """
+
+ JOINT = mujoco.mjtTrn.mjTRN_JOINT
+ JOINTINPARENT = mujoco.mjtTrn.mjTRN_JOINTINPARENT
+ SLIDERCRANK = mujoco.mjtTrn.mjTRN_SLIDERCRANK
+ TENDON = mujoco.mjtTrn.mjTRN_TENDON
+ BODY = mujoco.mjtTrn.mjTRN_BODY
+ SITE = mujoco.mjtTrn.mjTRN_SITE
+
+
+class DynType(enum.IntEnum):
+ """Type of actuator dynamics.
+
+ Attributes:
+ NONE: no internal dynamics; ctrl specifies force
+ INTEGRATOR: integrator: da/dt = u
+ FILTER: linear filter: da/dt = (u-a) / tau
+ FILTEREXACT: linear filter: da/dt = (u-a) / tau, with exact integration
+ MUSCLE: piece-wise linear filter with two time constants
+ """
+
+ NONE = mujoco.mjtDyn.mjDYN_NONE
+ INTEGRATOR = mujoco.mjtDyn.mjDYN_INTEGRATOR
+ FILTER = mujoco.mjtDyn.mjDYN_FILTER
+ FILTEREXACT = mujoco.mjtDyn.mjDYN_FILTEREXACT
+ MUSCLE = mujoco.mjtDyn.mjDYN_MUSCLE
+ # unsupported: USER
+
+
+class GainType(enum.IntEnum):
+ """Type of actuator gain.
+
+ Attributes:
+ FIXED: fixed gain
+ AFFINE: const + kp*length + kv*velocity
+ MUSCLE: muscle FLV curve computed by muscle_gain
+ """
+
+ FIXED = mujoco.mjtGain.mjGAIN_FIXED
+ AFFINE = mujoco.mjtGain.mjGAIN_AFFINE
+ MUSCLE = mujoco.mjtGain.mjGAIN_MUSCLE
+ # unsupported: USER
+
+
+class BiasType(enum.IntEnum):
+ """Type of actuator bias.
+
+ Attributes:
+ NONE: no bias
+ AFFINE: const + kp*length + kv*velocity
+ MUSCLE: muscle passive force computed by muscle_bias
+ """
+
+ NONE = mujoco.mjtBias.mjBIAS_NONE
+ AFFINE = mujoco.mjtBias.mjBIAS_AFFINE
+ MUSCLE = mujoco.mjtBias.mjBIAS_MUSCLE
+ # unsupported: USER
+
+
+class JointType(enum.IntEnum):
+ """Type of degree of freedom.
+
+ Attributes:
+ FREE: global position and orientation (quat) (7,)
+ BALL: orientation (quat) relative to parent (4,)
+ SLIDE: sliding distance along body-fixed axis (1,)
+ HINGE: rotation angle (rad) around body-fixed axis (1,)
+ """
+
+ FREE = mujoco.mjtJoint.mjJNT_FREE
+ BALL = mujoco.mjtJoint.mjJNT_BALL
+ SLIDE = mujoco.mjtJoint.mjJNT_SLIDE
+ HINGE = mujoco.mjtJoint.mjJNT_HINGE
+
+ def dof_width(self) -> int:
+ return {0: 6, 1: 3, 2: 1, 3: 1}[self.value]
+
+ def qpos_width(self) -> int:
+ return {0: 7, 1: 4, 2: 1, 3: 1}[self.value]
+
+
+class ConeType(enum.IntEnum):
+ """Type of friction cone.
+
+ Attributes:
+ PYRAMIDAL: pyramidal
+ ELLIPTIC: elliptic
+ """
+
+ PYRAMIDAL = mujoco.mjtCone.mjCONE_PYRAMIDAL
+ ELLIPTIC = mujoco.mjtCone.mjCONE_ELLIPTIC
+
+
+class IntegratorType(enum.IntEnum):
+ """Integrator mode.
+
+ Attributes:
+ EULER: semi-implicit Euler
+ RK4: 4th-order Runge Kutta
+ IMPLICITFAST: implicit in velocity, no rne derivative
+ """
+
+ EULER = mujoco.mjtIntegrator.mjINT_EULER
+ RK4 = mujoco.mjtIntegrator.mjINT_RK4
+ IMPLICITFAST = mujoco.mjtIntegrator.mjINT_IMPLICITFAST
+ # unsupported: IMPLICIT
+
+
+class GeomType(enum.IntEnum):
+ """Type of geometry.
+
+ Attributes:
+ PLANE: plane
+ HFIELD: heightfield
+ SPHERE: sphere
+ CAPSULE: capsule
+ ELLIPSOID: ellipsoid
+ CYLINDER: cylinder
+ BOX: box
+ MESH: mesh
+ SDF: sdf
+ """
+
+ PLANE = mujoco.mjtGeom.mjGEOM_PLANE
+ HFIELD = mujoco.mjtGeom.mjGEOM_HFIELD
+ SPHERE = mujoco.mjtGeom.mjGEOM_SPHERE
+ CAPSULE = mujoco.mjtGeom.mjGEOM_CAPSULE
+ ELLIPSOID = mujoco.mjtGeom.mjGEOM_ELLIPSOID
+ CYLINDER = mujoco.mjtGeom.mjGEOM_CYLINDER
+ BOX = mujoco.mjtGeom.mjGEOM_BOX
+ MESH = mujoco.mjtGeom.mjGEOM_MESH
+ SDF = mujoco.mjtGeom.mjGEOM_SDF
+ # unsupported: NGEOMTYPES, ARROW*, LINE, SKIN, LABEL, NONE
+
+
+class SolverType(enum.IntEnum):
+ """Constraint solver algorithm.
+
+ Attributes:
+ CG: Conjugate gradient (primal)
+ NEWTON: Newton (primal)
+ """
+
+ CG = mujoco.mjtSolver.mjSOL_CG
+ NEWTON = mujoco.mjtSolver.mjSOL_NEWTON
+ # unsupported: PGS
+
+
+class ConstraintType(enum.IntEnum):
+ """Type of constraint.
+
+ Attributes:
+ EQUALITY: equality constraint
+ FRICTION_DOF: dof friction
+ FRICTION_TENDON: tendon friction
+ LIMIT_JOINT: joint limit
+ LIMIT_TENDON: tendon limit
+ CONTACT_FRICTIONLESS: frictionless contact
+ CONTACT_PYRAMIDAL: frictional contact, pyramidal friction cone
+ CONTACT_ELLIPTIC: frictional contact, elliptic friction cone
+ """
+
+ EQUALITY = mujoco.mjtConstraint.mjCNSTR_EQUALITY
+ FRICTION_DOF = mujoco.mjtConstraint.mjCNSTR_FRICTION_DOF
+ FRICTION_TENDON = mujoco.mjtConstraint.mjCNSTR_FRICTION_TENDON
+ LIMIT_JOINT = mujoco.mjtConstraint.mjCNSTR_LIMIT_JOINT
+ LIMIT_TENDON = mujoco.mjtConstraint.mjCNSTR_LIMIT_TENDON
+ CONTACT_FRICTIONLESS = mujoco.mjtConstraint.mjCNSTR_CONTACT_FRICTIONLESS
+ CONTACT_PYRAMIDAL = mujoco.mjtConstraint.mjCNSTR_CONTACT_PYRAMIDAL
+ CONTACT_ELLIPTIC = mujoco.mjtConstraint.mjCNSTR_CONTACT_ELLIPTIC
+
+
+class SensorType(enum.IntEnum):
+ """Type of sensor.
+
+ Attributes:
+ MAGNETOMETER: magnetometer
+ CAMPROJECTION: camera projection
+ RANGEFINDER: scalar distance to nearest geom or site along z-axis
+ JOINTPOS: joint position
+ TENDONPOS: scalar tendon position
+ ACTUATORPOS: actuator position
+ BALLQUAT: ball joint orientation
+ JOINTLIMITPOS: joint limit distance-margin
+ TENDONLIMITPOS: tendon limit distance-margin
+ FRAMEPOS: frame position
+ FRAMEXAXIS: frame x-axis
+ FRAMEYAXIS: frame y-axis
+ FRAMEZAXIS: frame z-axis
+ FRAMEQUAT: frame orientation, represented as quaternion
+ SUBTREECOM: subtree center of mass
+ E_POTENTIAL: potential energy
+ E_KINETIC: kinetic energy
+ CLOCK: simulation time
+ VELOCIMETER: 3D linear velocity, in local frame
+ GYRO: 3D angular velocity, in local frame
+ JOINTVEL: joint velocity
+ TENDONVEL: scalar tendon velocity
+ ACTUATORVEL: actuator velocity
+ BALLANGVEL: ball joint angular velocity
+ JOINTLIMITVEL: joint limit velocity
+ TENDONLIMITVEL: tendon limit velocity
+ FRAMELINVEL: 3D linear velocity
+ FRAMEANGVEL: 3D angular velocity
+ SUBTREELINVEL: subtree linear velocity
+ SUBTREEANGMOM: subtree angular momentum
+ TOUCH: scalar contact normal forces summed over sensor zone
+ ACCELEROMETER: accelerometer
+ FORCE: force
+ TORQUE: torque
+ ACTUATORFRC: scalar actuator force, measured at the joint
+ TENDONACTFRC: scalar actuator force, measured at the tendon
+ JOINTACTFRC: scalar actuator force, measured at the joint
+ JOINTLIMITFRC: joint limit force
+ TENDONLIMITFRC: tendon limit force
+ FRAMELINACC: 3D linear acceleration
+ FRAMEANGACC: 3D angular acceleration
+ """
+
+ MAGNETOMETER = mujoco.mjtSensor.mjSENS_MAGNETOMETER
+ CAMPROJECTION = mujoco.mjtSensor.mjSENS_CAMPROJECTION
+ RANGEFINDER = mujoco.mjtSensor.mjSENS_RANGEFINDER
+ JOINTPOS = mujoco.mjtSensor.mjSENS_JOINTPOS
+ TENDONPOS = mujoco.mjtSensor.mjSENS_TENDONPOS
+ ACTUATORPOS = mujoco.mjtSensor.mjSENS_ACTUATORPOS
+ BALLQUAT = mujoco.mjtSensor.mjSENS_BALLQUAT
+ JOINTLIMITPOS = mujoco.mjtSensor.mjSENS_JOINTLIMITPOS
+ TENDONLIMITPOS = mujoco.mjtSensor.mjSENS_TENDONLIMITPOS
+ FRAMEPOS = mujoco.mjtSensor.mjSENS_FRAMEPOS
+ FRAMEXAXIS = mujoco.mjtSensor.mjSENS_FRAMEXAXIS
+ FRAMEYAXIS = mujoco.mjtSensor.mjSENS_FRAMEYAXIS
+ FRAMEZAXIS = mujoco.mjtSensor.mjSENS_FRAMEZAXIS
+ FRAMEQUAT = mujoco.mjtSensor.mjSENS_FRAMEQUAT
+ SUBTREECOM = mujoco.mjtSensor.mjSENS_SUBTREECOM
+ E_POTENTIAL = mujoco.mjtSensor.mjSENS_E_POTENTIAL
+ E_KINETIC = mujoco.mjtSensor.mjSENS_E_KINETIC
+ CLOCK = mujoco.mjtSensor.mjSENS_CLOCK
+ VELOCIMETER = mujoco.mjtSensor.mjSENS_VELOCIMETER
+ GYRO = mujoco.mjtSensor.mjSENS_GYRO
+ JOINTVEL = mujoco.mjtSensor.mjSENS_JOINTVEL
+ TENDONVEL = mujoco.mjtSensor.mjSENS_TENDONVEL
+ ACTUATORVEL = mujoco.mjtSensor.mjSENS_ACTUATORVEL
+ BALLANGVEL = mujoco.mjtSensor.mjSENS_BALLANGVEL
+ JOINTLIMITVEL = mujoco.mjtSensor.mjSENS_JOINTLIMITVEL
+ TENDONLIMITVEL = mujoco.mjtSensor.mjSENS_TENDONLIMITVEL
+ FRAMELINVEL = mujoco.mjtSensor.mjSENS_FRAMELINVEL
+ FRAMEANGVEL = mujoco.mjtSensor.mjSENS_FRAMEANGVEL
+ SUBTREELINVEL = mujoco.mjtSensor.mjSENS_SUBTREELINVEL
+ SUBTREEANGMOM = mujoco.mjtSensor.mjSENS_SUBTREEANGMOM
+ TOUCH = mujoco.mjtSensor.mjSENS_TOUCH
+ ACCELEROMETER = mujoco.mjtSensor.mjSENS_ACCELEROMETER
+ FORCE = mujoco.mjtSensor.mjSENS_FORCE
+ TORQUE = mujoco.mjtSensor.mjSENS_TORQUE
+ ACTUATORFRC = mujoco.mjtSensor.mjSENS_ACTUATORFRC
+ TENDONACTFRC = mujoco.mjtSensor.mjSENS_TENDONACTFRC
+ JOINTACTFRC = mujoco.mjtSensor.mjSENS_JOINTACTFRC
+ JOINTLIMITFRC = mujoco.mjtSensor.mjSENS_JOINTLIMITFRC
+ TENDONLIMITFRC = mujoco.mjtSensor.mjSENS_TENDONLIMITFRC
+ FRAMELINACC = mujoco.mjtSensor.mjSENS_FRAMELINACC
+ FRAMEANGACC = mujoco.mjtSensor.mjSENS_FRAMEANGACC
+
+
+class ObjType(enum.IntEnum):
+ """Type of object.
+
+ Attributes:
+ UNKNOWN: unknown object type
+ BODY: body
+ XBODY: body, used to access regular frame instead of i-frame
+ GEOM: geom
+ SITE: site
+ CAMERA: camera
+ """
+
+ UNKNOWN = mujoco.mjtObj.mjOBJ_UNKNOWN
+ BODY = mujoco.mjtObj.mjOBJ_BODY
+ XBODY = mujoco.mjtObj.mjOBJ_XBODY
+ GEOM = mujoco.mjtObj.mjOBJ_GEOM
+ SITE = mujoco.mjtObj.mjOBJ_SITE
+ CAMERA = mujoco.mjtObj.mjOBJ_CAMERA
+
+
+class EqType(enum.IntEnum):
+ """Type of equality constraint.
+
+ Attributes:
+ CONNECT: connect two bodies at a point (ball joint)
+ JOINT: couple the values of two scalar joints with cubic
+ WELD: fix relative position and orientation of two bodies
+ """
+
+ CONNECT = mujoco.mjtEq.mjEQ_CONNECT
+ WELD = mujoco.mjtEq.mjEQ_WELD
+ JOINT = mujoco.mjtEq.mjEQ_JOINT
+ TENDON = mujoco.mjtEq.mjEQ_TENDON
+ # unsupported: FLEX, DISTANCE
+
+
+class WrapType(enum.IntEnum):
+ """Type of tendon wrapping object.
+
+ Attributes:
+ JOINT: constant moment arm
+ PULLEY: pulley used to split tendon
+ SITE: pass through site
+ SPHERE: wrap around sphere
+ CYLINDER: wrap around (infinite) cylinder
+ """
+
+ JOINT = mujoco.mjtWrap.mjWRAP_JOINT
+ PULLEY = mujoco.mjtWrap.mjWRAP_PULLEY
+ SITE = mujoco.mjtWrap.mjWRAP_SITE
+ SPHERE = mujoco.mjtWrap.mjWRAP_SPHERE
+ CYLINDER = mujoco.mjtWrap.mjWRAP_CYLINDER
+
+
+class vec5f(wp.types.vector(length=5, dtype=float)):
+ pass
+
+
+class vec6f(wp.types.vector(length=6, dtype=float)):
+ pass
+
+
+class vec10f(wp.types.vector(length=10, dtype=float)):
+ pass
+
+
+class vec11f(wp.types.vector(length=11, dtype=float)):
+ pass
+
+
+vec5 = vec5f
+vec6 = vec6f
+vec10 = vec10f
+vec11 = vec11f
+
+
+class BroadphaseType(enum.IntEnum):
+ """Type of broadphase algorithm.
+
+ Attributes:
+ NXN: Broad phase checking all pairs
+ SAP_TILE: Sweep and prune broad phase using tile sort
+ SAP_SEGMENTED: Sweep and prune broad phase using segment sort
+ """
+
+ NXN = 0
+ SAP_TILE = 1
+ SAP_SEGMENTED = 2
+
+
+@dataclasses.dataclass
+class Option:
+ """Physics options.
+
+ Attributes:
+ timestep: simulation timestep
+ impratio: ratio of friction-to-normal contact impedance
+ tolerance: main solver tolerance
+ ls_tolerance: CG/Newton linesearch tolerance
+ gravity: gravitational acceleration
+ magnetic: global magnetic flux
+ integrator: integration mode (mjtIntegrator)
+ cone: type of friction cone (mjtCone)
+ solver: solver algorithm (mjtSolver)
+ iterations: number of main solver iterations
+ ls_iterations: maximum number of CG/Newton linesearch iterations
+ disableflags: bit flags for disabling standard features
+ enableflags: bit flags for enabling optional features
+ is_sparse: whether to use sparse representations
+ gjk_iterations: number of Gjk iterations in the convex narrowphase
+ epa_iterations: number of Epa iterations in the convex narrowphase
+ ls_parallel: evaluate engine solver step sizes in parallel
+ wind: wind (for lift, drag, and viscosity)
+ has_fluid: True if wind, density, or viscosity are non-zero at put_model time
+ density: density of medium
+ viscosity: viscosity of medium
+ broadphase: broadphase type, 0: nxn, 1: sap_tile, 2: sap_segmented
+ broadphase_filter: broadphase filter bitflag
+ graph_conditional: flag to use cuda graph conditional, should be False when JAX is used
+ sdf_initpoints: number of starting points for gradient descent
+ sdf_iterations: max number of iterations for gradient descent
+ run_collision_detection: if False, skips collision detection and allows user-populated
+ contacts during the physics step (as opposed to DisableBit.CONTACT which explicitly
+ zeros out the contacts at each step)
+ """
+
+ timestep: wp.array(dtype=float)
+ impratio: wp.array(dtype=float)
+ tolerance: wp.array(dtype=float)
+ ls_tolerance: wp.array(dtype=float)
+ gravity: wp.array(dtype=wp.vec3)
+ magnetic: wp.array(dtype=wp.vec3)
+ integrator: int
+ cone: int
+ solver: int
+ iterations: int
+ ls_iterations: int
+ disableflags: int
+ enableflags: int
+ is_sparse: bool
+ gjk_iterations: int # warp only
+ epa_iterations: int # warp only
+ ls_parallel: bool
+ wind: wp.array(dtype=wp.vec3)
+ has_fluid: bool
+ density: wp.array(dtype=float)
+ viscosity: wp.array(dtype=float)
+ broadphase: int # warp only
+ broadphase_filter: int # warp only
+ graph_conditional: bool # warp only
+ sdf_initpoints: int
+ sdf_iterations: int
+ run_collision_detection: bool # warp only
+
+
+@dataclasses.dataclass
+class Statistic:
+ """Model statistics (in qpos0).
+
+ Attributes:
+ meaninertia: mean diagonal inertia
+ """
+
+ meaninertia: float
+
+
+@dataclasses.dataclass
+class Constraint:
+ """Constraint data.
+
+ Attributes:
+ type: constraint type (mjtConstraint) (nworld, njmax)
+ id: id of object of specific type (nworld, njmax)
+ J: constraint Jacobian (nworld, njmax, nv)
+ pos: constraint position (equality, contact) (nworld, njmax)
+ margin: inclusion margin (contact) (nworld, njmax)
+ D: constraint mass (nworld, njmax)
+ vel: velocity in constraint space: J*qvel (nworld, njmax)
+ aref: reference pseudo-acceleration (nworld, njmax)
+ frictionloss: frictionloss (friction) (nworld, njmax)
+ force: constraint force in constraint space (nworld, njmax)
+ Jaref: Jac*qacc - aref (nworld, njmax)
+ Ma: M*qacc (nworld, nv)
+ grad: gradient of master cost (nworld, nv)
+ grad_dot: dot(grad, grad) (nworld,)
+ Mgrad: M / grad (nworld, nv)
+ search: linesearch vector (nworld, nv)
+ search_dot: dot(search, search) (nworld,)
+ gauss: gauss Cost (nworld,)
+ cost: constraint + Gauss cost (nworld,)
+ prev_cost: cost from previous iter (nworld,)
+ active: active (quadratic) constraints (nworld, njmax)
+ gtol: linesearch termination tolerance (nworld,)
+ mv: qM @ search (nworld, nv)
+ jv: efc_J @ search (nworld, njmax)
+ quad: quadratic cost coefficients (nworld, njmax, 3)
+ quad_gauss: quadratic cost gauss coefficients (nworld, 3)
+ h: cone hessian (nworld, nv, nv)
+ alpha: line search step size (nworld,)
+ prev_grad: previous grad (nworld, nv)
+ prev_Mgrad: previous Mgrad (nworld, nv)
+ beta: polak-ribiere beta (nworld,)
+ beta_num: numerator of beta (nworld,)
+ beta_den: denominator of beta (nworld,)
+ done: solver done (nworld,)
+ ls_done: linesearch done (nworld,)
+ p0: initial point (nworld, 3)
+ lo: low point bounding the line search interval (nworld, 3)
+ lo_alpha: alpha for low point (nworld,)
+ hi: high point bounding the line search interval (nworld, 3)
+ hi_alpha: alpha for high point (nworld,)
+ lo_next: next low point (nworld, 3)
+ lo_next_alpha: alpha for next low point (nworld,)
+ hi_next: next high point (nworld, 3)
+ hi_next_alpha: alpha for next high point (nworld,)
+ mid: loss at mid_alpha (nworld, 3)
+ mid_alpha: midpoint between lo_alpha and hi_alpha (nworld,)
+ cost_candidate: costs associated with step sizes (nworld, nlsp)
+ u: friction cone (normal and tangents) (nconmax, 6)
+ uu: elliptic cone variables (nconmax,)
+ uv: elliptic cone variables (nconmax,)
+ vv: elliptic cone variables (nconmax,)
+ condim: if contact: condim, else: -1 (nworld, njmax)
+ """
+
+ type: wp.array2d(dtype=int)
+ id: wp.array2d(dtype=int)
+ J: wp.array3d(dtype=float)
+ pos: wp.array2d(dtype=float)
+ margin: wp.array2d(dtype=float)
+ D: wp.array2d(dtype=float)
+ vel: wp.array2d(dtype=float)
+ aref: wp.array2d(dtype=float)
+ frictionloss: wp.array2d(dtype=float)
+ force: wp.array2d(dtype=float)
+ Jaref: wp.array2d(dtype=float)
+ Ma: wp.array2d(dtype=float)
+ grad: wp.array2d(dtype=float)
+ cholesky_L_tmp: wp.array3d(dtype=float)
+ cholesky_y_tmp: wp.array2d(dtype=float)
+ grad_dot: wp.array(dtype=float)
+ Mgrad: wp.array2d(dtype=float)
+ search: wp.array2d(dtype=float)
+ search_dot: wp.array(dtype=float)
+ gauss: wp.array(dtype=float)
+ cost: wp.array(dtype=float)
+ prev_cost: wp.array(dtype=float)
+ active: wp.array2d(dtype=bool)
+ gtol: wp.array(dtype=float)
+ mv: wp.array2d(dtype=float)
+ jv: wp.array2d(dtype=float)
+ quad: wp.array2d(dtype=wp.vec3)
+ quad_gauss: wp.array(dtype=wp.vec3)
+ h: wp.array3d(dtype=float)
+ alpha: wp.array(dtype=float)
+ prev_grad: wp.array2d(dtype=float)
+ prev_Mgrad: wp.array2d(dtype=float)
+ beta: wp.array(dtype=float)
+ beta_num: wp.array(dtype=float)
+ beta_den: wp.array(dtype=float)
+ done: wp.array(dtype=bool)
+ # linesearch
+ ls_done: wp.array(dtype=bool)
+ p0: wp.array(dtype=wp.vec3)
+ lo: wp.array(dtype=wp.vec3)
+ lo_alpha: wp.array(dtype=float)
+ hi: wp.array(dtype=wp.vec3)
+ hi_alpha: wp.array(dtype=float)
+ lo_next: wp.array(dtype=wp.vec3)
+ lo_next_alpha: wp.array(dtype=float)
+ hi_next: wp.array(dtype=wp.vec3)
+ hi_next_alpha: wp.array(dtype=float)
+ mid: wp.array(dtype=wp.vec3)
+ mid_alpha: wp.array(dtype=float)
+ cost_candidate: wp.array2d(dtype=float)
+ # elliptic cone
+ u: wp.array(dtype=vec6)
+ uu: wp.array(dtype=float)
+ uv: wp.array(dtype=float)
+ vv: wp.array(dtype=float)
+ condim: wp.array2d(dtype=int)
+
+
+@dataclasses.dataclass
+class TileSet:
+ """Tiling configuration for decomposable block diagonal matrix.
+
+ For non-square, non-block-diagonal tiles, use two tilesets.
+
+ Attributes:
+ adr: address of each tile in the set
+ size: size of all the tiles in this set
+ """
+
+ adr: wp.array(dtype=int)
+ size: int
+
+
+# TODO(team): make Model/Data fields sort order match mujoco
+
+
+@dataclasses.dataclass
+class Model:
+ """Model definition and parameters.
+
+ Attributes:
+ nq: number of generalized coordinates
+ nv: number of degrees of freedom
+ nu: number of actuators/controls
+ na: number of activation states
+ nbody: number of bodies
+ njnt: number of joints
+ ngeom: number of geoms
+ nsite: number of sites
+ ncam: number of cameras
+ nlight: number of lights
+ nexclude: number of excluded geom pairs
+ neq: number of equality constraints
+ nmocap: number of mocap bodies
+ ngravcomp: number of bodies with nonzero gravcomp
+ nM: number of non-zeros in sparse inertia matrix
+ nC: number of non-zeros in sparse reduced dof-dof matrix
+ ntendon: number of tendons
+ nwrap: number of wrap objects in all tendon paths
+ nsensor: number of sensors
+ nsensordata: number of elements in sensor data vector
+ nmeshvert: number of vertices for all meshes
+ nmeshface: number of faces for all meshes
+ nmeshgraph: number of ints in mesh auxiliary data
+ nmeshpoly: number of polygons in all meshes
+ nmeshpolyvert: number of vertices in all polygons
+ nmeshpolymap: number of polygons in vertex map
+ nlsp: number of step sizes for parallel linsearch
+ npair: number of predefined geom pairs
+ nhfield: number of heightfields
+ nhfielddata: size of elevation data
+ opt: physics options
+ stat: model statistics
+ qpos0: qpos values at default pose (nworld, nq)
+ qpos_spring: reference pose for springs (nworld, nq)
+ qM_fullm_i: sparse mass matrix addressing
+ qM_fullm_j: sparse mass matrix addressing
+ qM_mulm_i: sparse mass matrix addressing
+ qM_mulm_j: sparse mass matrix addressing
+ qM_madr_ij: sparse mass matrix addressing
+ qLD_update_tree: dof tree ordering for qLD updates
+ qLD_update_treeadr: index of each dof tree level
+ M_rownnz: number of non-zeros in each row of qM (nv,)
+ M_rowadr: index of each row in qM (nv,)
+ M_colind: column indices of non-zeros in qM (nM,)
+ mapM2M: index mapping from M (legacy) to M (CSR) (nC)
+ qM_tiles: tiling configuration
+ body_tree: list of body ids by tree level
+ 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_mocapid: id of mocap data; -1: none (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_geomnum: number of geoms (nbody,)
+ body_geomadr: start addr of geoms; -1: no geoms (nbody,)
+ body_pos: position offset rel. to parent body (nworld, nbody, 3)
+ body_quat: orientation offset rel. to parent body (nworld, nbody, 4)
+ body_ipos: local position of center of mass (nworld, nbody, 3)
+ body_iquat: local orientation of inertia ellipsoid (nworld, nbody, 4)
+ body_mass: mass (nworld, nbody,)
+ body_subtreemass: mass of subtree starting at this body (nworld, nbody,)
+ subtree_mass: mass of subtree (nworld, nbody,)
+ body_inertia: diagonal inertia in ipos/iquat frame (nworld, nbody, 3)
+ body_invweight0: mean inv inert in qpos0 (trn, rot) (nworld, nbody, 2)
+ body_contype: OR over all geom contypes (nbody,)
+ body_conaffinity: OR over all geom conaffinities (nbody,)
+ body_gravcomp: antigravity force, units of body weight (nworld, nbody)
+ 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_limited: does joint have limits (njnt,)
+ jnt_actfrclimited: does joint have actuator force limits (njnt,)
+ jnt_solref: constraint solver reference: limit (nworld, njnt, mjNREF)
+ jnt_solimp: constraint solver impedance: limit (nworld, njnt, mjNIMP)
+ jnt_pos: local anchor position (nworld, njnt, 3)
+ jnt_axis: local joint axis (nworld, njnt, 3)
+ jnt_stiffness: stiffness coefficient (nworld, njnt)
+ jnt_range: joint limits (nworld, njnt, 2)
+ jnt_actfrcrange: range of total actuator force (nworld, njnt, 2)
+ jnt_margin: min distance for limit detection (nworld, njnt)
+ jnt_limited_slide_hinge_adr: limited/slide/hinge jntadr
+ jnt_limited_ball_adr: limited/ball jntadr
+ jnt_actgravcomp: is gravcomp force applied via actuators (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_Madr: dof address in M-diagonal (nv,)
+ dof_armature: dof armature inertia/mass (nworld, nv)
+ dof_damping: damping coefficient (nworld, nv)
+ dof_invweight0: diag. inverse inertia in qpos0 (nworld, nv)
+ dof_frictionloss: dof friction loss (nworld, nv)
+ dof_solimp: constraint solver impedance: frictionloss (nworld, nv, NIMP)
+ dof_solref: constraint solver reference: frictionloss (nworld, nv, NREF)
+ dof_tri_row: np.tril_indices (mjm.nv)[0]
+ dof_tri_col: np.tril_indices (mjm.nv)[1]
+ 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: geom group inclusion/exclusion mask (ngeom,)
+ geom_matid: material id for rendering (nworld, ngeom,)
+ geom_priority: geom contact priority (ngeom,)
+ geom_solmix: mixing coef for solref/imp in geom pair (nworld, ngeom,)
+ geom_solref: constraint solver reference: contact (nworld, ngeom, mjNREF)
+ geom_solimp: constraint solver impedance: contact (nworld, ngeom, mjNIMP)
+ geom_size: geom-specific size parameters (ngeom, 3)
+ geom_aabb: bounding box, (center, size) (ngeom, 6)
+ geom_rbound: radius of bounding sphere (nworld, ngeom,)
+ geom_pos: local position offset rel. to body (nworld, ngeom, 3)
+ geom_quat: local orientation offset rel. to body (nworld, ngeom, 4)
+ geom_friction: friction for (slide, spin, roll) (nworld, ngeom, 3)
+ geom_margin: detect contact if dist bool:
+ """Check for intersection of two 2D line segments.
+
+ Args:
+ p1: 2D point from segment 1
+ p2: 2D point from segment 1
+ p3: 2D point from segment 2
+ p4: 2D point from segment 2
+
+ Returns:
+ intersection status of line segments
+ """
+ # compute determinant, check
+ det = (p4[1] - p3[1]) * (p2[0] - p1[0]) - (p4[0] - p3[0]) * (p2[1] - p1[1])
+
+ if wp.abs(det) < MJ_MINVAL:
+ return False
+
+ # compute intersection point on each line
+ a = ((p4[0] - p3[0]) * (p1[1] - p3[1]) - (p4[1] - p3[1]) * (p1[0] - p3[0])) / det
+ b = ((p2[0] - p1[0]) * (p1[1] - p3[1]) - (p2[1] - p1[1]) * (p1[0] - p3[0])) / det
+
+ if a >= 0 and a <= 1.0 and b >= 0.0 and b <= 1.0:
+ return True
+ else:
+ return False
+
+
+@wp.func
+def halton(index: int, base: int) -> float:
+ n0 = index
+ b = float(base)
+ f = float(1.0) / b
+ hn = float(0.0)
+
+ while n0 > 0:
+ n1 = n0 // base
+ r = n0 - n1 * base
+ hn += f * float(r)
+ f /= b
+ n0 = n1
+
+ return hn
+
+
+@wp.func
+def length_circle(p0: wp.vec2, p1: wp.vec2, ind: int, radius: float) -> float:
+ """Curve length along circle.
+
+ Args:
+ p0: 2D point
+ p1: 2D point
+ ind: input for flip
+ radius: circle radius
+
+ Returns:
+ curve length
+ """
+ # compute angle between 0 and pi
+ p0n, _ = math.normalize_with_norm(p0)
+ p1n, _ = math.normalize_with_norm(p1)
+
+ angle = wp.acos(wp.dot(p0n, p1n))
+
+ # flip if necessary
+ cross = p0[1] * p1[0] - p0[0] * p1[1]
+ if (cross > 0.0 and ind != 0) or (cross < 0.0 and ind == 0):
+ angle = 2.0 * wp.pi - angle
+
+ return radius * angle
+
+
+@wp.func
+def wrap_circle(end: wp.vec4, side: wp.vec2, radius: float) -> Tuple[float, wp.vec2, wp.vec2]:
+ """2D circle wrap.
+
+ Args:
+ end: two 2D points
+ side: optional 2D side point, no side point: wp.vec2(wp.inf)
+ radius: circle radius
+
+ Returns:
+ length of circular wrap or -1.0 if no wrap, pair of 2D wrap points
+ """
+ valid_side = wp.norm_l2(side) < wp.inf
+
+ end0 = wp.vec2(end[0], end[1])
+ end1 = wp.vec2(end[2], end[3])
+
+ sqlen0 = wp.dot(end0, end0)
+ sqlen1 = wp.dot(end1, end1)
+ sqrad = radius * radius
+
+ # either point inside circle or circle too small: no wrap
+ if (sqlen0 < sqrad) or (sqlen1 < sqrad) or (radius < MJ_MINVAL):
+ return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf)
+
+ # points too close: no wrap
+ dif = end1 - end0
+ dd = wp.dot(dif, dif)
+ if dd < MJ_MINVAL:
+ return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf)
+
+ # find nearest point on line segment to origin: a * dif + d0
+ a = -wp.dot(dif, end0) / dd
+ a = wp.clamp(a, 0.0, 1.0)
+
+ # check for intersection and side
+ tmp = a * dif + end0
+ if (wp.dot(tmp, tmp) > sqrad) and (not valid_side or wp.dot(side, tmp) >= 0.0):
+ return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf)
+
+ sqrt0 = wp.sqrt(sqlen0 - sqrad)
+ sqrt1 = wp.sqrt(sqlen1 - sqrad)
+
+ # construct the two solutions, compute goodness
+ sol00 = wp.vec2(
+ (end[0] * sqrad + radius * end[1] * sqrt0) / sqlen0,
+ (end[1] * sqrad - radius * end[0] * sqrt0) / sqlen0,
+ )
+ sol01 = wp.vec2(
+ (end[2] * sqrad - radius * end[3] * sqrt1) / sqlen1,
+ (end[3] * sqrad + radius * end[2] * sqrt1) / sqlen1,
+ )
+
+ sol10 = wp.vec2(
+ (end[0] * sqrad - radius * end[1] * sqrt0) / sqlen0,
+ (end[1] * sqrad + radius * end[0] * sqrt0) / sqlen0,
+ )
+ sol11 = wp.vec2(
+ (end[2] * sqrad + radius * end[3] * sqrt1) / sqlen1,
+ (end[3] * sqrad - radius * end[2] * sqrt1) / sqlen1,
+ )
+
+ # goodness: close to sd, or shorter path
+ if valid_side:
+ tmp0, _ = math.normalize_with_norm(sol00 + sol01)
+ good0 = wp.dot(tmp0, side)
+ tmp1, _ = math.normalize_with_norm(sol10 + sol11)
+ good1 = wp.dot(tmp1, side)
+ else:
+ tmp0 = sol00 - sol01
+ good0 = -wp.dot(tmp0, tmp0)
+ tmp1 = sol10 - sol11
+ good1 = -wp.dot(tmp1, tmp1)
+
+ # penalize for intersection
+ if is_intersect(end0, sol00, end1, sol01):
+ good0 = -10000.0
+ if is_intersect(end0, sol10, end1, sol11):
+ good1 = -10000.0
+
+ # select the better solution
+ if good0 > good1:
+ pnt0 = sol00
+ pnt1 = sol01
+ ind = 0
+ else:
+ pnt0 = sol10
+ pnt1 = sol11
+ ind = 1
+
+ # check for intersection
+ if is_intersect(end0, pnt0, end1, pnt1):
+ return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf)
+
+ # return curve length
+ return length_circle(pnt0, pnt1, ind, radius), pnt0, pnt1
+
+
+@wp.func
+def wrap_inside(
+ # In:
+ end: wp.vec4,
+ radius: float,
+ # TODO(team): update kernel analyzer to allow defaults
+ maxiter: int = 20, # kernel_analyzer: ignore
+ zinit: float = 1.0 - 1.0e-7, # kernel_analyzer: ignore
+ tolerance: float = 1.0e-6, # kernel_analyzer: ignore
+) -> Tuple[float, wp.vec2, wp.vec2]:
+ """2D inside wrap.
+
+ Args:
+ end: two 2D points
+ radius: circle radius
+ maxiter: maximum number of solver iterations
+ zinit: initialization for solver
+ tolerance: solver convergence tolerance
+
+ Returns:
+ 0.0 if wrap else -1.0, pair of 2D wrap points
+ """
+
+ end0 = wp.vec2(end[0], end[1])
+ end1 = wp.vec2(end[2], end[3])
+
+ # constants
+ len0 = wp.norm_l2(end0)
+ len1 = wp.norm_l2(end1)
+ dif = end1 - end0
+ dd = wp.dot(dif, dif)
+
+ # either point inside circle or circle too small: no wrap
+ if (len0 <= radius) or (len1 <= radius) or (radius < MJ_MINVAL) or (len0 < MJ_MINVAL) or (len1 < MJ_MINVAL):
+ return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf)
+
+ # segment-circle intersection: no wrap
+ if dd > MJ_MINVAL:
+ # find nearest point on line segment to origin: d0 + a * dif
+ a = -wp.dot(dif, end0) / dd
+
+ # in segment
+ if (a > 0.0) and (a < 1.0):
+ tmp = end0 + a * dif
+ if wp.norm_l2(tmp) <= radius:
+ return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf)
+
+ # prepare default in case of numerical failure: average
+ pnt = 0.5 * (end0 + end1)
+ pnt, _ = math.normalize_with_norm(pnt)
+ pnt *= radius
+
+ # compute function parameters: asin(A * z) + asin(B * z) - 2 * asin(z) + G = 0
+ A = radius / len0
+ B = radius / len1
+ sq_A = A * A
+ sq_B = B * B
+ cosG = (len0 * len0 + len1 * len1 - dd) / (2.0 * len0 * len1)
+ if cosG < -1.0 + MJ_MINVAL:
+ return -1.0, pnt, pnt
+ elif cosG > 1.0 - MJ_MINVAL:
+ return 0.0, pnt, pnt
+ G = wp.acos(cosG)
+
+ # init
+ z = zinit
+ f = wp.asin(A * z) + wp.asin(B * z) - 2.0 * wp.asin(z) + G
+
+ # make sure init is not on the other side
+ if f > 0.0:
+ return 0.0, pnt, pnt
+
+ # Newton method
+ iter = int(0)
+
+ while (iter < maxiter) and (wp.abs(f) > tolerance):
+ # derivative
+ sq_z = z * z
+ df = (
+ A / wp.max(MJ_MINVAL, wp.sqrt(1.0 - sq_z * sq_A))
+ + B / wp.max(MJ_MINVAL, wp.sqrt(1.0 - sq_z * sq_B))
+ - 2.0 / wp.max(MJ_MINVAL, wp.sqrt(1.0 - sq_z))
+ )
+
+ # check sign; SHOULD NOT OCCUR
+ if df > -MJ_MINVAL:
+ return 0.0, pnt, pnt
+
+ # new point
+ z1 = z - f / df
+
+ # make sure we are moving to the left; SHOULD NOT OCCUR
+ if z1 > z:
+ return 0.0, pnt, pnt
+
+ # update solution
+ z = z1
+ f = wp.asin(A * z) + wp.asin(B * z) - 2.0 * wp.asin(z) + G
+
+ # exit if positive: SHOULD NOT OCCUR
+ if f > tolerance:
+ return 0.0, pnt, pnt
+
+ iter += 1
+
+ # check convergence
+ if iter >= maxiter:
+ return 0.0, pnt, pnt
+
+ # finalize: rotation by ang from vec = a or b, depending on cross(a, b) sign
+ if end[0] * end[3] - end[1] * end[2] > 0.0:
+ vec = end0
+ ang = wp.asin(z) - wp.asin(A * z)
+ else:
+ vec = end1
+ ang = wp.asin(z) - wp.asin(B * z)
+
+ vec, _ = math.normalize_with_norm(vec)
+ pnt = wp.vec2(
+ radius * (wp.cos(ang) * vec[0] - wp.sin(ang) * vec[1]),
+ radius * (wp.sin(ang) * vec[0] + wp.cos(ang) * vec[1]),
+ )
+
+ return 0.0, pnt, pnt
+
+
+@wp.func
+def wrap(
+ x0: wp.vec3, x1: wp.vec3, pos: wp.vec3, mat: wp.mat33, radius: float, geomtype: int, side: wp.vec3
+) -> Tuple[float, wp.vec3, wp.vec3]:
+ """Wrap tendons around spheres and cylinders.
+
+ Args:
+ x0: 3D endpoint
+ x1: 3D endpoint
+ pos: position of geom
+ mat: orientation of geom
+ radius: geom radius
+ type: wrap type (mjtWrap)
+ side: 3D position for sidesite, no side point: wp.vec3(wp.inf)
+
+ Returns:
+ length of circular wrap else -1.0 if no wrap, pair of 3D wrap points
+ """
+ # check object type
+ if geomtype != int(WrapType.SPHERE.value) and geomtype != int(WrapType.CYLINDER.value):
+ return wp.inf, wp.vec3(wp.inf), wp.vec3(wp.inf)
+
+ # map sites to wrap object's local frame
+ matT = wp.transpose(mat)
+ p0 = matT @ (x0 - pos)
+ p1 = matT @ (x1 - pos)
+
+ # too close to origin: return
+ if (wp.norm_l2(p0) < MJ_MINVAL) or (wp.norm_l2(p1) < MJ_MINVAL):
+ return -1.0, wp.vec3(wp.inf), wp.vec3(wp.inf)
+
+ # construct 2D frame for circle wrap
+ if geomtype == int(WrapType.SPHERE.value):
+ # 1st axis = p0
+ axis0, _ = math.normalize_with_norm(p0)
+
+ # normal to p0-0-p1 plane = cross(p0, p1)
+ normal = wp.cross(p0, p1)
+ normal, nrm = math.normalize_with_norm(normal)
+
+ # if (p0, p1) parallel: different normal
+ if nrm < MJ_MINVAL:
+ # find max component of axis0
+ axis0_abs = wp.abs(axis0)
+ i = int(0)
+ if (axis0_abs[1] > axis0_abs[0]) and (axis0_abs[1] > axis0_abs[2]):
+ i = 1
+ if (axis0_abs[2] > axis0_abs[0]) and (axis0_abs[2] > axis0_abs[1]):
+ i = 2
+
+ # init second axis: 0 at i; 1 elsewhere
+ axis1 = wp.vec3(1.0)
+ axis1[i] = 0.0
+
+ # recompute normal
+ normal = wp.cross(axis0, axis1)
+ normal, _ = math.normalize_with_norm(normal)
+
+ # 2nd axis = cross(normal, p0)
+ axis1 = wp.cross(normal, axis0)
+ axis1, _ = math.normalize_with_norm(axis1)
+ else: # WrapType.CYLINDER
+ # 1st axis = x
+ axis0 = wp.vec3(1.0, 0.0, 0.0)
+
+ # 2nd axis = y
+ axis1 = wp.vec3(0.0, 1.0, 0.0)
+
+ # project points in 2D frame: p => end
+ end = wp.vec4(
+ wp.dot(p0, axis0),
+ wp.dot(p0, axis1),
+ wp.dot(p1, axis0),
+ wp.dot(p1, axis1),
+ )
+
+ # handle sidesite
+ valid_side = wp.norm_l2(side) < wp.inf
+
+ if valid_side:
+ # side point: apply same projection as x0, x1
+ sidepnt = matT @ (side - pos)
+
+ # side point: project and rescale
+ sidepnt_proj = wp.vec2(
+ wp.dot(sidepnt, axis0),
+ wp.dot(sidepnt, axis1),
+ )
+
+ sidepnt_proj, _ = math.normalize_with_norm(sidepnt_proj)
+ sidepnt_proj *= radius
+ else:
+ sidepnt_proj = wp.vec2(wp.inf)
+
+ # apply inside wrap
+ if valid_side and wp.norm_l2(sidepnt) < radius:
+ wlen, pnt0, pnt1 = wrap_inside(end, radius)
+ else: # apply circle wrap
+ wlen, pnt0, pnt1 = wrap_circle(end, sidepnt_proj, radius)
+
+ # no wrap: return
+ if wlen < 0.0:
+ return -1.0, wp.vec3(wp.inf), wp.vec3(wp.inf)
+
+ # reconstruct 3D points in local frame: res
+ res0 = axis0 * pnt0[0] + axis1 * pnt0[1]
+ res1 = axis0 * pnt1[0] + axis1 * pnt1[1]
+
+ # cylinder: correct along z
+ if geomtype == int(WrapType.CYLINDER.value):
+ # set vertical coordinates
+ L0 = wp.sqrt((p0[0] - res0[0]) * (p0[0] - res0[0]) + (p0[1] - res0[1]) * (p0[1] - res0[1]))
+ L1 = wp.sqrt((p1[0] - res1[0]) * (p1[0] - res1[0]) + (p1[1] - res1[1]) * (p1[1] - res1[1]))
+ res0[2] = p0[2] + (p1[2] - p0[2]) * L0 / (L0 + wlen + L1)
+ res1[2] = p0[2] + (p1[2] - p0[2]) * (L0 + wlen) / (L0 + wlen + L1)
+
+ # correct wlen for height
+ height = wp.abs(res1[2] - res0[2])
+ wlen = wp.sqrt(wlen * wlen + height * height)
+
+ # map back to global frame: wpnt
+ wpnt0 = mat @ res0 + pos
+ wpnt1 = mat @ res1 + pos
+
+ return wlen, wpnt0, wpnt1
+
+
+@wp.func
+def muscle_gain_length(length: float, lmin: float, lmax: float) -> float:
+ """Normalized muscle length-gain curve."""
+
+ if (lmin > length) or (length > lmax):
+ return 0.0
+
+ # mid-ranges (maximum is at 1.0)
+ a = 0.5 * (lmin + 1.0)
+ b = 0.5 * (1.0 + lmax)
+
+ if length <= a:
+ x = (length - lmin) / wp.max(MJ_MINVAL, a - lmin)
+ return 0.5 * x * x
+ elif length <= 1.0:
+ x = (1.0 - length) / wp.max(MJ_MINVAL, 1.0 - a)
+ return 1.0 - 0.5 * x * x
+ elif length <= b:
+ x = (length - 1.0) / wp.max(MJ_MINVAL, b - 1.0)
+ return 1.0 - 0.5 * x * x
+ else:
+ x = (lmax - length) / wp.max(MJ_MINVAL, lmax - b)
+ return 0.5 * x * x
+
+
+@wp.func
+def muscle_gain(len: float, vel: float, lengthrange: wp.vec2, acc0: float, prm: vec10) -> float:
+ """Muscle active force, prm = (range[2], force, scale, lmin, lmax, vmax, fpmax, fvmax)."""
+
+ # unpack parameters
+ range_ = wp.vec2(prm[0], prm[1])
+ force = prm[2]
+ scale = prm[3]
+ lmin = prm[4]
+ lmax = prm[5]
+ vmax = prm[6]
+ fvmax = prm[8]
+
+ # scale force if negative
+ if force < 0.0:
+ force = scale / wp.max(MJ_MINVAL, acc0)
+
+ # optimum length
+ L0 = (lengthrange[1] - lengthrange[0]) / wp.max(MJ_MINVAL, range_[1] - range_[0])
+
+ # normalized length and velocity
+ L = range_[0] + (len - lengthrange[0]) / wp.max(MJ_MINVAL, L0)
+ V = vel / wp.max(MJ_MINVAL, L0 * vmax)
+
+ # length curve
+ FL = muscle_gain_length(L, lmin, lmax)
+
+ # velocity curve
+ y = fvmax - 1.0
+ if V <= -1.0:
+ FV = 0.0
+ elif V <= 0.0:
+ FV = (V + 1.0) * (V + 1.0)
+ elif V <= y:
+ FV = fvmax - (y - V) * (y - V) / wp.max(MJ_MINVAL, y)
+ else:
+ FV = fvmax
+
+ # compute FVL and scale, make it negative
+ return -force * FL * FV
+
+
+@wp.func
+def muscle_bias(len: float, lengthrange: wp.vec2, acc0: float, prm: vec10) -> float:
+ """Calculates muscle passive force.
+
+ prm = (range[2], force, scale, lmin, lmax, vmax, fpmax, fvmax)."""
+
+ # unpack parameters
+ range_ = wp.vec2(prm[0], prm[1])
+ force = prm[2]
+ scale = prm[3]
+ lmax = prm[5]
+ fpmax = prm[7]
+
+ # scale force if negative
+ if force < 0.0:
+ force = scale / wp.max(MJ_MINVAL, acc0)
+
+ # optimum length
+ L0 = (lengthrange[1] - lengthrange[0]) / wp.max(MJ_MINVAL, range_[1] - range_[0])
+
+ # normalized length
+ L = range_[0] + (len - lengthrange[0]) / wp.max(MJ_MINVAL, L0)
+
+ # half-quadratic to (L0 + lmax) / 2, linear beyond
+ b = 0.5 * (1.0 + lmax)
+ if L <= 1.0:
+ return 0.0
+ elif L <= b:
+ x = (L - 1.0) / wp.max(MJ_MINVAL, b - 1.0)
+ return -force * fpmax * 0.5 * x * x
+ else:
+ x = (L - b) / wp.max(MJ_MINVAL, b - 1.0)
+ return -force * fpmax * (0.5 + x)
+
+
+@wp.func
+def _sigmoid(x: float) -> float:
+ """Sigmoid function over 0 <= x <= 1 using quintic polynomial."""
+
+ if x <= 0.0:
+ return 0.0
+
+ if x >= 1.0:
+ return 1.0
+
+ # sigmoid f(x) = 6 * x^5 - 15 * x^4 + 10 * x^3
+ # solution of f(0) = f'(0) = f''(0) = 0, f(1) = 1, f'(1) = f''(1) = 0
+ return x * x * x * (3.0 * x * (2.0 * x - 5.0) + 10.0)
+
+
+@wp.func
+def muscle_dynamics_timescale(dctrl: float, tau_act: float, tau_deact: float, smooth_width: float) -> float:
+ """Muscle time constant with optional smoothing."""
+
+ # hard switching
+ if smooth_width < MJ_MINVAL:
+ if dctrl > 0.0:
+ return tau_act
+ else:
+ return tau_deact
+ else: # smooth switching
+ # scale by width, center around 0.5 midpoint, rescale to bounds
+ return tau_deact + (tau_act - tau_deact) * _sigmoid(dctrl / smooth_width + 0.5)
+
+
+@wp.func
+def muscle_dynamics(control: float, activation: float, prm: vec10) -> float:
+ """Muscle activation dynamics, prm = (tau_act, tau_deact, smooth_width)."""
+
+ # clamp control
+ ctrlclamp = wp.clamp(control, 0.0, 1.0)
+
+ # clamp activation
+ actclamp = wp.clamp(activation, 0.0, 1.0)
+
+ # compute timescales as in Millard et al. (2013) https://doi.org/10.1115/1.4023390
+ tau_act = prm[0] * (0.5 + 1.5 * actclamp) # activation timescale
+ tau_deact = prm[1] / (0.5 + 1.5 * actclamp) # deactivation timescale
+ smooth_width = prm[2] # width of smoothing sigmoid
+ dctrl = ctrlclamp - activation # excess excitation
+
+ tau = muscle_dynamics_timescale(dctrl, tau_act, tau_deact, smooth_width)
+
+ # filter output
+ return dctrl / wp.max(MJ_MINVAL, tau)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc_test.py
new file mode 100644
index 00000000..dbd43212
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc_test.py
@@ -0,0 +1,573 @@
+# Copyright 2025 The Newton Developers
+#
+# 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 miscellaneous utilities."""
+
+from typing import Tuple
+
+import numpy as np
+import warp as wp
+from absl.testing import absltest
+from absl.testing import parameterized
+
+from mujoco.mjx.third_party.mujoco_warp._src import util_misc
+from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
+from mujoco.mjx.third_party.mujoco_warp._src.types import WrapType
+from mujoco.mjx.third_party.mujoco_warp._src.types import vec10
+
+
+def _assert_eq(a, b, name):
+ tol = 1e-3 # avoid test noise
+ err_msg = f"mismatch: {name}"
+ np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+
+
+def _is_intersect(p1: np.array, p2: np.array, p3: np.array, p4: np.array) -> bool:
+ intersect = wp.empty(1, dtype=bool)
+
+ @wp.kernel
+ def is_intersect(
+ # In:
+ p1: wp.vec2,
+ p2: wp.vec2,
+ p3: wp.vec2,
+ p4: wp.vec2,
+ # Out:
+ intersect_out: wp.array(dtype=bool),
+ ):
+ intersect_out[0] = util_misc.is_intersect(p1, p2, p3, p4)
+
+ wp.launch(
+ is_intersect,
+ dim=(1,),
+ inputs=[
+ wp.vec2(p1[0], p1[1]),
+ wp.vec2(p2[0], p2[1]),
+ wp.vec2(p3[0], p3[1]),
+ wp.vec2(p4[0], p4[1]),
+ ],
+ outputs=[
+ intersect,
+ ],
+ )
+ return intersect.numpy()[0]
+
+
+def _length_circle(p0: np.array, p1: np.array, ind: int, radius: float) -> float:
+ length = wp.empty(1, dtype=float)
+
+ @wp.kernel
+ def length_circle(
+ # In:
+ p0: wp.vec2,
+ p1: wp.vec2,
+ ind: int,
+ radius: float,
+ # Out:
+ length_out: wp.array(dtype=float),
+ ):
+ length_out[0] = util_misc.length_circle(p0, p1, ind, radius)
+
+ wp.launch(
+ length_circle,
+ dim=(1,),
+ inputs=[wp.vec2(p0[0], p0[1]), wp.vec2(p1[0], p1[1]), ind, radius],
+ outputs=[
+ length,
+ ],
+ )
+ return length.numpy()[0]
+
+
+def _wrap_circle(end: np.array, side: np.array, radius: float) -> Tuple[float, np.array, np.array]:
+ length = wp.empty(1, dtype=float)
+ wpnt0 = wp.empty(1, dtype=wp.vec2)
+ wpnt1 = wp.empty(1, dtype=wp.vec2)
+
+ @wp.kernel
+ def wrap_circle(
+ # In:
+ end: wp.vec4,
+ side: wp.vec2,
+ radius: float,
+ # Out:
+ length_out: wp.array(dtype=float),
+ wpnt0_out: wp.array(dtype=wp.vec2),
+ wpnt1_out: wp.array(dtype=wp.vec2),
+ ):
+ length_, wpnt0_, wpnt1_ = util_misc.wrap_circle(end, side, radius)
+ length_out[0] = length_
+ wpnt0_out[0] = wpnt0_
+ wpnt1_out[0] = wpnt1_
+
+ wp.launch(
+ wrap_circle,
+ dim=(1,),
+ inputs=[
+ wp.vec4(end[0], end[1], end[2], end[3]),
+ wp.vec2(side[0], side[1]),
+ radius,
+ ],
+ outputs=[
+ length,
+ wpnt0,
+ wpnt1,
+ ],
+ )
+ return length.numpy()[0], wpnt0.numpy()[0], wpnt1.numpy()[0]
+
+
+def _wrap_inside(end: np.array, radius: float) -> Tuple[float, np.array, np.array]:
+ length = wp.empty(1, dtype=float)
+ wpnt0 = wp.empty(1, dtype=wp.vec2)
+ wpnt1 = wp.empty(1, dtype=wp.vec2)
+
+ @wp.kernel
+ def wrap_inside(
+ # In:
+ end: wp.vec4,
+ radius: float,
+ # Out:
+ length_out: wp.array(dtype=float),
+ wpnt0_out: wp.array(dtype=wp.vec2),
+ wpnt1_out: wp.array(dtype=wp.vec2),
+ ):
+ length_, wpnt0_, wpnt1_ = util_misc.wrap_inside(end, radius)
+ length_out[0] = length_
+ wpnt0_out[0] = wpnt0_
+ wpnt1_out[0] = wpnt1_
+
+ wp.launch(
+ wrap_inside,
+ dim=(1,),
+ inputs=[wp.vec4(end[0], end[1], end[2], end[3]), radius],
+ outputs=[
+ length,
+ wpnt0,
+ wpnt1,
+ ],
+ )
+ return length.numpy()[0], wpnt0.numpy()[0], wpnt1.numpy()[0]
+
+
+def _wrap(
+ x0: np.array,
+ x1: np.array,
+ xpos: np.array,
+ xmat: np.array,
+ radius: float,
+ geomtype: int,
+ side: np.array,
+) -> Tuple[float, np.array, np.array]:
+ length = wp.empty(1, dtype=float)
+ wpnt0 = wp.empty(1, dtype=wp.vec3)
+ wpnt1 = wp.empty(1, dtype=wp.vec3)
+
+ @wp.kernel
+ def wrap(
+ # In:
+ x0: wp.vec3,
+ x1: wp.vec3,
+ pos: wp.vec3,
+ mat: wp.mat33,
+ radius: float,
+ geomtype: int,
+ side: wp.vec3,
+ # Out:
+ length_out: wp.array(dtype=float),
+ wpnt0_out: wp.array(dtype=wp.vec3),
+ wpnt1_out: wp.array(dtype=wp.vec3),
+ ):
+ length_, wpnt0_, wpnt1_ = util_misc.wrap(x0, x1, pos, mat, radius, geomtype, side)
+ length_out[0] = length_
+ wpnt0_out[0] = wpnt0_
+ wpnt1_out[0] = wpnt1_
+
+ wp.launch(
+ wrap,
+ dim=(1,),
+ inputs=[
+ wp.vec3(x0[0], x0[1], x0[2]),
+ wp.vec3(x1[0], x1[1], x1[2]),
+ wp.vec3(xpos[0], xpos[1], xpos[2]),
+ wp.mat33(
+ xmat[0, 0],
+ xmat[0, 1],
+ xmat[0, 2],
+ xmat[1, 0],
+ xmat[1, 1],
+ xmat[1, 2],
+ xmat[2, 0],
+ xmat[2, 1],
+ xmat[2, 2],
+ ),
+ radius,
+ geomtype,
+ wp.vec3(side[0], side[1], side[2]),
+ ],
+ outputs=[
+ length,
+ wpnt0,
+ wpnt1,
+ ],
+ )
+ return length.numpy()[0], wpnt0.numpy()[0], wpnt1.numpy()[0]
+
+
+def _muscle_dynamics_millard(ctrl, act, prm):
+ """Compute time constant as in Millard et al. (2013) https://doi.org/10.1115/1.4023390."""
+
+ # clamp control
+ ctrlclamp = np.clip(ctrl, 0.0, 1.0)
+
+ # clamp activation
+ actclamp = np.clip(act, 0.0, 1.0)
+
+ if ctrlclamp > act:
+ tau = prm[0] * (0.5 + 1.5 * actclamp)
+ else:
+ tau = prm[1] / (0.5 + 1.5 * actclamp)
+
+ # filter output
+ return (ctrlclamp - act) / np.maximum(MJ_MINVAL, tau)
+
+
+def _muscle_dynamics(ctrl, act, prm):
+ @wp.kernel
+ def muscle_dynamics(control: float, activation: float, prm: vec10, dynamics_out: wp.array(dtype=float)):
+ dynamics_out[0] = util_misc.muscle_dynamics(control, activation, prm)
+
+ output = wp.empty(1, dtype=float)
+ wp.launch(
+ muscle_dynamics,
+ dim=(1,),
+ inputs=[
+ ctrl,
+ act,
+ vec10(prm[0], prm[1], prm[2], 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0),
+ ],
+ outputs=[output],
+ )
+
+ return output.numpy()[0]
+
+
+def _muscle_gain_length(length, lmin, lmax):
+ @wp.kernel
+ def muscle_gain_length(length: float, lmin: float, lmax: float, gain_length_out: wp.array(dtype=float)):
+ gain_length_out[0] = util_misc.muscle_gain_length(length, lmin, lmax)
+
+ output = wp.empty(1, dtype=float)
+ wp.launch(muscle_gain_length, dim=(1,), inputs=[length, lmin, lmax], outputs=[output])
+
+ return output.numpy()[0]
+
+
+def _muscle_dynamics_timescale(dctrl, tau_act, tau_deact, smooth_width):
+ @wp.kernel
+ def muscle_gain_length(
+ dctrl: float, tau_act: float, tau_deact: float, smooth_width: float, dynamics_timescale_out: wp.array(dtype=float)
+ ):
+ dynamics_timescale_out[0] = util_misc.muscle_dynamics_timescale(dctrl, tau_act, tau_deact, smooth_width)
+
+ output = wp.empty(1, dtype=float)
+ wp.launch(
+ muscle_gain_length,
+ dim=(1,),
+ inputs=[dctrl, tau_act, tau_deact, smooth_width],
+ outputs=[output],
+ )
+
+ return output.numpy()[0]
+
+
+class UtilMiscTest(parameterized.TestCase):
+ def test_is_intersect(self):
+ self.assertFalse(
+ _is_intersect(
+ np.array([0, 0]),
+ np.array([1, 0]),
+ np.array([0, 1]),
+ np.array([1, 1]),
+ )
+ )
+
+ self.assertTrue(
+ _is_intersect(
+ np.array([0, 0]),
+ np.array([1, 0]),
+ np.array([0.5, -1]),
+ np.array([0.5, 1]),
+ )
+ )
+
+ self.assertFalse(
+ _is_intersect(
+ np.array([0, 0]),
+ np.array([0, 0]),
+ np.array([0, 0]),
+ np.array([0, 0]),
+ )
+ )
+
+ def test_length_circle(self):
+ _assert_eq(
+ _length_circle(np.array([0, 1]), np.array([1, 0]), 0, 1.0),
+ 0.5 * np.pi,
+ "length_circle",
+ )
+ _assert_eq(
+ _length_circle(np.array([0, 1]), np.array([1, 0]), 1, 1.0),
+ 1.5 * np.pi,
+ "length_circle",
+ )
+ _assert_eq(
+ _length_circle(np.array([1, 0]), np.array([0, 1]), 0, 1.0),
+ 1.5 * np.pi,
+ "length_circle",
+ )
+ _assert_eq(
+ _length_circle(np.array([1, 0]), np.array([0, 1]), 1, 1.0),
+ 0.5 * np.pi,
+ "length_circle",
+ )
+
+ def test_wrap_circle(self):
+ # no wrap
+ wlen, wpnt0, wpnt1 = _wrap_circle(np.array([1, 0, 0, 1]), np.array([np.inf, np.inf]), 0.1)
+ _assert_eq(wlen, -1.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1")
+
+ # no wrap
+ wlen, wpnt0, wpnt1 = _wrap_circle(np.array([1, 0, 0, 1]), np.array([0.0, 0.0]), 0.1)
+ _assert_eq(wlen, -1.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1")
+
+ # wrap
+ wlen, wpnt0, wpnt1 = _wrap_circle(np.array([np.sqrt(2.0), 0, 0, np.sqrt(2.0)]), np.array([np.inf, np.inf]), 1.0 + 5e-4)
+ _assert_eq(wlen, 0.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.sqrt(2.0) / 2.0, np.sqrt(2.0) / 2.0]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.sqrt(2.0) / 2.0, np.sqrt(2.0) / 2.0]), "wpnt1")
+
+ # wrap
+ wlen, wpnt0, wpnt1 = _wrap_circle(np.array([np.sqrt(2.0), 0, 0, np.sqrt(2.0)]), np.array([0.0, 0.0]), 1.0 + 5e-4)
+ _assert_eq(wlen, 0.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.sqrt(2.0) / 2.0, np.sqrt(2.0) / 2.0]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.sqrt(2.0) / 2.0, np.sqrt(2.0) / 2.0]), "wpnt1")
+
+ # wrap
+ wlen, wpnt0, wpnt1 = _wrap_circle(np.array([1.0, 0, 0, 1.0]), np.array([0.0, 0.0]), 1.0)
+ _assert_eq(wlen, 0.5 * np.pi, "wlen")
+ _assert_eq(wpnt0, np.array([1.0, 0.0]), "wpnt0")
+ _assert_eq(wpnt1, np.array([0.0, 1.0]), "wpnt1")
+
+ # wrap w/ sidesite
+ wlen, wpnt0, wpnt1 = _wrap_circle(np.array([0, -100, 0, 100]), np.array([0.2, 0.0]), 0.1)
+
+ # wlen, wpnt0[1], wpnt1[1] are ~0
+ _assert_eq(wlen, 0.0, "wlen")
+ _assert_eq(wpnt0, np.array([0.1, 0]), "wpnt0")
+ _assert_eq(wpnt1, np.array([0.1, 0]), "wpnt1")
+
+ wlen, wpnt0, wpnt1 = _wrap_circle(np.array([0, -100, 0, 100]), np.array([-0.2, 0.0]), 0.1)
+
+ # wlen, wpnt0[1], wpnt1[1] are ~0
+ _assert_eq(wlen, 0.0, "wlen")
+ _assert_eq(wpnt0, np.array([-0.1, 0]), "wpnt0")
+ _assert_eq(wpnt1, np.array([-0.1, 0]), "wpnt1")
+
+ def test_wrap_inside(self):
+ wlen, wpnt0, wpnt1 = _wrap_inside(np.array([1, 0, 0, 1]), 0.7071)
+ _assert_eq(wlen, 0.0, "wlen")
+ _assert_eq(wpnt0, np.array([0.5, 0.5]), "wpnt0")
+ _assert_eq(wpnt1, np.array([0.5, 0.5]), "wpnt1")
+
+ wlen, wpnt0, wpnt1 = _wrap_inside(np.array([0, 0, 1, 0]), 1.0)
+ _assert_eq(wlen, -1.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1")
+
+ wlen, wpnt0, wpnt1 = _wrap_inside(np.array([1, 0, 0, 0]), 1.0)
+ _assert_eq(wlen, -1.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1")
+
+ wlen, wpnt0, wpnt1 = _wrap_inside(np.array([0, 0, 0, 0]), 1.0)
+ _assert_eq(wlen, -1.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1")
+
+ wlen, wpnt0, wpnt1 = _wrap_inside(np.array([1, 0, 0, 0]), 2.0)
+ _assert_eq(wlen, -1.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1")
+
+ wlen, wpnt0, wpnt1 = _wrap_inside(np.array([0, 0, 1, 0]), 2.0)
+ _assert_eq(wlen, -1.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1")
+
+ wlen, wpnt0, wpnt1 = _wrap_inside(np.array([1, 1, 1, 1]), 0.1 * MJ_MINVAL)
+ _assert_eq(wlen, -1.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1")
+
+ wlen, wpnt0, wpnt1 = _wrap_inside(np.array([-1, 0, 1, 0]), 0.1)
+ _assert_eq(wlen, -1.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1")
+
+ wlen, wpnt0, wpnt1 = _wrap_inside(np.array([-1, 0.2, 1, 0.2]), 0.1)
+ _assert_eq(wlen, 0.0, "wlen")
+ _assert_eq(wpnt0, np.array([0, 0.1]), "wpnt0")
+ _assert_eq(wpnt1, np.array([0, 0.1]), "wpnt1")
+
+ @parameterized.parameters(WrapType.SPHERE, WrapType.CYLINDER)
+ def test_wrap(self, wraptype):
+ # no wrap
+ x0 = np.array([1, 1, 1])
+ x1 = np.array([2, 2, 2])
+ xpos = np.array([0, 0, 0])
+ xmat = np.eye(3)
+ radius = 0.1
+ side = np.array([np.inf, np.inf, np.inf])
+
+ wlen, wpnt0, wpnt1 = _wrap(x0, x1, xpos, xmat, radius, wraptype, side)
+
+ _assert_eq(wlen, -1.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.inf, np.inf, np.inf]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.inf, np.inf, np.inf]), "wpnt1")
+
+ # wrap
+ x0 = np.array([0.1, -1.0, 0.0])
+ x1 = np.array([0.1, 1.0, 0.0])
+ xpos = np.array([0, 0, 0])
+ xmat = np.eye(3)
+ radius = 0.1
+ side = np.array([np.inf, np.inf, np.inf])
+
+ wlen, wpnt0, wpnt1 = _wrap(x0, x1, xpos, xmat, radius, wraptype, side)
+
+ _assert_eq(wlen, 0.0, "wlen")
+ _assert_eq(wpnt0, np.array([0.1, 0.0, 0.0]), "wpnt0")
+ _assert_eq(wpnt1, np.array([0.1, 0.0, 0.0]), "wpnt1")
+
+ # outside wrap w/ sidesite
+ x0 = np.array([MJ_MINVAL, -100.0, 0.0])
+ x1 = np.array([MJ_MINVAL, 100.0, 0.0])
+ xpos = np.array([0, 0, 0])
+ xmat = np.eye(3)
+ radius = 0.1
+ side = np.array([radius + 10 * MJ_MINVAL, 0, 0])
+
+ wlen, wpnt0, wpnt1 = _wrap(x0, x1, xpos, xmat, radius, wraptype, side)
+
+ # wlen, wpnt0[1], wpnt1[1] are ~0
+ _assert_eq(wlen, 0, "wlen")
+ _assert_eq(wpnt0, np.array([0.1, 0, 0]), "wpnt0")
+ _assert_eq(wpnt1, np.array([0.1, 0, 0]), "wpnt1")
+
+ # inside no wrap w/ sidesite
+ x0 = np.array([0.0, -1.0, 0.0])
+ x1 = np.array([0.0, 1.0, 0.0])
+ xpos = np.array([0, 0, 0])
+ xmat = np.eye(3)
+ radius = 0.1
+ wraptype = WrapType.CYLINDER
+ side = np.array([0, 0, 0])
+
+ wlen, wpnt0, wpnt1 = _wrap(x0, x1, xpos, xmat, radius, wraptype, side)
+
+ _assert_eq(wlen, -1.0, "wlen")
+ _assert_eq(wpnt0, np.array([np.inf, np.inf, np.inf]), "wpnt0")
+ _assert_eq(wpnt1, np.array([np.inf, np.inf, np.inf]), "wpnt1")
+
+ # inside wrap w/ sidesite
+ x0 = np.array([1.0, -1.0, 0.0])
+ x1 = np.array([1.0, 1.0, 0.0])
+ xpos = np.array([0, 0, 0])
+ xmat = np.eye(3)
+ radius = 0.1
+ side = np.array([0.0, 0.0, 0.0])
+
+ wlen, wpnt0, wpnt1 = _wrap(x0, x1, xpos, xmat, radius, wraptype, side)
+
+ _assert_eq(wlen, 0.0, "wlen")
+ _assert_eq(wpnt0, np.array([0.1, 0.0, 0.0]), "wpnt0")
+ _assert_eq(wpnt1, np.array([0.1, 0.0, 0.0]), "wpnt1")
+
+ @parameterized.product(ctrl=[-0.1, 0.0, 0.4, 0.5, 1.0, 1.1], act=[-0.1, 0.0, 0.4, 0.5, 1.0, 1.1])
+ def test_muscle_dynamics_tausmooth0(self, ctrl, act):
+ # exact equality if tau_smooth = 0
+ prm = np.array([0.01, 0.04, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0])
+
+ actdot_old = _muscle_dynamics_millard(ctrl, act, prm)
+ actdot_new = _muscle_dynamics(ctrl, act, prm)
+
+ _assert_eq(actdot_new, actdot_old, "actdot")
+
+ def test_muscle_dynamics_tausmooth_positive(self):
+ # positive tau_smooth
+ prm = np.array([0.01, 0.04, 0.2, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0])
+ act = 0.5
+ eps = 1.0e-6
+
+ ctrl = 0.4 - eps # smaller than act by just over 0.5 * tau_smooth
+ _assert_eq(
+ _muscle_dynamics(ctrl, act, prm),
+ _muscle_dynamics_millard(ctrl, act, prm),
+ "actdot",
+ )
+
+ ctrl = 0.6 + eps # larger than act by just over 0.5 * tau_smooth
+ _assert_eq(
+ _muscle_dynamics(ctrl, act, prm),
+ _muscle_dynamics_millard(ctrl, act, prm),
+ "actdot",
+ )
+
+ @parameterized.parameters(0.0, 0.1, 0.2, 1.0, 1.1)
+ def test_muscle_dynamics_timescale(self, dctrl):
+ # right in the middle should give average of time constants
+ tau_smooth = 0.2
+ tau_act = 0.2
+ tau_deact = 0.3
+
+ lower = _muscle_dynamics_timescale(-dctrl, tau_act, tau_deact, tau_smooth)
+ upper = _muscle_dynamics_timescale(dctrl, tau_act, tau_deact, tau_smooth)
+
+ _assert_eq(0.5 * (lower + upper), 0.5 * (tau_act + tau_deact), "muscle_dynamics_timescale")
+
+ @parameterized.parameters(
+ (0.0, 0.0),
+ (0.5, 0.0),
+ (0.75, 0.5),
+ (1.0, 1.0),
+ (1.25, 0.5),
+ (1.5, 0.0),
+ (2.0, 0.0),
+ )
+ def test_muscle_gain_length(self, input, output):
+ _assert_eq(_muscle_gain_length(input, 0.5, 1.5), output, "length-gain")
+
+ # TODO(team): test util_misc.muscle_gain
+ # TODO(team): test util_misc.muscle_bias
+
+
+if __name__ == "__main__":
+ wp.init()
+ absltest.main()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/warp_util.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/warp_util.py
new file mode 100644
index 00000000..572d2312
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/warp_util.py
@@ -0,0 +1,196 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+
+import functools
+from typing import Callable, Optional
+
+import warp as wp
+from warp.context import Module
+from warp.context import assert_conditional_graph_support
+from warp.context import get_module
+
+_STACK = None
+
+
+class EventTracer:
+ """Calculates elapsed times of functions annotated with `event_scope`.
+
+ Use as a context manager like so:
+
+ @event_trace
+ def my_warp_function(...):
+ ...
+
+ with EventTracer() as tracer:
+ my_warp_function(...)
+ print(tracer.trace())
+ """
+
+ def __init__(self, enabled: bool = True):
+ global _STACK
+ if _STACK is not None:
+ raise ValueError("only one EventTracer can run at a time")
+ if enabled:
+ _STACK = {}
+
+ def __enter__(self):
+ return self
+
+ def trace(self) -> dict:
+ """Calculates elapsed times for every node of the trace."""
+ global _STACK
+
+ if _STACK is None:
+ return {}
+
+ ret = {}
+
+ for k, v in _STACK.items():
+ events, sub_stack = v
+ # push into next level of stack
+ saved_stack, _STACK = _STACK, sub_stack
+ sub_trace = self.trace()
+ # pop!
+ _STACK = saved_stack
+ events = tuple(wp.get_event_elapsed_time(beg, end) for beg, end in events)
+ ret[k] = (events, sub_trace)
+
+ return ret
+
+ def __exit__(self, type, value, traceback):
+ global _STACK
+ _STACK = None
+
+
+def _merge(a: dict, b: dict) -> dict:
+ """Merges two event trace stacks."""
+ ret = {}
+ if not a or not b:
+ return dict(**a, **b)
+ if set(a) != set(b):
+ raise ValueError("incompatible stacks")
+ for key in a:
+ a1_events, a1_substack = a[key]
+ a2_events, a2_substack = b[key]
+ ret[key] = (a1_events + a2_events, _merge(a1_substack, a2_substack))
+ return ret
+
+
+def event_scope(fn, name: str = ""):
+ """Wraps a function and records an event before and after the function invocation."""
+ name = name or getattr(fn, "__name__")
+
+ @functools.wraps(fn)
+ def wrapper(*args, **kwargs):
+ global _STACK
+ if _STACK is None:
+ return fn(*args, **kwargs)
+ # push into next level of stack
+ saved_stack, _STACK = _STACK, {}
+ beg = wp.Event(enable_timing=True)
+ end = wp.Event(enable_timing=True)
+ wp.record_event(beg)
+ res = fn(*args, **kwargs)
+ wp.record_event(end)
+ # pop back up to current level
+ sub_stack, _STACK = _STACK, saved_stack
+ # append events and substack
+ prev_events, prev_substack = _STACK.get(name, ((), {}))
+ events = prev_events + ((beg, end),)
+ sub_stack = _merge(prev_substack, sub_stack)
+ _STACK[name] = (events, sub_stack)
+ return res
+
+ return wrapper
+
+
+# @kernel decorator to automatically set up modules based on nested
+# function names
+def kernel(
+ f: Optional[Callable] = None,
+ *,
+ enable_backward: Optional[bool] = None,
+ module: Optional[Module] = None,
+):
+ """
+ Decorator to register a Warp kernel from a Python function.
+ The function must be defined with type annotations for all arguments.
+ The function must not return anything.
+
+ Example::
+
+ @kernel
+ def my_kernel(a: wp.array(dtype=float), b: wp.array(dtype=float)):
+ tid = wp.tid()
+ b[tid] = a[tid] + 1.0
+
+
+ @kernel(enable_backward=False)
+ def my_kernel_no_backward(a: wp.array(dtype=float, ndim=2), x: float):
+ # the backward pass will not be generated
+ i, j = wp.tid()
+ a[i, j] = x
+
+
+ @kernel(module="unique")
+ def my_kernel_unique_module(a: wp.array(dtype=float), b: wp.array(dtype=float)):
+ # the kernel will be registered in new unique module created just for this
+ # kernel and its dependent functions and structs
+ tid = wp.tid()
+ b[tid] = a[tid] + 1.0
+
+ Args:
+ f: The function to be registered as a kernel.
+ enable_backward: If False, the backward pass will not be generated.
+ module: The :class:`warp.context.Module` to which the kernel belongs. Alternatively,
+ if a string `"unique"` is provided, the kernel is assigned to a new module
+ named after the kernel name and hash. If None, the module is inferred from
+ the function's module.
+
+ Returns:
+ The registered kernel.
+ """
+ if module is None:
+ # create a module name based on the name of the nested function
+ # get the qualified name, e.g. "main..nested_kernel"
+ qualname = f.__qualname__
+ parts = [part for part in qualname.split(".") if part != ""]
+ outer_functions = parts[:-1]
+ module = get_module(".".join([f.__module__] + outer_functions))
+
+ return wp.kernel(f, enable_backward=enable_backward, module=module)
+
+
+_KERNEL_CACHE = {}
+
+
+def cache_kernel(func):
+ # caching kernels to avoid crashes in graph_conditional code
+ @functools.wraps(func)
+ def wrapper(*args):
+ key = tuple(a.size if hasattr(a, "size") else hash(a) for a in args) + (hash(func.__name__),)
+ if key not in _KERNEL_CACHE:
+ _KERNEL_CACHE[key] = func(*args)
+ return _KERNEL_CACHE[key]
+
+ return wrapper
+
+
+def conditional_graph_supported():
+ try:
+ assert_conditional_graph_support()
+ except Exception:
+ return False
+ return True
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuation.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuation.xml
new file mode 100644
index 00000000..2e532171
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuation.xml
@@ -0,0 +1,40 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuators.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuators.xml
new file mode 100644
index 00000000..24c86ae9
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuators.xml
@@ -0,0 +1,28 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/adhesion.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/adhesion.xml
new file mode 100644
index 00000000..fb891b39
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/adhesion.xml
@@ -0,0 +1,21 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/muscle.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/muscle.xml
new file mode 100644
index 00000000..b07e7f1a
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/muscle.xml
@@ -0,0 +1,22 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/position.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/position.xml
new file mode 100644
index 00000000..16e9d7a6
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/position.xml
@@ -0,0 +1,40 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/site.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/site.xml
new file mode 100644
index 00000000..fcff6eea
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/site.xml
@@ -0,0 +1,31 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/slidercrank.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/slidercrank.xml
new file mode 100644
index 00000000..75a34153
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/slidercrank.xml
@@ -0,0 +1,15 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/tendon_force_limit.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/tendon_force_limit.xml
new file mode 100644
index 00000000..c7c0ed8b
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/tendon_force_limit.xml
@@ -0,0 +1,26 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision.xml
new file mode 100644
index 00000000..5cc041e2
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision.xml
@@ -0,0 +1,43 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/bolt.py b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/bolt.py
new file mode 100644
index 00000000..07a99a1e
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/bolt.py
@@ -0,0 +1,70 @@
+import warp as wp
+
+
+@wp.func
+def Fract(x: float) -> float:
+ return x - wp.floor(x)
+
+
+@wp.func
+def Subtraction(a: float, b: float) -> float:
+ return wp.max(a, -b)
+
+
+@wp.func
+def Union(a: float, b: float) -> float:
+ return wp.min(a, b)
+
+
+@wp.func
+def Intersection(a: float, b: float) -> float:
+ return wp.max(a, b)
+
+
+@wp.func
+def bolt(p: wp.vec3, attr: wp.vec3) -> float:
+ screw = 12.0
+ radius = wp.sqrt(p[0] * p[0] + p[1] * p[1]) - attr[0]
+ sqrt12 = wp.sqrt(2.0) / 2.0
+
+ azimuth = wp.atan2(p[1], p[0])
+ triangle = wp.abs(Fract(p[2] * screw - azimuth / wp.pi / 2.0) - 0.5)
+ thread = (radius - triangle / screw) * sqrt12
+
+ bolt_val = Subtraction(thread, 0.5 - wp.abs(p[2] + 0.5))
+ cone = (p[2] - radius) * sqrt12
+
+ bolt_val = Subtraction(bolt_val, cone + 1.0 * sqrt12)
+
+ point2D = wp.vec2(p[0], p[1])
+ k = 6.0 / wp.pi / 2.0
+ angle = -wp.floor((wp.atan2(point2D[1], point2D[0])) * k + 0.5) / k
+ s = wp.vec2(wp.sin(angle), wp.sin(angle + wp.pi * 0.5))
+
+ res = wp.vec2(s[1] * point2D[0] - s[0] * point2D[1], s[0] * point2D[0] + s[1] * point2D[1])
+ point3D = wp.vec3(res[0], res[1], p[2])
+ head = point3D[0] - 0.5
+
+ head = Intersection(head, wp.abs(point3D[2] + 0.25) - 0.25)
+ head = Intersection(head, (point3D[2] + radius - 0.22) * sqrt12)
+
+ return Union(bolt_val, head)
+
+
+@wp.func
+def bolt_sdf_grad(p: wp.vec3, attr: wp.vec3) -> wp.vec3:
+ grad = wp.vec3()
+ eps = 1e-6
+ f_original = bolt(p, attr)
+ x_plus = wp.vec3(p[0] + eps, p[1], p[2])
+ f_plus = bolt(x_plus, attr)
+ grad[0] = (f_plus - f_original) / eps
+
+ x_plus = wp.vec3(p[0], p[1] + eps, p[2])
+ f_plus = bolt(x_plus, attr)
+ grad[1] = (f_plus - f_original) / eps
+
+ x_plus = wp.vec3(p[0], p[1], p[2] + eps)
+ f_plus = bolt(x_plus, attr)
+ grad[2] = (f_plus - f_original) / eps
+ return grad
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nut.py b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nut.py
new file mode 100644
index 00000000..47b8cc73
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nut.py
@@ -0,0 +1,65 @@
+import warp as wp
+
+
+@wp.func
+def Fract(x: float) -> float:
+ return x - wp.floor(x)
+
+
+@wp.func
+def Subtraction(a: float, b: float) -> float:
+ return wp.max(a, -b)
+
+
+@wp.func
+def Union(a: float, b: float) -> float:
+ return wp.min(a, b)
+
+
+@wp.func
+def Intersection(a: float, b: float) -> float:
+ return wp.max(a, b)
+
+
+@wp.func
+def nut(p: wp.vec3, attr: wp.vec3) -> float:
+ screw = 12.0
+ radius2 = wp.sqrt(p[0] * p[0] + p[1] * p[1]) - attr[0]
+ sqrt12 = wp.sqrt(2.0) / 2.0
+ azimuth = wp.atan2(p[1], p[0])
+ triangle = wp.abs(Fract(p[2] * screw - azimuth / (wp.pi * 2.0)) - 0.5)
+ thread2 = (radius2 - triangle / screw) * sqrt12
+ cone2 = (p[2] - radius2) * sqrt12
+ hole = Subtraction(thread2, cone2 + 0.5 * sqrt12)
+ hole = Union(hole, -cone2 - 0.05 * sqrt12)
+ k = 6.0 / wp.pi / 2.0
+ angle = -wp.floor((wp.atan2(p[1], p[0])) * k + 0.5) / k
+ s0 = wp.sin(angle)
+ s1 = wp.sin(angle + wp.pi * 0.5)
+ res0 = s1 * p[0] - s0 * p[1]
+ res1 = s0 * p[0] + s1 * p[1]
+ point3D0 = res0
+ point3D2 = p[2]
+ head = point3D0 - 0.5
+ head = Intersection(head, wp.abs(point3D2 + 0.25) - 0.25)
+ head = Intersection(head, (point3D2 + radius2 - 0.22) * sqrt12)
+ return Subtraction(head, hole)
+
+
+@wp.func
+def nut_sdf_grad(p: wp.vec3, attr: wp.vec3) -> wp.vec3:
+ grad = wp.vec3()
+ eps = 1e-6
+ f_original = nut(p, attr)
+ x_plus = wp.vec3(p[0] + eps, p[1], p[2])
+ f_plus = nut(x_plus, attr)
+ grad[0] = (f_plus - f_original) / eps
+
+ x_plus = wp.vec3(p[0], p[1] + eps, p[2])
+ f_plus = nut(x_plus, attr)
+ grad[1] = (f_plus - f_original) / eps
+
+ x_plus = wp.vec3(p[0], p[1], p[2] + eps)
+ f_plus = nut(x_plus, attr)
+ grad[2] = (f_plus - f_original) / eps
+ return grad
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nutbolt.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nutbolt.xml
new file mode 100644
index 00000000..d0fa2d26
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nutbolt.xml
@@ -0,0 +1,56 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/scene.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/scene.xml
new file mode 100644
index 00000000..1591519e
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/scene.xml
@@ -0,0 +1,38 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/utils.py b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/utils.py
new file mode 100644
index 00000000..76a12f89
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/utils.py
@@ -0,0 +1,86 @@
+# Copyright 2025 The Newton Developers
+#
+# 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.
+# ==============================================================================
+"""Utility functions for SDF collision handling."""
+
+import enum
+from typing import Dict
+
+import mujoco
+import warp as wp
+
+from .bolt import bolt
+from .bolt import bolt_sdf_grad
+from .nut import nut
+from .nut import nut_sdf_grad
+
+
+class SDFType(enum.Enum):
+ """Enum for SDF types."""
+
+ NUT = "NUT"
+ BOLT = "BOLT"
+
+
+def register_sdf_plugins(collision_sdf) -> Dict[str, int]:
+ xml = """
+
+
+
+
+
+
+
+
+
+
+
+
+ """
+
+ try:
+ m = mujoco.MjModel.from_xml_string(xml)
+ except Exception as e:
+ raise ValueError(f"Failed to create MuJoCo model from XML: {e}")
+
+ sdf_types = {}
+
+ for i in range(m.ngeom):
+ name = mujoco.mj_id2name(m, mujoco.mjtObj.mjOBJ_GEOM, i)
+ if name == "ng":
+ sdf_types[SDFType.NUT.value] = int(m.plugin[i])
+ elif name == "bg":
+ sdf_types[SDFType.BOLT.value] = int(m.plugin[i])
+
+ @wp.func
+ def user_sdf(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> float:
+ result = 0.0
+ if sdf_type == wp.static(sdf_types[SDFType.NUT.value]):
+ result = nut(p, attr)
+ elif sdf_type == wp.static(sdf_types[SDFType.BOLT.value]):
+ result = bolt(p, attr)
+ return result
+
+ @wp.func
+ def user_sdf_grad(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> wp.vec3:
+ if sdf_type == wp.static(sdf_types[SDFType.NUT.value]):
+ return nut_sdf_grad(p, attr)
+ elif sdf_type == wp.static(sdf_types[SDFType.BOLT.value]):
+ return bolt_sdf_grad(p, attr)
+ return wp.vec3()
+
+ collision_sdf.user_sdf = user_sdf
+ collision_sdf.user_sdf_grad = user_sdf_grad
+
+ return sdf_types
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/constraints.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/constraints.xml
new file mode 100644
index 00000000..8e5be10f
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/constraints.xml
@@ -0,0 +1,114 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/cloth.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/cloth.xml
new file mode 100644
index 00000000..89e0adfa
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/cloth.xml
@@ -0,0 +1,45 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/floppy.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/floppy.xml
new file mode 100644
index 00000000..6430e70d
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/floppy.xml
@@ -0,0 +1,45 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/mannequin.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/mannequin.xml
new file mode 100644
index 00000000..9237b943
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/mannequin.xml
@@ -0,0 +1,174 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/scene.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/scene.xml
new file mode 100644
index 00000000..ef314930
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/scene.xml
@@ -0,0 +1,41 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/hfield/hfield.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/hfield/hfield.xml
new file mode 100644
index 00000000..27d179ee
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/hfield/hfield.xml
@@ -0,0 +1,84 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
\ No newline at end of file
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/humanoid/humanoid.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/humanoid/humanoid.xml
new file mode 100644
index 00000000..196aa5c9
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/humanoid/humanoid.xml
@@ -0,0 +1,252 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/meshes/dodecahedron.stl b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/meshes/dodecahedron.stl
new file mode 100644
index 00000000..1a3f9f69
Binary files /dev/null and b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/meshes/dodecahedron.stl differ
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/meshes/tetrahedron.stl b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/meshes/tetrahedron.stl
new file mode 100644
index 00000000..68fed6b8
Binary files /dev/null and b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/meshes/tetrahedron.stl differ
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/pendula.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/pendula.xml
new file mode 100644
index 00000000..b3ea0453
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/pendula.xml
@@ -0,0 +1,167 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/ray.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/ray.xml
new file mode 100644
index 00000000..46bc3070
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/ray.xml
@@ -0,0 +1,21 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/armature.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/armature.xml
new file mode 100644
index 00000000..6ab99fda
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/armature.xml
@@ -0,0 +1,19 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/damping.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/damping.xml
new file mode 100644
index 00000000..e4a6d380
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/damping.xml
@@ -0,0 +1,20 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed.xml
new file mode 100644
index 00000000..42b6a164
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed.xml
@@ -0,0 +1,29 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed_site.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed_site.xml
new file mode 100644
index 00000000..261b900b
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed_site.xml
@@ -0,0 +1,35 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_fixed_site.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_fixed_site.xml
new file mode 100644
index 00000000..25b3cab5
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_fixed_site.xml
@@ -0,0 +1,36 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site.xml
new file mode 100644
index 00000000..0b4b351c
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site.xml
@@ -0,0 +1,46 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site_fixed.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site_fixed.xml
new file mode 100644
index 00000000..8d7bf020
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site_fixed.xml
@@ -0,0 +1,32 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_wrap.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_wrap.xml
new file mode 100644
index 00000000..1dff6600
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_wrap.xml
@@ -0,0 +1,42 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site.xml
new file mode 100644
index 00000000..da50fe17
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site.xml
@@ -0,0 +1,44 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site_fixed.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site_fixed.xml
new file mode 100644
index 00000000..af85b662
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site_fixed.xml
@@ -0,0 +1,31 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/tendon_limit.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/tendon_limit.xml
new file mode 100644
index 00000000..e6f243c8
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/tendon_limit.xml
@@ -0,0 +1,29 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/wrap.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/wrap.xml
new file mode 100644
index 00000000..ae9d19a1
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/wrap.xml
@@ -0,0 +1,40 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py
new file mode 100644
index 00000000..89044207
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py
@@ -0,0 +1,16 @@
+# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+#
+# 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.
+
+from .custom_call import jax_kernel
diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py
new file mode 100644
index 00000000..97faaec5
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py
@@ -0,0 +1,363 @@
+# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+#
+# 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.
+
+import ctypes
+
+import warp as wp
+from warp.context import type_str
+from warp.jax import get_jax_device
+from warp.types import array_t, launch_bounds_t, strides_from_shape
+
+_jax_warp_p = None
+
+# Holder for the custom callback to keep it alive.
+_cc_callback = None
+_registered_kernels = [None]
+_registered_kernel_to_id = {}
+
+
+def jax_kernel(kernel, launch_dims=None):
+ """Create a Jax primitive from a Warp kernel.
+
+ NOTE: This is an experimental feature under development.
+
+ Args:
+ kernel: The Warp kernel to be wrapped.
+ launch_dims: Optional. Specify the kernel launch dimensions. If None,
+ dimensions are inferred from the shape of the first argument.
+ This option when set will specify the output dimensions.
+
+ Limitations:
+ - All kernel arguments must be contiguous arrays.
+ - Input arguments are followed by output arguments in the Warp kernel definition.
+ - There must be at least one input argument and at least one output argument.
+ - Only the CUDA backend is supported.
+ """
+
+ if _jax_warp_p is None:
+ # Create and register the primitive
+ _create_jax_warp_primitive()
+ if kernel not in _registered_kernel_to_id:
+ id = len(_registered_kernels)
+ _registered_kernels.append(kernel)
+ _registered_kernel_to_id[kernel] = id
+ else:
+ id = _registered_kernel_to_id[kernel]
+
+ def bind(*args):
+ return _jax_warp_p.bind(*args, kernel=id, launch_dims=launch_dims)
+
+ return bind
+
+
+def _warp_custom_callback(stream, buffers, opaque, opaque_len):
+ # The descriptor is the form
+ # ||
+ # Example: 42|16,32|16,32;100;16,32
+ kernel_id_str, dim_str, args_str = opaque.decode().split("|")
+
+ # Get the kernel from the registry.
+ kernel_id = int(kernel_id_str)
+ kernel = _registered_kernels[kernel_id]
+
+ # Parse launch dimensions.
+ dims = [int(d) for d in dim_str.split(",")]
+ bounds = launch_bounds_t(dims)
+
+ # Parse arguments.
+ arg_strings = args_str.split(";")
+ num_args = len(arg_strings)
+ assert num_args == len(kernel.adj.args), "Incorrect number of arguments"
+
+ # First param is the launch bounds.
+ kernel_params = (ctypes.c_void_p * (1 + num_args))()
+ kernel_params[0] = ctypes.addressof(bounds)
+
+ # Parse array descriptors.
+ args = []
+ for i in range(num_args):
+ dtype = kernel.adj.args[i].type.dtype
+ shape = [int(d) for d in arg_strings[i].split(",")]
+ strides = strides_from_shape(shape, dtype)
+
+ arr = array_t(buffers[i], 0, len(shape), shape, strides)
+ args.append(arr) # keep a reference
+ arg_ptr = ctypes.addressof(arr)
+
+ kernel_params[i + 1] = arg_ptr
+
+ # Get current device.
+ device = wp.device_from_jax(get_jax_device())
+
+ # Get kernel hooks.
+ # Note: module was loaded during jit lowering.
+ hooks = kernel.module.get_kernel_hooks(kernel, device)
+ assert hooks.forward, "Failed to find kernel entry point"
+
+ # Launch the kernel.
+ wp.context.runtime.core.wp_cuda_launch_kernel(
+ device.context, hooks.forward, bounds.size, 0, 256, hooks.forward_smem_bytes, kernel_params, stream
+ )
+
+
+def _create_jax_warp_primitive():
+ from functools import reduce
+
+ import jax
+ from jax._src.interpreters import batching
+ from jax.interpreters import mlir
+ from jax.interpreters.mlir import ir
+ from jax.jaxlib.hlo_helpers import custom_call
+
+ global _jax_warp_p
+ global _cc_callback
+
+ # Create and register the primitive.
+ # TODO add default implementation that calls the kernel via warp.
+ try:
+ # newer JAX versions
+ import jax.extend
+
+ _jax_warp_p = jax.extend.core.Primitive("jax_warp")
+ except (ImportError, AttributeError):
+ # older JAX versions
+ _jax_warp_p = jax.core.Primitive("jax_warp")
+ _jax_warp_p.multiple_results = True
+
+ # TODO Just launch the kernel directly, but make sure the argument
+ # shapes are massaged the same way as below so that vmap works.
+ def impl(*args):
+ raise Exception("Not implemented")
+
+ _jax_warp_p.def_impl(impl)
+
+ # Auto-batching. Make sure all the arguments are fully broadcasted
+ # so that Warp is not confused about dimensions.
+ def vectorized_multi_batcher(args, dims, **params):
+ # Figure out the number of outputs.
+ wp_kernel = _registered_kernels[params["kernel"]]
+ output_count = len(wp_kernel.adj.args) - len(args)
+ shape, dim = next((a.shape, d) for a, d in zip(args, dims) if d is not None)
+ size = shape[dim]
+ args = [batching.bdim_at_front(a, d, size) if len(a.shape) else a for a, d in zip(args, dims)]
+ # Create the batched primitive.
+ return _jax_warp_p.bind(*args, **params), [dims[0]] * output_count
+
+ batching.primitive_batchers[_jax_warp_p] = vectorized_multi_batcher
+
+ def get_vecmat_shape(warp_type):
+ if hasattr(warp_type.dtype, "_shape_"):
+ return warp_type.dtype._shape_
+ return []
+
+ def strip_vecmat_dimensions(warp_arg, actual_shape):
+ shape = get_vecmat_shape(warp_arg.type)
+ for i, s in enumerate(reversed(shape)):
+ item = actual_shape[-i - 1]
+ if s != item:
+ raise Exception(f"The vector/matrix shape for argument {warp_arg.label} does not match")
+ return actual_shape[: len(actual_shape) - len(shape)]
+
+ def collapse_into_leading_dimension(warp_arg, actual_shape):
+ if len(actual_shape) < warp_arg.type.ndim:
+ raise Exception(f"Argument {warp_arg.label} has too few non-matrix/vector dimensions")
+ index_rest = len(actual_shape) - warp_arg.type.ndim + 1
+ leading_size = reduce(lambda x, y: x * y, actual_shape[:index_rest])
+ return [leading_size] + actual_shape[index_rest:]
+
+ # Infer array dimensions from input type.
+ def infer_dimensions(warp_arg, actual_shape):
+ actual_shape = strip_vecmat_dimensions(warp_arg, actual_shape)
+ return collapse_into_leading_dimension(warp_arg, actual_shape)
+
+ def base_type_to_jax(warp_dtype):
+ if hasattr(warp_dtype, "_wp_scalar_type_"):
+ return wp.dtype_to_jax(warp_dtype._wp_scalar_type_)
+ return wp.dtype_to_jax(warp_dtype)
+
+ def base_type_to_jax_ir(warp_dtype):
+ warp_to_jax_dict = {
+ wp.float16: ir.F16Type.get(),
+ wp.float32: ir.F32Type.get(),
+ wp.float64: ir.F64Type.get(),
+ wp.int8: ir.IntegerType.get_signless(8),
+ wp.int16: ir.IntegerType.get_signless(16),
+ wp.int32: ir.IntegerType.get_signless(32),
+ wp.int64: ir.IntegerType.get_signless(64),
+ wp.uint8: ir.IntegerType.get_unsigned(8),
+ wp.uint16: ir.IntegerType.get_unsigned(16),
+ wp.uint32: ir.IntegerType.get_unsigned(32),
+ wp.uint64: ir.IntegerType.get_unsigned(64),
+ }
+ if hasattr(warp_dtype, "_wp_scalar_type_"):
+ warp_dtype = warp_dtype._wp_scalar_type_
+ jax_dtype = warp_to_jax_dict.get(warp_dtype)
+ if jax_dtype is None:
+ raise TypeError(f"Invalid or unsupported data type: {warp_dtype}")
+ return jax_dtype
+
+ def base_type_is_compatible(warp_type, jax_ir_type):
+ jax_ir_to_warp = {
+ "f16": wp.float16,
+ "f32": wp.float32,
+ "f64": wp.float64,
+ "i8": wp.int8,
+ "i16": wp.int16,
+ "i32": wp.int32,
+ "i64": wp.int64,
+ "ui8": wp.uint8,
+ "ui16": wp.uint16,
+ "ui32": wp.uint32,
+ "ui64": wp.uint64,
+ }
+ expected_warp_type = jax_ir_to_warp.get(str(jax_ir_type))
+ if expected_warp_type is not None:
+ if hasattr(warp_type, "_wp_scalar_type_"):
+ return warp_type._wp_scalar_type_ == expected_warp_type
+ else:
+ return warp_type == expected_warp_type
+ else:
+ raise TypeError(f"Invalid or unsupported data type: {jax_ir_type}")
+
+ # Abstract evaluation.
+ def jax_warp_abstract(*args, kernel=None, launch_dims=None):
+ wp_kernel = _registered_kernels[kernel]
+ # All the extra arguments to the warp kernel are outputs.
+ warp_outputs = [o.type for o in wp_kernel.adj.args[len(args) :]]
+
+ if launch_dims is None:
+ # Use the first input dimension to infer the output's dimensions if launch_dims is not provided
+ dims = strip_vecmat_dimensions(wp_kernel.adj.args[0], list(args[0].shape))
+ else:
+ dims = launch_dims
+
+ jax_outputs = []
+ for o in warp_outputs:
+ shape = list(dims) + list(get_vecmat_shape(o))
+ dtype = base_type_to_jax(o.dtype)
+ jax_outputs.append(jax.core.ShapedArray(shape, dtype))
+ return jax_outputs
+
+ _jax_warp_p.def_abstract_eval(jax_warp_abstract)
+
+ # Lowering to MLIR.
+
+ # Create python-land custom call target.
+ CCALLFUNC = ctypes.CFUNCTYPE(
+ ctypes.c_voidp, ctypes.c_void_p, ctypes.POINTER(ctypes.c_void_p), ctypes.c_char_p, ctypes.c_size_t
+ )
+ _cc_callback = CCALLFUNC(_warp_custom_callback)
+ ccall_address = ctypes.cast(_cc_callback, ctypes.c_void_p)
+
+ # Put the custom call into a capsule, as required by XLA.
+ PyCapsule_Destructor = ctypes.CFUNCTYPE(None, ctypes.py_object)
+ PyCapsule_New = ctypes.pythonapi.PyCapsule_New
+ PyCapsule_New.restype = ctypes.py_object
+ PyCapsule_New.argtypes = (ctypes.c_void_p, ctypes.c_char_p, PyCapsule_Destructor)
+ capsule = PyCapsule_New(ccall_address.value, b"xla._CUSTOM_CALL_TARGET", PyCapsule_Destructor(0))
+
+ # Register the callback in XLA.
+ try:
+ # newer JAX versions
+ jax.ffi.register_ffi_target("warp_call", capsule, platform="gpu", api_version=0)
+ except AttributeError:
+ # older JAX versions
+ jax.lib.xla_client.register_custom_call_target("warp_call", capsule, platform="gpu")
+
+ def default_layout(shape):
+ return range(len(shape) - 1, -1, -1)
+
+ def warp_call_lowering(ctx, *args, kernel=None, launch_dims=None):
+ if not kernel:
+ raise Exception("Unknown kernel id " + str(kernel))
+ wp_kernel = _registered_kernels[kernel]
+
+ # TODO This may not be necessary, but it is perhaps better not to be
+ # mucking with kernel loading while already running the workload.
+ module = wp_kernel.module
+ device = wp.device_from_jax(get_jax_device())
+ if not module.load(device):
+ raise Exception("Could not load kernel on device")
+
+ if launch_dims is None:
+ # Infer dimensions from the first input.
+ warp_arg0 = wp_kernel.adj.args[0]
+ actual_shape0 = ir.RankedTensorType(args[0].type).shape
+ dims = strip_vecmat_dimensions(warp_arg0, actual_shape0)
+ warp_dims = collapse_into_leading_dimension(warp_arg0, dims)
+ else:
+ dims = launch_dims
+ warp_dims = launch_dims
+ # Figure out the types and shapes of the input arrays.
+ arg_strings = []
+ operand_layouts = []
+ for actual, warg in zip(args, wp_kernel.adj.args):
+ wtype = warg.type
+ rtt = ir.RankedTensorType(actual.type)
+
+ if not isinstance(wtype, wp.array):
+ raise Exception("Only contiguous arrays are supported for Jax kernel arguments")
+
+ if not base_type_is_compatible(wtype.dtype, rtt.element_type):
+ raise TypeError(
+ f"Incompatible data type for argument '{warg.label}', expected {type_str(wtype.dtype)}, got {rtt.element_type}"
+ )
+
+ # Infer array dimension (by removing the vector/matrix dimensions and
+ # collapsing the initial dimensions).
+ shape = infer_dimensions(warg, rtt.shape)
+
+ if len(shape) != wtype.ndim:
+ raise TypeError(f"Incompatible array dimensionality for argument '{warg.label}'")
+
+ arg_strings.append(",".join([str(d) for d in shape]))
+ operand_layouts.append(default_layout(rtt.shape))
+
+ # Figure out the types and shapes of the output arrays.
+ result_types = []
+ result_layouts = []
+ for warg in wp_kernel.adj.args[len(args) :]:
+ wtype = warg.type
+
+ if not isinstance(wtype, wp.array):
+ raise Exception("Only contiguous arrays are supported for Jax kernel arguments")
+
+ # Infer dimensions from the first input.
+ arg_strings.append(",".join([str(d) for d in warp_dims]))
+
+ result_shape = list(dims) + list(get_vecmat_shape(wtype))
+ result_types.append(ir.RankedTensorType.get(result_shape, base_type_to_jax_ir(wtype.dtype)))
+ result_layouts.append(default_layout(result_shape))
+
+ # Build opaque descriptor for callback.
+ shape_str = ",".join([str(d) for d in warp_dims])
+ args_str = ";".join(arg_strings)
+ descriptor = f"{kernel}|{shape_str}|{args_str}"
+
+ out = custom_call(
+ b"warp_call",
+ result_types=result_types,
+ operands=args,
+ backend_config=descriptor.encode("utf-8"),
+ operand_layouts=operand_layouts,
+ result_layouts=result_layouts,
+ ).results
+ return out
+
+ mlir.register_lowering(
+ _jax_warp_p,
+ warp_call_lowering,
+ platform="gpu",
+ )
diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py
new file mode 100644
index 00000000..82b6a8b8
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py
@@ -0,0 +1,804 @@
+# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+#
+# 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.
+
+import ctypes
+import threading
+import traceback
+from typing import Callable, Optional
+
+import jax
+
+import warp as wp
+from warp.codegen import get_full_arg_spec, make_full_qualified_name
+from warp.jax import get_jax_device
+from warp.types import array_t, launch_bounds_t, strides_from_shape, type_to_warp
+
+from .xla_ffi import *
+
+
+class FfiArg:
+ def __init__(self, name, type, in_out=False):
+ self.name = name
+ self.type = type
+ self.in_out = in_out
+ self.is_array = isinstance(type, wp.array)
+
+ if self.is_array:
+ if hasattr(type.dtype, "_wp_scalar_type_"):
+ self.dtype_shape = type.dtype._shape_
+ self.dtype_ndim = len(self.dtype_shape)
+ self.jax_scalar_type = wp.dtype_to_jax(type.dtype._wp_scalar_type_)
+ self.jax_ndim = type.ndim + self.dtype_ndim
+ elif type.dtype in wp.types.value_types:
+ self.dtype_ndim = 0
+ self.dtype_shape = ()
+ self.jax_scalar_type = wp.dtype_to_jax(type.dtype)
+ self.jax_ndim = type.ndim
+ else:
+ raise TypeError(f"Invalid data type for array argument '{name}', expected scalar, vector, or matrix")
+ self.warp_ndim = type.ndim
+ elif type in wp.types.value_types:
+ self.dtype_ndim = 0
+ self.dtype_shape = ()
+ self.jax_scalar_type = wp.dtype_to_jax(type_to_warp(type))
+ self.jax_ndim = 0
+ self.warp_ndim = 0
+ else:
+ raise TypeError(f"Invalid type for argument '{name}', expected array or scalar, got {type}")
+
+
+class FfiLaunchDesc:
+ def __init__(self, static_inputs, launch_dims):
+ self.static_inputs = static_inputs
+ self.launch_dims = launch_dims
+
+
+class FfiKernel:
+ def __init__(self, kernel, num_outputs, vmap_method, launch_dims, output_dims, in_out_argnames):
+ self.kernel = kernel
+ self.name = generate_unique_name(kernel.func)
+ self.num_outputs = num_outputs
+ self.vmap_method = vmap_method
+ self.launch_dims = launch_dims
+ self.output_dims = output_dims
+ self.first_array_arg = None
+ self.launch_id = 0
+ self.launch_descriptors = {}
+
+ in_out_argnames_list = in_out_argnames or []
+ in_out_argnames = set(in_out_argnames_list)
+ if len(in_out_argnames_list) != len(in_out_argnames):
+ raise AssertionError("in_out_argnames must not contain duplicate names")
+
+ self.num_kernel_args = len(kernel.adj.args)
+ self.num_in_out = len(in_out_argnames)
+ self.num_inputs = self.num_kernel_args - num_outputs + self.num_in_out
+ if self.num_outputs < 1:
+ raise ValueError("At least one output is required")
+ if self.num_outputs > self.num_kernel_args:
+ raise ValueError("Number of outputs cannot be greater than the number of kernel arguments")
+ if self.num_outputs < self.num_in_out:
+ raise ValueError("Number of outputs cannot be smaller than the number of in_out_argnames")
+
+ # process input args
+ self.input_args = []
+ for i in range(self.num_inputs):
+ arg_name = kernel.adj.args[i].label
+ arg = FfiArg(arg_name, kernel.adj.args[i].type, arg_name in in_out_argnames)
+ if arg_name in in_out_argnames:
+ in_out_argnames.remove(arg_name)
+ if arg.is_array:
+ # keep track of the first input array argument
+ if self.first_array_arg is None:
+ self.first_array_arg = i
+ self.input_args.append(arg)
+
+ # process output args
+ self.output_args = []
+ for i in range(self.num_inputs, self.num_kernel_args):
+ arg_name = kernel.adj.args[i].label
+ if arg_name in in_out_argnames:
+ raise AssertionError(
+ f"Expected an output-only argument for argument {arg_name}."
+ " in_out arguments should be placed before output-only arguments."
+ )
+ arg = FfiArg(arg_name, kernel.adj.args[i].type, False)
+ if not arg.is_array:
+ raise TypeError("All output arguments must be arrays")
+ self.output_args.append(arg)
+
+ if in_out_argnames:
+ raise ValueError(f"in_out_argnames: '{in_out_argnames}' did not match any function argument names.")
+
+ # Build input output aliases.
+ out_id = 0
+ input_output_aliases = {}
+ for in_id, arg in enumerate(self.input_args):
+ if not arg.in_out:
+ continue
+ input_output_aliases[in_id] = out_id
+ out_id += 1
+ self.input_output_aliases = input_output_aliases
+
+ # register the callback
+ FFI_CCALLFUNC = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame))
+ self.callback_func = FFI_CCALLFUNC(lambda call_frame: self.ffi_callback(call_frame))
+ ffi_ccall_address = ctypes.cast(self.callback_func, ctypes.c_void_p)
+ ffi_capsule = jax.ffi.pycapsule(ffi_ccall_address.value)
+ jax.ffi.register_ffi_target(self.name, ffi_capsule, platform="CUDA")
+
+ def __call__(self, *args, output_dims=None, launch_dims=None, vmap_method=None):
+ num_inputs = len(args)
+ if num_inputs != self.num_inputs:
+ raise ValueError(f"Expected {self.num_inputs} inputs, but got {num_inputs}")
+
+ # default argument fallback
+ if launch_dims is None:
+ launch_dims = self.launch_dims
+ if output_dims is None:
+ output_dims = self.output_dims
+ if vmap_method is None:
+ vmap_method = self.vmap_method
+
+ # output types
+ out_types = []
+
+ # process inputs
+ static_inputs = {}
+ for i in range(num_inputs):
+ input_arg = self.input_args[i]
+ input_value = args[i]
+ if input_arg.is_array:
+ # check dtype
+ if input_value.dtype != input_arg.jax_scalar_type:
+ raise TypeError(
+ f"Invalid data type for array argument '{input_arg.name}', expected {input_arg.jax_scalar_type}, got {input_value.dtype}"
+ )
+ # check ndim
+ if input_value.ndim != input_arg.jax_ndim:
+ raise TypeError(
+ f"Invalid dimensionality for array argument '{input_arg.name}', expected {input_arg.jax_ndim} dimensions, got {input_value.ndim}"
+ )
+ # check inner dims
+ for d in range(input_arg.dtype_ndim):
+ if input_value.shape[input_arg.type.ndim + d] != input_arg.dtype_shape[d]:
+ raise TypeError(
+ f"Invalid inner dimensions for array argument '{input_arg.name}', expected {input_arg.dtype_shape}, got {input_value.shape[-input_arg.dtype_ndim :]}"
+ )
+ else:
+ # make sure scalar is not a traced variable, should be static
+ if isinstance(input_value, jax.core.Tracer):
+ raise ValueError(f"Argument '{input_arg.name}' must be a static value")
+ # stash the value to be retrieved by callback
+ static_inputs[input_arg.name] = input_arg.type(input_value)
+
+ # append in-out arg to output types
+ if input_arg.in_out:
+ out_types.append(get_jax_output_type(input_arg, input_value.shape))
+
+ # launch dimensions
+ if launch_dims is None:
+ # use the shape of the first input array
+ if self.first_array_arg is not None:
+ launch_dims = get_warp_shape(self.input_args[self.first_array_arg], args[self.first_array_arg].shape)
+ else:
+ raise RuntimeError("Failed to determine launch dimensions")
+ elif isinstance(launch_dims, int):
+ launch_dims = (launch_dims,)
+ else:
+ launch_dims = tuple(launch_dims)
+
+ # output shapes
+ if isinstance(output_dims, dict):
+ # assume a dictionary of shapes keyed on argument name
+ for output_arg in self.output_args:
+ dims = output_dims.get(output_arg.name)
+ if dims is None:
+ raise ValueError(f"Missing output dimensions for argument '{output_arg.name}'")
+ out_types.append(get_jax_output_type(output_arg, dims))
+ else:
+ if output_dims is None:
+ # use launch dimensions
+ output_dims = launch_dims
+ elif isinstance(output_dims, int):
+ output_dims = (output_dims,)
+ # assume same dimensions for all outputs
+ for output_arg in self.output_args:
+ out_types.append(get_jax_output_type(output_arg, output_dims))
+
+ call = jax.ffi.ffi_call(
+ self.name,
+ out_types,
+ vmap_method=vmap_method,
+ input_output_aliases=self.input_output_aliases,
+ )
+
+ # ensure the kernel module is loaded before the callback, otherwise graph capture may fail
+ device = wp.device_from_jax(get_jax_device())
+ self.kernel.module.load(device)
+
+ # save launch data to be retrieved by callback
+ launch_id = self.launch_id
+ self.launch_descriptors[launch_id] = FfiLaunchDesc(static_inputs, launch_dims)
+ self.launch_id += 1
+
+ return call(*args, launch_id=launch_id)
+
+ def ffi_callback(self, call_frame):
+ try:
+ # On the first call, XLA runtime will query the API version and traits
+ # metadata using the |extension| field. Let us respond to that query
+ # if the metadata extension is present.
+ extension = call_frame.contents.extension_start
+ if extension:
+ # Try to set the version metadata.
+ if extension.contents.type == XLA_FFI_Extension_Type.Metadata:
+ metadata_ext = ctypes.cast(extension, ctypes.POINTER(XLA_FFI_Metadata_Extension))
+ metadata_ext.contents.metadata.contents.api_version.major_version = 0
+ metadata_ext.contents.metadata.contents.api_version.minor_version = 1
+ # Turn on CUDA graphs for this handler.
+ metadata_ext.contents.metadata.contents.traits = (
+ XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE
+ )
+ return None
+
+ # retrieve call info
+ attrs = decode_attrs(call_frame.contents.attrs)
+ launch_id = int(attrs["launch_id"])
+ launch_desc = self.launch_descriptors[launch_id]
+
+ num_inputs = call_frame.contents.args.size
+ inputs = ctypes.cast(call_frame.contents.args.args, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer)))
+
+ num_outputs = call_frame.contents.rets.size
+ outputs = ctypes.cast(call_frame.contents.rets.rets, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer)))
+
+ assert num_inputs == self.num_inputs
+ assert num_outputs == self.num_outputs
+
+ launch_bounds = launch_bounds_t(launch_desc.launch_dims)
+
+ # first kernel param is the launch bounds
+ kernel_params = (ctypes.c_void_p * (1 + self.num_kernel_args))()
+ kernel_params[0] = ctypes.addressof(launch_bounds)
+
+ arg_refs = []
+
+ # input and in-out args
+ for i, input_arg in enumerate(self.input_args):
+ if input_arg.is_array:
+ buffer = inputs[i].contents
+ shape = buffer.dims[: input_arg.type.ndim]
+ strides = strides_from_shape(shape, input_arg.type.dtype)
+ arg = array_t(buffer.data, 0, input_arg.type.ndim, shape, strides)
+ kernel_params[i + 1] = ctypes.addressof(arg)
+ arg_refs.append(arg) # keep a reference
+ else:
+ # scalar argument, get stashed value
+ value = launch_desc.static_inputs[input_arg.name]
+ arg = input_arg.type._type_(value)
+ kernel_params[i + 1] = ctypes.addressof(arg)
+ arg_refs.append(arg) # keep a reference
+
+ # pure output args (skip in-out FFI buffers)
+ for i, output_arg in enumerate(self.output_args):
+ buffer = outputs[i + self.num_in_out].contents
+ shape = buffer.dims[: output_arg.type.ndim]
+ strides = strides_from_shape(shape, output_arg.type.dtype)
+ arg = array_t(buffer.data, 0, output_arg.type.ndim, shape, strides)
+ kernel_params[num_inputs + i + 1] = ctypes.addressof(arg)
+ arg_refs.append(arg) # keep a reference
+
+ # get device and stream
+ device = wp.device_from_jax(get_jax_device())
+ stream = get_stream_from_callframe(call_frame.contents)
+
+ # get kernel hooks
+ hooks = self.kernel.module.get_kernel_hooks(self.kernel, device)
+ assert hooks.forward, "Failed to find kernel entry point"
+
+ # launch the kernel
+ wp.context.runtime.core.wp_cuda_launch_kernel(
+ device.context,
+ hooks.forward,
+ launch_bounds.size,
+ 0,
+ 256,
+ hooks.forward_smem_bytes,
+ kernel_params,
+ stream,
+ )
+
+ except Exception as e:
+ print(traceback.format_exc())
+ return create_ffi_error(
+ call_frame.contents.api, XLA_FFI_Error_Code.UNKNOWN, f"FFI callback error: {type(e).__name__}: {e}"
+ )
+
+
+class FfiCallDesc:
+ def __init__(self, static_inputs):
+ self.static_inputs = static_inputs
+
+
+class FfiCallable:
+ def __init__(self, func, num_outputs, graph_compatible, vmap_method, output_dims, in_out_argnames):
+ self.func = func
+ self.name = generate_unique_name(func)
+ self.num_outputs = num_outputs
+ self.vmap_method = vmap_method
+ self.graph_compatible = graph_compatible
+ self.output_dims = output_dims
+ self.first_array_arg = None
+ self.call_id = 0
+ self.call_descriptors = {}
+
+ in_out_argnames_list = in_out_argnames or []
+ in_out_argnames = set(in_out_argnames_list)
+ if len(in_out_argnames_list) != len(in_out_argnames):
+ raise AssertionError("in_out_argnames must not contain duplicate names")
+
+ # get arguments and annotations
+ argspec = get_full_arg_spec(func)
+
+ num_args = len(argspec.args)
+ self.num_in_out = len(in_out_argnames)
+ self.num_inputs = num_args - num_outputs + self.num_in_out
+ if self.num_outputs < 1:
+ raise ValueError("At least one output is required")
+ if self.num_outputs > num_args:
+ raise ValueError("Number of outputs cannot be greater than the number of kernel arguments")
+ if self.num_outputs < self.num_in_out:
+ raise ValueError("Number of outputs cannot be smaller than the number of in_out_argnames")
+
+ if len(argspec.annotations) < num_args:
+ raise RuntimeError(f"Incomplete argument annotations on function {self.name}")
+
+ # parse type annotations
+ self.args = []
+ arg_idx = 0
+ for arg_name, arg_type in argspec.annotations.items():
+ if arg_name == "return":
+ if arg_type is not None:
+ raise TypeError("Function must not return a value")
+ else:
+ arg = FfiArg(arg_name, arg_type, arg_name in in_out_argnames)
+ if arg_name in in_out_argnames:
+ in_out_argnames.remove(arg_name)
+ if arg.is_array:
+ if arg_idx < self.num_inputs and self.first_array_arg is None:
+ self.first_array_arg = arg_idx
+ self.args.append(arg)
+
+ if arg.in_out and arg_idx >= self.num_inputs:
+ raise AssertionError(
+ f"Expected an output-only argument for argument {arg_name}."
+ " in_out arguments should be placed before output-only arguments."
+ )
+
+ arg_idx += 1
+
+ if in_out_argnames:
+ raise ValueError(f"in_out_argnames: '{in_out_argnames}' did not match any function argument names.")
+
+ self.input_args = self.args[: self.num_inputs] # includes in-out args
+ self.output_args = self.args[self.num_inputs :] # pure output args
+
+ # Build input output aliases.
+ out_id = 0
+ input_output_aliases = {}
+ for in_id, arg in enumerate(self.input_args):
+ if not arg.in_out:
+ continue
+ input_output_aliases[in_id] = out_id
+ out_id += 1
+ self.input_output_aliases = input_output_aliases
+
+ # register the callback
+ FFI_CCALLFUNC = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame))
+ self.callback_func = FFI_CCALLFUNC(lambda call_frame: self.ffi_callback(call_frame))
+ ffi_ccall_address = ctypes.cast(self.callback_func, ctypes.c_void_p)
+ ffi_capsule = jax.ffi.pycapsule(ffi_ccall_address.value)
+ jax.ffi.register_ffi_target(self.name, ffi_capsule, platform="CUDA")
+
+ def __call__(self, *args, output_dims=None, vmap_method=None):
+ num_inputs = len(args)
+ if num_inputs != self.num_inputs:
+ input_names = ", ".join(arg.name for arg in self.input_args)
+ s = "" if self.num_inputs == 1 else "s"
+ raise ValueError(f"Expected {self.num_inputs} input{s} ({input_names}), but got {num_inputs}")
+
+ # default argument fallback
+ if vmap_method is None:
+ vmap_method = self.vmap_method
+ if output_dims is None:
+ output_dims = self.output_dims
+
+ # output types
+ out_types = []
+
+ # process inputs
+ static_inputs = {}
+ for i in range(num_inputs):
+ input_arg = self.input_args[i]
+ input_value = args[i]
+ if input_arg.is_array:
+ # check dtype
+ if input_value.dtype != input_arg.jax_scalar_type:
+ raise TypeError(
+ f"Invalid data type for array argument '{input_arg.name}', expected {input_arg.jax_scalar_type}, got {input_value.dtype}"
+ )
+ # check ndim
+ if input_value.ndim != input_arg.jax_ndim:
+ raise TypeError(
+ f"Invalid dimensionality for array argument '{input_arg.name}', expected {input_arg.jax_ndim} dimensions, got {input_value.ndim}"
+ )
+ # check inner dims
+ for d in range(input_arg.dtype_ndim):
+ if input_value.shape[input_arg.type.ndim + d] != input_arg.dtype_shape[d]:
+ raise TypeError(
+ f"Invalid inner dimensions for array argument '{input_arg.name}', expected {input_arg.dtype_shape}, got {input_value.shape[-input_arg.dtype_ndim :]}"
+ )
+ else:
+ # make sure scalar is not a traced variable, should be static
+ if isinstance(input_value, jax.core.Tracer):
+ raise ValueError(f"Argument '{input_arg.name}' must be a static value")
+ # stash the value to be retrieved by callback
+ static_inputs[input_arg.name] = input_arg.type(input_value)
+
+ # append in-out arg to output types
+ if input_arg.in_out:
+ out_types.append(get_jax_output_type(input_arg, input_value.shape))
+
+ # output shapes
+ if isinstance(output_dims, dict):
+ # assume a dictionary of shapes keyed on argument name
+ for output_arg in self.output_args:
+ dims = output_dims.get(output_arg.name)
+ if dims is None:
+ raise ValueError(f"Missing output dimensions for argument '{output_arg.name}'")
+ out_types.append(get_jax_output_type(output_arg, dims))
+ else:
+ if output_dims is None:
+ if self.first_array_arg is None:
+ raise ValueError("Unable to determine output dimensions")
+ output_dims = get_warp_shape(self.input_args[self.first_array_arg], args[self.first_array_arg].shape)
+ elif isinstance(output_dims, int):
+ output_dims = (output_dims,)
+ # assume same dimensions for all outputs
+ for output_arg in self.output_args:
+ out_types.append(get_jax_output_type(output_arg, output_dims))
+
+ call = jax.ffi.ffi_call(
+ self.name,
+ out_types,
+ vmap_method=vmap_method,
+ input_output_aliases=self.input_output_aliases,
+ # has_side_effect=True, # force this function to execute even if outputs aren't used
+ )
+
+ # load the module
+ # NOTE: if the target function uses kernels from different modules, they will not be loaded here
+ device = wp.device_from_jax(get_jax_device())
+ module = wp.get_module(self.func.__module__)
+ module.load(device)
+
+ # save call data to be retrieved by callback
+ call_id = self.call_id
+ self.call_descriptors[call_id] = FfiCallDesc(static_inputs)
+ self.call_id += 1
+ return call(*args, call_id=call_id)
+
+ def ffi_callback(self, call_frame):
+ try:
+ # TODO Try-catch around the body and return XLA_FFI_Error on error.
+ extension = call_frame.contents.extension_start
+ # On the first call, XLA runtime will query the API version and traits
+ # metadata using the |extension| field. Let us respond to that query
+ # if the metadata extension is present.
+ if extension:
+ # Try to set the version metadata.
+ if extension.contents.type == XLA_FFI_Extension_Type.Metadata:
+ metadata_ext = ctypes.cast(extension, ctypes.POINTER(XLA_FFI_Metadata_Extension))
+ metadata_ext.contents.metadata.contents.api_version.major_version = 0
+ metadata_ext.contents.metadata.contents.api_version.minor_version = 1
+ # Turn on CUDA graphs for this handler.
+ if self.graph_compatible:
+ metadata_ext.contents.metadata.contents.traits = (
+ XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE
+ )
+ return None
+
+ # retrieve call info
+ attrs = decode_attrs(call_frame.contents.attrs)
+ call_id = int(attrs["call_id"])
+ call_desc = self.call_descriptors[call_id]
+
+ num_inputs = call_frame.contents.args.size
+ inputs = ctypes.cast(call_frame.contents.args.args, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer)))
+
+ num_outputs = call_frame.contents.rets.size
+ outputs = ctypes.cast(call_frame.contents.rets.rets, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer)))
+
+ assert num_inputs == self.num_inputs
+ assert num_outputs == self.num_outputs
+
+ device = wp.device_from_jax(get_jax_device())
+ cuda_stream = get_stream_from_callframe(call_frame.contents)
+ stream = wp.Stream(device, cuda_stream=cuda_stream)
+
+ # reconstruct the argument list
+ arg_list = []
+
+ # input and in-out args
+ for i, arg in enumerate(self.input_args):
+ if arg.is_array:
+ buffer = inputs[i].contents
+ shape = buffer.dims[: buffer.rank - arg.dtype_ndim]
+ arr = wp.array(ptr=buffer.data, dtype=arg.type.dtype, shape=shape, device=device)
+ arg_list.append(arr)
+ else:
+ # scalar argument, get stashed value
+ value = call_desc.static_inputs[arg.name]
+ arg_list.append(value)
+
+ # pure output args (skip in-out FFI buffers)
+ for i, arg in enumerate(self.output_args):
+ buffer = outputs[i + self.num_in_out].contents
+ shape = buffer.dims[: buffer.rank - arg.dtype_ndim]
+ arr = wp.array(ptr=buffer.data, dtype=arg.type.dtype, shape=shape, device=device)
+ arg_list.append(arr)
+
+ # call the Python function with reconstructed arguments
+ with wp.ScopedStream(stream, sync_enter=False):
+ if stream.is_capturing:
+ with wp.ScopedCapture(stream=stream, external=True) as capture:
+ self.func(*arg_list)
+ # keep a reference to the capture object to prevent required modules getting unloaded
+ call_desc.capture = capture
+ else:
+ self.func(*arg_list)
+
+ except Exception as e:
+ print(traceback.format_exc())
+ return create_ffi_error(
+ call_frame.contents.api, XLA_FFI_Error_Code.UNKNOWN, f"FFI callback error: {type(e).__name__}: {e}"
+ )
+
+ return None
+
+
+# Holders for the custom callbacks to keep them alive.
+_FFI_CALLABLE_REGISTRY: dict[str, FfiCallable] = {}
+_FFI_KERNEL_REGISTRY: dict[str, FfiKernel] = {}
+_FFI_REGISTRY_LOCK = threading.Lock()
+
+
+def jax_kernel(
+ kernel, num_outputs=1, vmap_method="broadcast_all", launch_dims=None, output_dims=None, in_out_argnames=None
+):
+ """Create a JAX callback from a Warp kernel.
+
+ NOTE: This is an experimental feature under development.
+
+ Args:
+ kernel: The Warp kernel to launch.
+ num_outputs: Optional. Specify the number of output arguments if greater than 1.
+ This must include the number of ``in_out_arguments``.
+ vmap_method: Optional. String specifying how the callback transforms under ``vmap()``.
+ This argument can also be specified for individual calls.
+ launch_dims: Optional. Specify the default kernel launch dimensions. If None, launch
+ dimensions are inferred from the shape of the first array argument.
+ This argument can also be specified for individual calls.
+ output_dims: Optional. Specify the default dimensions of output arrays. If None, output
+ dimensions are inferred from the launch dimensions.
+ This argument can also be specified for individual calls.
+ in_out_argnames: Optional. Names of input-output arguments.
+
+ Limitations:
+ - All kernel arguments must be contiguous arrays or scalars.
+ - Scalars must be static arguments in JAX.
+ - Input and input-output arguments must precede the output arguments in the ``kernel`` definition.
+ - There must be at least one output or input-output argument.
+ - Only the CUDA backend is supported.
+ """
+ key = (
+ kernel.func,
+ num_outputs,
+ vmap_method,
+ tuple(launch_dims) if launch_dims else launch_dims,
+ tuple(sorted(output_dims.items())) if output_dims else output_dims,
+ )
+
+ with _FFI_REGISTRY_LOCK:
+ if key not in _FFI_KERNEL_REGISTRY:
+ new_kernel = FfiKernel(kernel, num_outputs, vmap_method, launch_dims, output_dims, in_out_argnames)
+ _FFI_KERNEL_REGISTRY[key] = new_kernel
+
+ return _FFI_KERNEL_REGISTRY[key]
+
+
+def jax_callable(
+ func: Callable,
+ num_outputs: int = 1,
+ graph_compatible: bool = True,
+ vmap_method: Optional[str] = "broadcast_all",
+ output_dims=None,
+ in_out_argnames=None,
+):
+ """Create a JAX callback from an annotated Python function.
+
+ The Python function arguments must have type annotations like Warp kernels.
+
+ NOTE: This is an experimental feature under development.
+
+ Args:
+ func: The Python function to call.
+ num_outputs: Optional. Specify the number of output arguments if greater than 1.
+ This must include the number of ``in_out_arguments``.
+ graph_compatible: Optional. Whether the function can be called during CUDA graph capture.
+ vmap_method: Optional. String specifying how the callback transforms under ``vmap()``.
+ This argument can also be specified for individual calls.
+ output_dims: Optional. Specify the default dimensions of output arrays.
+ If ``None``, output dimensions are inferred from the launch dimensions.
+ This argument can also be specified for individual calls.
+ in_out_argnames: Optional. Names of input-output arguments.
+
+ Limitations:
+ - All kernel arguments must be contiguous arrays or scalars.
+ - Scalars must be static arguments in JAX.
+ - Input and input-output arguments must precede the output arguments in the ``func`` definition.
+ - There must be at least one output or input-output argument.
+ - Only the CUDA backend is supported.
+ """
+ key = (
+ func,
+ num_outputs,
+ graph_compatible,
+ vmap_method,
+ tuple(sorted(output_dims.items())) if output_dims else output_dims,
+ )
+
+ with _FFI_REGISTRY_LOCK:
+ if key not in _FFI_CALLABLE_REGISTRY:
+ new_callable = FfiCallable(func, num_outputs, graph_compatible, vmap_method, output_dims, in_out_argnames)
+ _FFI_CALLABLE_REGISTRY[key] = new_callable
+
+ return _FFI_CALLABLE_REGISTRY[key]
+
+
+###############################################################################
+#
+# Generic FFI callbacks for Python functions of the form
+# func(inputs, outputs, attrs, ctx)
+#
+###############################################################################
+
+
+def register_ffi_callback(name: str, func: Callable, graph_compatible: bool = True) -> None:
+ """Create a JAX callback from a Python function.
+
+ The Python function must have the form ``func(inputs, outputs, attrs, ctx)``.
+
+ NOTE: This is an experimental feature under development.
+
+ Args:
+ name: A unique FFI callback name.
+ func: The Python function to call.
+ graph_compatible: Optional. Whether the function can be called during CUDA graph capture.
+ """
+
+ # TODO check that the name is not already registered
+
+ def ffi_callback(call_frame):
+ try:
+ # TODO Try-catch around the body and return XLA_FFI_Error on error.
+ extension = call_frame.contents.extension_start
+ # On the first call, XLA runtime will query the API version and traits
+ # metadata using the |extension| field. Let us respond to that query
+ # if the metadata extension is present.
+ if extension:
+ # Try to set the version metadata.
+ if extension.contents.type == XLA_FFI_Extension_Type.Metadata:
+ metadata_ext = ctypes.cast(extension, ctypes.POINTER(XLA_FFI_Metadata_Extension))
+ metadata_ext.contents.metadata.contents.api_version.major_version = 0
+ metadata_ext.contents.metadata.contents.api_version.minor_version = 1
+ if graph_compatible:
+ # Turn on CUDA graphs for this handler.
+ metadata_ext.contents.metadata.contents.traits = (
+ XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE
+ )
+ return None
+
+ attrs = decode_attrs(call_frame.contents.attrs)
+
+ input_count = call_frame.contents.args.size
+ inputs = ctypes.cast(call_frame.contents.args.args, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer)))
+ inputs = [FfiBuffer(inputs[i].contents) for i in range(input_count)]
+
+ output_count = call_frame.contents.rets.size
+ outputs = ctypes.cast(call_frame.contents.rets.rets, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer)))
+ outputs = [FfiBuffer(outputs[i].contents) for i in range(output_count)]
+
+ ctx = ExecutionContext(call_frame.contents)
+
+ func(inputs, outputs, attrs, ctx)
+ except Exception as e:
+ print(traceback.format_exc())
+ return create_ffi_error(
+ call_frame.contents.api, XLA_FFI_Error_Code.UNKNOWN, f"FFI callback error: {type(e).__name__}: {e}"
+ )
+
+ return None
+
+ FFI_CCALLFUNC = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame))
+ callback_func = FFI_CCALLFUNC(ffi_callback)
+ with _FFI_REGISTRY_LOCK:
+ _FFI_CALLABLE_REGISTRY[name] = callback_func
+ ffi_ccall_address = ctypes.cast(callback_func, ctypes.c_void_p)
+ ffi_capsule = jax.ffi.pycapsule(ffi_ccall_address.value)
+ jax.ffi.register_ffi_target(name, ffi_capsule, platform="CUDA")
+
+
+###############################################################################
+#
+# Utilities
+#
+###############################################################################
+
+# ensure unique FFI callback names
+ffi_name_counts = {}
+
+
+def generate_unique_name(func) -> str:
+ key = make_full_qualified_name(func)
+ unique_id = ffi_name_counts.get(key, 0)
+ ffi_name_counts[key] = unique_id + 1
+ return f"{key}_{unique_id}"
+
+
+def get_warp_shape(arg, dims):
+ if arg.dtype_ndim > 0:
+ # vector/matrix array
+ return dims[: arg.warp_ndim]
+ else:
+ # scalar array
+ return dims
+
+
+def get_jax_output_type(arg, dims):
+ if isinstance(dims, int):
+ dims = (dims,)
+
+ ndim = len(dims)
+
+ if arg.dtype_ndim > 0:
+ # vector/matrix array
+ if ndim == arg.warp_ndim:
+ return jax.ShapeDtypeStruct((*dims, *arg.dtype_shape), arg.jax_scalar_type)
+ elif ndim == arg.jax_ndim:
+ # make sure inner dimensions match
+ inner_dims = dims[-arg.dtype_ndim :]
+ for i in range(arg.dtype_ndim):
+ if inner_dims[i] != arg.dtype_shape[i]:
+ raise ValueError(f"Invalid output dimensions for argument '{arg.name}': {dims}")
+ return jax.ShapeDtypeStruct(dims, arg.jax_scalar_type)
+ else:
+ raise ValueError(f"Invalid output dimensions for argument '{arg.name}': {dims}")
+ else:
+ # scalar array
+ if ndim != arg.warp_ndim:
+ raise ValueError(f"Invalid output dimensions for argument '{arg.name}': {dims}")
+ return jax.ShapeDtypeStruct(dims, arg.jax_scalar_type)
diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py
new file mode 100644
index 00000000..8f33030b
--- /dev/null
+++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py
@@ -0,0 +1,615 @@
+# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+#
+# 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.
+
+import ctypes
+import enum
+
+import jax.numpy as jnp
+import numpy as np
+
+import warp as wp
+
+#######################################################################
+# ctypes structures and enums for XLA's FFI API:
+# https://github.com/openxla/xla/blob/a1a5e62fbffa3a3b6c409d72607456cf5b353a22/xla/ffi/api/c_api.h
+#######################################################################
+
+
+# typedef enum {
+# XLA_FFI_Extension_Metadata = 1,
+# } XLA_FFI_Extension_Type;
+class XLA_FFI_Extension_Type(enum.IntEnum):
+ Metadata = 1
+
+
+# typedef struct XLA_FFI_Extension_Base {
+# size_t struct_size;
+# XLA_FFI_Extension_Type type;
+# struct XLA_FFI_Extension_Base* next;
+# } XLA_FFI_Extension_Base;
+class XLA_FFI_Extension_Base(ctypes.Structure):
+ pass
+
+
+XLA_FFI_Extension_Base._fields_ = [
+ ("struct_size", ctypes.c_size_t),
+ ("type", ctypes.c_int), # XLA_FFI_Extension_Type
+ ("next", ctypes.POINTER(XLA_FFI_Extension_Base)),
+]
+
+
+# typedef enum {
+# XLA_FFI_ExecutionStage_INSTANTIATE = 0,
+# XLA_FFI_ExecutionStage_PREPARE = 1,
+# XLA_FFI_ExecutionStage_INITIALIZE = 2,
+# XLA_FFI_ExecutionStage_EXECUTE = 3,
+# } XLA_FFI_ExecutionStage;
+class XLA_FFI_ExecutionStage(enum.IntEnum):
+ INSTANTIATE = 0
+ PREPARE = 1
+ INITIALIZE = 2
+ EXECUTE = 3
+
+
+# typedef enum {
+# XLA_FFI_DataType_INVALID = 0,
+# XLA_FFI_DataType_PRED = 1,
+# XLA_FFI_DataType_S8 = 2,
+# XLA_FFI_DataType_S16 = 3,
+# XLA_FFI_DataType_S32 = 4,
+# XLA_FFI_DataType_S64 = 5,
+# XLA_FFI_DataType_U8 = 6,
+# XLA_FFI_DataType_U16 = 7,
+# XLA_FFI_DataType_U32 = 8,
+# XLA_FFI_DataType_U64 = 9,
+# XLA_FFI_DataType_F16 = 10,
+# XLA_FFI_DataType_F32 = 11,
+# XLA_FFI_DataType_F64 = 12,
+# XLA_FFI_DataType_BF16 = 16,
+# XLA_FFI_DataType_C64 = 15,
+# XLA_FFI_DataType_C128 = 18,
+# XLA_FFI_DataType_TOKEN = 17,
+# XLA_FFI_DataType_F8E5M2 = 19,
+# XLA_FFI_DataType_F8E3M4 = 29,
+# XLA_FFI_DataType_F8E4M3 = 28,
+# XLA_FFI_DataType_F8E4M3FN = 20,
+# XLA_FFI_DataType_F8E4M3B11FNUZ = 23,
+# XLA_FFI_DataType_F8E5M2FNUZ = 24,
+# XLA_FFI_DataType_F8E4M3FNUZ = 25,
+# XLA_FFI_DataType_F4E2M1FN = 32,
+# XLA_FFI_DataType_F8E8M0FNU = 33,
+# } XLA_FFI_DataType;
+class XLA_FFI_DataType(enum.IntEnum):
+ INVALID = 0
+ PRED = 1
+ S8 = 2
+ S16 = 3
+ S32 = 4
+ S64 = 5
+ U8 = 6
+ U16 = 7
+ U32 = 8
+ U64 = 9
+ F16 = 10
+ F32 = 11
+ F64 = 12
+ BF16 = 16
+ C64 = 15
+ C128 = 18
+ TOKEN = 17
+ F8E5M2 = 19
+ F8E3M4 = 29
+ F8E4M3 = 28
+ F8E4M3FN = 20
+ F8E4M3B11FNUZ = 23
+ F8E5M2FNUZ = 24
+ F8E4M3FNUZ = 25
+ F4E2M1FN = 32
+ F8E8M0FNU = 33
+
+
+# struct XLA_FFI_Buffer {
+# size_t struct_size;
+# XLA_FFI_Extension_Base* extension_start;
+#
+# XLA_FFI_DataType dtype;
+# void* data;
+# int64_t rank;
+# int64_t* dims; // length == rank
+# };
+class XLA_FFI_Buffer(ctypes.Structure):
+ _fields_ = (
+ ("struct_size", ctypes.c_size_t),
+ ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)),
+ ("dtype", ctypes.c_int), # XLA_FFI_DataType
+ ("data", ctypes.c_void_p),
+ ("rank", ctypes.c_int64),
+ ("dims", ctypes.POINTER(ctypes.c_int64)),
+ )
+
+
+# typedef enum {
+# XLA_FFI_ArgType_BUFFER = 1,
+# } XLA_FFI_ArgType;
+class XLA_FFI_ArgType(enum.IntEnum):
+ BUFFER = 1
+
+
+# typedef enum {
+# XLA_FFI_RetType_BUFFER = 1,
+# } XLA_FFI_RetType;
+class XLA_FFI_RetType(enum.IntEnum):
+ BUFFER = 1
+
+
+# struct XLA_FFI_Args {
+# size_t struct_size;
+# XLA_FFI_Extension_Base* extension_start;
+# int64_t size;
+# XLA_FFI_ArgType* types; // length == size
+# void** args; // length == size
+# };
+class XLA_FFI_Args(ctypes.Structure):
+ _fields_ = (
+ ("struct_size", ctypes.c_size_t),
+ ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)),
+ ("size", ctypes.c_int64),
+ ("types", ctypes.POINTER(ctypes.c_int)), # XLA_FFI_ArgType*
+ ("args", ctypes.POINTER(ctypes.c_void_p)),
+ )
+
+
+# struct XLA_FFI_Rets {
+# size_t struct_size;
+# XLA_FFI_Extension_Base* extension_start;
+# int64_t size;
+# XLA_FFI_RetType* types; // length == size
+# void** rets; // length == size
+# };
+class XLA_FFI_Rets(ctypes.Structure):
+ _fields_ = (
+ ("struct_size", ctypes.c_size_t),
+ ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)),
+ ("size", ctypes.c_int64),
+ ("types", ctypes.POINTER(ctypes.c_int)), # XLA_FFI_RetType*
+ ("rets", ctypes.POINTER(ctypes.c_void_p)),
+ )
+
+
+# typedef struct XLA_FFI_ByteSpan {
+# const char* ptr;
+# size_t len;
+# } XLA_FFI_ByteSpan;
+class XLA_FFI_ByteSpan(ctypes.Structure):
+ _fields_ = (
+ ("ptr", ctypes.POINTER(ctypes.c_char)),
+ ("len", ctypes.c_size_t),
+ )
+
+
+# typedef struct XLA_FFI_Scalar {
+# XLA_FFI_DataType dtype;
+# void* value;
+# } XLA_FFI_Scalar;
+class XLA_FFI_Scalar(ctypes.Structure):
+ _fields_ = (
+ ("dtype", ctypes.c_int),
+ ("value", ctypes.c_void_p),
+ )
+
+
+# typedef struct XLA_FFI_Array {
+# XLA_FFI_DataType dtype;
+# size_t size;
+# void* data;
+# } XLA_FFI_Array;
+class XLA_FFI_Array(ctypes.Structure):
+ _fields_ = (
+ ("dtype", ctypes.c_int),
+ ("size", ctypes.c_size_t),
+ ("data", ctypes.c_void_p),
+ )
+
+
+# typedef enum {
+# XLA_FFI_AttrType_ARRAY = 1,
+# XLA_FFI_AttrType_DICTIONARY = 2,
+# XLA_FFI_AttrType_SCALAR = 3,
+# XLA_FFI_AttrType_STRING = 4,
+# } XLA_FFI_AttrType;
+class XLA_FFI_AttrType(enum.IntEnum):
+ ARRAY = 1
+ DICTIONARY = 2
+ SCALAR = 3
+ STRING = 4
+
+
+# struct XLA_FFI_Attrs {
+# size_t struct_size;
+# XLA_FFI_Extension_Base* extension_start;
+# int64_t size;
+# XLA_FFI_AttrType* types; // length == size
+# XLA_FFI_ByteSpan** names; // length == size
+# void** attrs; // length == size
+# };
+class XLA_FFI_Attrs(ctypes.Structure):
+ _fields_ = (
+ ("struct_size", ctypes.c_size_t),
+ ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)),
+ ("size", ctypes.c_int64),
+ ("types", ctypes.POINTER(ctypes.c_int)), # XLA_FFI_AttrType*
+ ("names", ctypes.POINTER(ctypes.POINTER(XLA_FFI_ByteSpan))),
+ ("attrs", ctypes.POINTER(ctypes.c_void_p)),
+ )
+
+
+# struct XLA_FFI_Api_Version {
+# size_t struct_size;
+# XLA_FFI_Extension_Base* extension_start;
+# int major_version; // out
+# int minor_version; // out
+# };
+class XLA_FFI_Api_Version(ctypes.Structure):
+ _fields_ = (
+ ("struct_size", ctypes.c_size_t),
+ ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)),
+ ("major_version", ctypes.c_int),
+ ("minor_version", ctypes.c_int),
+ )
+
+
+# enum XLA_FFI_Handler_TraitsBits {
+# // Calls to FFI handler are safe to trace into the command buffer. It means
+# // that calls to FFI handler always launch exactly the same device operations
+# // (can depend on attribute values) that can be captured and then replayed.
+# XLA_FFI_HANDLER_TRAITS_COMMAND_BUFFER_COMPATIBLE = 1u << 0,
+# };
+class XLA_FFI_Handler_TraitsBits(enum.IntEnum):
+ COMMAND_BUFFER_COMPATIBLE = 1 << 0
+
+
+# struct XLA_FFI_Metadata {
+# size_t struct_size;
+# XLA_FFI_Api_Version api_version;
+# XLA_FFI_Handler_Traits traits;
+# };
+class XLA_FFI_Metadata(ctypes.Structure):
+ _fields_ = (
+ ("struct_size", ctypes.c_size_t),
+ ("api_version", XLA_FFI_Api_Version), # XLA_FFI_Extension_Type
+ ("traits", ctypes.c_uint32), # XLA_FFI_Handler_Traits
+ )
+
+
+# struct XLA_FFI_Metadata_Extension {
+# XLA_FFI_Extension_Base extension_base;
+# XLA_FFI_Metadata* metadata;
+# };
+class XLA_FFI_Metadata_Extension(ctypes.Structure):
+ _fields_ = (
+ ("extension_base", XLA_FFI_Extension_Base),
+ ("metadata", ctypes.POINTER(XLA_FFI_Metadata)),
+ )
+
+
+# typedef enum {
+# XLA_FFI_Error_Code_OK = 0,
+# XLA_FFI_Error_Code_CANCELLED = 1,
+# XLA_FFI_Error_Code_UNKNOWN = 2,
+# XLA_FFI_Error_Code_INVALID_ARGUMENT = 3,
+# XLA_FFI_Error_Code_DEADLINE_EXCEEDED = 4,
+# XLA_FFI_Error_Code_NOT_FOUND = 5,
+# XLA_FFI_Error_Code_ALREADY_EXISTS = 6,
+# XLA_FFI_Error_Code_PERMISSION_DENIED = 7,
+# XLA_FFI_Error_Code_RESOURCE_EXHAUSTED = 8,
+# XLA_FFI_Error_Code_FAILED_PRECONDITION = 9,
+# XLA_FFI_Error_Code_ABORTED = 10,
+# XLA_FFI_Error_Code_OUT_OF_RANGE = 11,
+# XLA_FFI_Error_Code_UNIMPLEMENTED = 12,
+# XLA_FFI_Error_Code_INTERNAL = 13,
+# XLA_FFI_Error_Code_UNAVAILABLE = 14,
+# XLA_FFI_Error_Code_DATA_LOSS = 15,
+# XLA_FFI_Error_Code_UNAUTHENTICATED = 16
+# } XLA_FFI_Error_Code;
+class XLA_FFI_Error_Code(enum.IntEnum):
+ OK = 0
+ CANCELLED = 1
+ UNKNOWN = 2
+ INVALID_ARGUMENT = 3
+ DEADLINE_EXCEEDED = 4
+ NOT_FOUND = 5
+ ALREADY_EXISTS = 6
+ PERMISSION_DENIED = 7
+ RESOURCE_EXHAUSTED = 8
+ FAILED_PRECONDITION = 9
+ ABORTED = 10
+ OUT_OF_RANGE = 11
+ UNIMPLEMENTED = 12
+ INTERNAL = 13
+ UNAVAILABLE = 14
+ DATA_LOSS = 15
+ UNAUTHENTICATED = 16
+
+
+# struct XLA_FFI_Error_Create_Args {
+# size_t struct_size;
+# XLA_FFI_Extension_Base* extension_start;
+# const char* message;
+# XLA_FFI_Error_Code errc;
+# };
+class XLA_FFI_Error_Create_Args(ctypes.Structure):
+ _fields_ = (
+ ("struct_size", ctypes.c_size_t),
+ ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)),
+ ("message", ctypes.c_char_p),
+ ("errc", ctypes.c_int),
+ ) # XLA_FFI_Error_Code
+
+
+XLA_FFI_Error_Create = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_Error_Create_Args))
+
+
+# struct XLA_FFI_Stream_Get_Args {
+# size_t struct_size;
+# XLA_FFI_Extension_Base* extension_start;
+# XLA_FFI_ExecutionContext* ctx;
+# void* stream; // out
+# };
+class XLA_FFI_Stream_Get_Args(ctypes.Structure):
+ _fields_ = (
+ ("struct_size", ctypes.c_size_t),
+ ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)),
+ ("ctx", ctypes.c_void_p), # XLA_FFI_ExecutionContext*
+ ("stream", ctypes.c_void_p),
+ ) # // out
+
+
+XLA_FFI_Stream_Get = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_Stream_Get_Args))
+
+
+# struct XLA_FFI_Api {
+# size_t struct_size;
+# XLA_FFI_Extension_Base* extension_start;
+#
+# XLA_FFI_Api_Version api_version;
+# XLA_FFI_InternalApi* internal_api;
+#
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Error_Create);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Error_GetMessage);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Error_Destroy);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Handler_Register);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Stream_Get);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_TypeId_Register);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_ExecutionContext_Get);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_State_Set);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_State_Get);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_DeviceMemory_Allocate);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_DeviceMemory_Free);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_ThreadPool_Schedule);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_ThreadPool_NumThreads);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Future_Create);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Future_SetAvailable);
+# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Future_SetError);
+# };
+class XLA_FFI_Api(ctypes.Structure):
+ _fields_ = (
+ ("struct_size", ctypes.c_size_t),
+ ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)),
+ ("api_version", XLA_FFI_Api_Version),
+ ("internal_api", ctypes.c_void_p), # XLA_FFI_InternalApi*
+ ("XLA_FFI_Error_Create", XLA_FFI_Error_Create), # XLA_FFI_Error_Create
+ ("XLA_FFI_Error_GetMessage", ctypes.c_void_p), # XLA_FFI_Error_GetMessage
+ ("XLA_FFI_Error_Destroy", ctypes.c_void_p), # XLA_FFI_Error_Destroy
+ ("XLA_FFI_Handler_Register", ctypes.c_void_p), # XLA_FFI_Handler_Register
+ ("XLA_FFI_Stream_Get", XLA_FFI_Stream_Get), # XLA_FFI_Stream_Get
+ ("XLA_FFI_TypeId_Register", ctypes.c_void_p), # XLA_FFI_TypeId_Register
+ ("XLA_FFI_ExecutionContext_Get", ctypes.c_void_p), # XLA_FFI_ExecutionContext_Get
+ ("XLA_FFI_State_Set", ctypes.c_void_p), # XLA_FFI_State_Set
+ ("XLA_FFI_State_Get", ctypes.c_void_p), # XLA_FFI_State_Get
+ ("XLA_FFI_DeviceMemory_Allocate", ctypes.c_void_p), # XLA_FFI_DeviceMemory_Allocate
+ ("XLA_FFI_DeviceMemory_Free", ctypes.c_void_p), # XLA_FFI_DeviceMemory_Free
+ ("XLA_FFI_ThreadPool_Schedule", ctypes.c_void_p), # XLA_FFI_ThreadPool_Schedule
+ ("XLA_FFI_ThreadPool_NumThreads", ctypes.c_void_p), # XLA_FFI_ThreadPool_NumThreads
+ ("XLA_FFI_Future_Create", ctypes.c_void_p), # XLA_FFI_Future_Create
+ ("XLA_FFI_Future_SetAvailable", ctypes.c_void_p), # XLA_FFI_Future_SetAvailable
+ ("XLA_FFI_Future_SetError", ctypes.c_void_p), # XLA_FFI_Future_SetError
+ )
+
+
+# struct XLA_FFI_CallFrame {
+# size_t struct_size;
+# XLA_FFI_Extension_Base* extension_start;
+# const XLA_FFI_Api* api;
+# XLA_FFI_ExecutionContext* ctx;
+# XLA_FFI_ExecutionStage stage;
+# XLA_FFI_Args args;
+# XLA_FFI_Rets rets;
+# XLA_FFI_Attrs attrs;
+#
+# // XLA FFI handler implementation can use `future` to signal a result of
+# // asynchronous computation to the XLA runtime. XLA runtime will keep all
+# // arguments, results and attributes alive until `future` is completed.
+# XLA_FFI_Future* future; // out
+# };
+class XLA_FFI_CallFrame(ctypes.Structure):
+ _fields_ = (
+ ("struct_size", ctypes.c_size_t),
+ ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)),
+ ("api", ctypes.POINTER(XLA_FFI_Api)),
+ ("ctx", ctypes.c_void_p), # XLA_FFI_ExecutionContext*
+ ("stage", ctypes.c_int), # XLA_FFI_ExecutionStage
+ ("args", XLA_FFI_Args),
+ ("rets", XLA_FFI_Rets),
+ ("attrs", XLA_FFI_Attrs),
+ ("future", ctypes.c_void_p), # XLA_FFI_Future* // out
+ )
+
+
+_xla_data_type_to_constructor = {
+ # XLA_FFI_DataType.INVALID
+ XLA_FFI_DataType.PRED: jnp.bool,
+ XLA_FFI_DataType.S8: jnp.int8,
+ XLA_FFI_DataType.S16: jnp.int16,
+ XLA_FFI_DataType.S32: jnp.int32,
+ XLA_FFI_DataType.S64: jnp.int64,
+ XLA_FFI_DataType.U8: jnp.uint8,
+ XLA_FFI_DataType.U16: jnp.uint16,
+ XLA_FFI_DataType.U32: jnp.uint32,
+ XLA_FFI_DataType.U64: jnp.uint64,
+ XLA_FFI_DataType.F16: jnp.float16,
+ XLA_FFI_DataType.F32: jnp.float32,
+ XLA_FFI_DataType.F64: jnp.float64,
+ XLA_FFI_DataType.BF16: jnp.bfloat16,
+ XLA_FFI_DataType.C64: jnp.complex64,
+ XLA_FFI_DataType.C128: jnp.complex128,
+ # XLA_FFI_DataType.TOKEN
+ XLA_FFI_DataType.F8E5M2: jnp.float8_e5m2,
+ XLA_FFI_DataType.F8E3M4: jnp.float8_e3m4,
+ XLA_FFI_DataType.F8E4M3: jnp.float8_e4m3,
+ XLA_FFI_DataType.F8E4M3FN: jnp.float8_e4m3fn,
+ XLA_FFI_DataType.F8E4M3B11FNUZ: jnp.float8_e4m3b11fnuz,
+ XLA_FFI_DataType.F8E5M2FNUZ: jnp.float8_e5m2fnuz,
+ XLA_FFI_DataType.F8E4M3FNUZ: jnp.float8_e4m3fnuz,
+ # XLA_FFI_DataType.F4E2M1FN: jnp.float4_e2m1fn.dtype,
+ # XLA_FFI_DataType.F8E8M0FNU: jnp.float8_e8m0fnu.dtype,
+}
+
+
+########################################################################
+# Helpers for translating between ctypes and python types
+#######################################################################
+
+
+def decode_bytespan(span: XLA_FFI_ByteSpan):
+ len = span.len
+ chars = ctypes.cast(span.ptr, ctypes.POINTER(ctypes.c_char * len))
+ return chars.contents.value.decode("utf-8")
+
+
+def decode_scalar(scalar: XLA_FFI_Scalar):
+ # TODO validate if dtype supported
+ dtype = jnp.dtype(_xla_data_type_to_constructor[scalar.dtype])
+ bytes = ctypes.string_at(scalar.value, dtype.itemsize)
+ return np.frombuffer(bytes, dtype=dtype).reshape(())
+
+
+def decode_array(array: XLA_FFI_Array):
+ # TODO validate if dtype supported
+ dtype = jnp.dtype(_xla_data_type_to_constructor[array.dtype])
+ bytes = ctypes.string_at(array.data, dtype.itemsize * array.size)
+ return np.frombuffer(bytes, dtype=dtype)
+
+
+def decode_attrs(attrs: XLA_FFI_Attrs):
+ result = {}
+ for i in range(attrs.size):
+ attr_name = decode_bytespan(attrs.names[i].contents)
+ attr_type = attrs.types[i]
+ if attr_type == XLA_FFI_AttrType.STRING:
+ bytespan = ctypes.cast(attrs.attrs[i], ctypes.POINTER(XLA_FFI_ByteSpan))
+ attr_value = decode_bytespan(bytespan.contents)
+ elif attr_type == XLA_FFI_AttrType.SCALAR:
+ attr_value = ctypes.cast(attrs.attrs[i], ctypes.POINTER(XLA_FFI_Scalar))
+ attr_value = decode_scalar(attr_value.contents)
+ elif attr_type == XLA_FFI_AttrType.ARRAY:
+ attr_value = ctypes.cast(attrs.attrs[i], ctypes.POINTER(XLA_FFI_Array))
+ attr_value = decode_array(attr_value.contents)
+ elif attr_type == XLA_FFI_AttrType.DICTIONARY:
+ attr_value = ctypes.cast(attrs.attrs[i], ctypes.POINTER(XLA_FFI_Attrs))
+ attr_value = decode_attrs(attr_value.contents)
+ else:
+ raise Exception("Unexpected attr type")
+ result[attr_name] = attr_value
+ return result
+
+
+# error-string to XLA_FFI_Error
+def create_ffi_error(api, errc, message):
+ create_args = XLA_FFI_Error_Create_Args(
+ ctypes.sizeof(XLA_FFI_Error_Create_Args),
+ ctypes.POINTER(XLA_FFI_Extension_Base)(),
+ ctypes.c_char_p(message.encode("utf-8")),
+ errc,
+ )
+ return api.contents.XLA_FFI_Error_Create(create_args)
+
+
+def create_invalid_argument_ffi_error(api, message):
+ return create_ffi_error(api, XLA_FFI_Error_Code.INVALID_ARGUMENT, message)
+
+
+# Extract CUDA stream from XLA_FFI_CallFrame.
+def get_stream_from_callframe(call_frame):
+ api = call_frame.api
+ get_stream_args = XLA_FFI_Stream_Get_Args(
+ ctypes.sizeof(XLA_FFI_Stream_Get_Args), ctypes.POINTER(XLA_FFI_Extension_Base)(), call_frame.ctx, None
+ )
+ api.contents.XLA_FFI_Stream_Get(get_stream_args)
+ # TODO check result
+ return get_stream_args.stream
+
+
+_dtype_from_ffi = {
+ XLA_FFI_DataType.S8: wp.int8,
+ XLA_FFI_DataType.S16: wp.int16,
+ XLA_FFI_DataType.S32: wp.int32,
+ XLA_FFI_DataType.S64: wp.int64,
+ XLA_FFI_DataType.U8: wp.uint8,
+ XLA_FFI_DataType.U16: wp.uint16,
+ XLA_FFI_DataType.U32: wp.uint32,
+ XLA_FFI_DataType.U64: wp.uint64,
+ XLA_FFI_DataType.F16: wp.float16,
+ XLA_FFI_DataType.F32: wp.float32,
+ XLA_FFI_DataType.F64: wp.float64,
+}
+
+
+def dtype_from_ffi(ffi_dtype):
+ return _dtype_from_ffi.get(ffi_dtype)
+
+
+def jax_dtype_from_ffi(ffi_dtype):
+ return _xla_data_type_to_constructor.get(ffi_dtype)
+
+
+# Execution context (stream, stage)
+class ExecutionContext:
+ stage: XLA_FFI_ExecutionStage
+ stream: int
+
+ def __init__(self, callframe: XLA_FFI_CallFrame):
+ self.stage = XLA_FFI_ExecutionStage(callframe.stage)
+ self.stream = get_stream_from_callframe(callframe)
+
+
+class FfiBuffer:
+ dtype: str
+ data: int
+ shape: tuple[int]
+
+ def __init__(self, xla_buffer):
+ # TODO check if valid
+ self.dtype = jnp.dtype(_xla_data_type_to_constructor[xla_buffer.dtype])
+ self.shape = tuple(xla_buffer.dims[i] for i in range(xla_buffer.rank))
+ self.data = xla_buffer.data
+
+ @property
+ def __cuda_array_interface__(self):
+ return {
+ "shape": self.shape,
+ "typestr": self.dtype.char,
+ "data": (self.data, False),
+ "version": 2,
+ }
diff --git a/mjx/mujoco/mjx/viewer.py b/mjx/mujoco/mjx/viewer.py
index 5cdce300..b8bbff9d 100644
--- a/mjx/mujoco/mjx/viewer.py
+++ b/mjx/mujoco/mjx/viewer.py
@@ -15,7 +15,10 @@
"""An example integration of MJX with the MuJoCo viewer."""
import logging
-import time
+import os
+
+os.environ['XLA_FLAGS'] = '--xla_gpu_graph_min_graph_size=1'
+import time # pylint: disable=g-import-not-at-top
from typing import Sequence
from absl import app
@@ -25,13 +28,29 @@ from jax import numpy as jp
import mujoco
from mujoco import mjx
import mujoco.viewer
+import warp as wp
_JIT = flags.DEFINE_bool('jit', True, 'To jit or not to jit.')
_MODEL_PATH = flags.DEFINE_string(
'mjcf', None, 'Path to a MuJoCo MJCF file.', required=True
)
-
+_IMPL = flags.DEFINE_string('impl', 'jax', 'MJX implementation.')
+_WP_KERNEL_CACHE_DIR = flags.DEFINE_string(
+ 'wp_kernel_cache_dir',
+ None,
+ 'Path to the Warp kernel cache directory.',
+)
+_NCONMAX = flags.DEFINE_integer(
+ 'nconmax',
+ None,
+ 'Maximum number of contacts to simulate, warp only.',
+)
+_NJMAX = flags.DEFINE_integer(
+ 'njmax',
+ None,
+ 'Maximum number of constraints to simulate, warp only.',
+)
_VIEWER_GLOBAL_STATE = {
'running': True,
@@ -49,6 +68,9 @@ def _main(argv: Sequence[str]) -> None:
if len(argv) > 1:
raise app.UsageError('Too many command-line arguments.')
+ if _WP_KERNEL_CACHE_DIR.value:
+ wp.config.kernel_cache_dir = _WP_KERNEL_CACHE_DIR.value
+
jax.config.update('jax_debug_nans', True)
print(f'Loading model from: {_MODEL_PATH.value}.')
@@ -57,8 +79,16 @@ def _main(argv: Sequence[str]) -> None:
else:
m = mujoco.MjModel.from_xml_path(_MODEL_PATH.value)
d = mujoco.MjData(m)
- mx = mjx.put_model(m)
- dx = mjx.put_data(m, d)
+ mx = mjx.put_model(m, impl=_IMPL.value)
+ if _IMPL.value == 'warp':
+ # TODO(btaba): use put_data.
+ dx = mjx.make_data(
+ m, impl=_IMPL.value, nconmax=_NCONMAX.value, njmax=_NJMAX.value
+ )
+ else:
+ dx = mjx.put_data(
+ m, d, impl=_IMPL.value, nconmax=_NCONMAX.value, njmax=_NJMAX.value
+ )
print(f'Default backend: {jax.default_backend()}')
step_fn = mjx.step
diff --git a/mjx/mujoco/mjx/warp/__init__.py b/mjx/mujoco/mjx/warp/__init__.py
new file mode 100644
index 00000000..067adf1e
--- /dev/null
+++ b/mjx/mujoco/mjx/warp/__init__.py
@@ -0,0 +1,70 @@
+# Copyright 2025 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.
+# ==============================================================================
+
+import typing
+from typing import Any
+from mujoco.mjx.warp import types
+
+if not typing.TYPE_CHECKING:
+ # Runtime.
+ warp: Any = None
+ mujoco_warp: Any = None
+ mjwp_types: Any = None
+ WARP_INSTALLED: bool = False
+
+ # pylint: disable=g-import-not-at-top
+ try:
+ import warp
+ WARP_INSTALLED = True
+ except ImportError:
+ WARP_INSTALLED = False
+ try:
+ import mujoco.mjx.third_party.mujoco_warp as mujoco_warp
+ from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
+ except ImportError as e:
+ pass
+ # pylint: enable=g-import-not-at-top
+else:
+ # Only used for type checking.
+ class _WpStub:
+
+ def ScopedDevice(self, device: str): # pylint: disable=invalid-name
+ pass
+
+ def types(self):
+ pass
+
+ class _MjwpStub:
+
+ def put_model(self, *args, **kwargs):
+ pass
+
+ def make_data(self, *args, **kwargs):
+ pass
+
+ def step(self, *args, **kwargs):
+ pass
+
+ class _MjwpTypesStub:
+ def TileSet(self, *args, **kwargs): # pylint: disable=invalid-name
+ pass
+
+ def BlockDim(self, *args, **kwargs): # pylint: disable=invalid-name
+ pass
+
+ WARP_INSTALLED: bool = True
+ warp: Any = _WpStub()
+ mujoco_warp: Any = _MjwpStub()
+ mjwp_types: Any = _MjwpTypesStub()
diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py
new file mode 100644
index 00000000..0b86774e
--- /dev/null
+++ b/mjx/mujoco/mjx/warp/collision_driver.py
@@ -0,0 +1,498 @@
+# Copyright 2025 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.
+# ==============================================================================
+
+"""DO NOT EDIT. This file is auto-generated."""
+import dataclasses
+import jax
+from mujoco.mjx._src import types
+from mujoco.mjx.warp import ffi
+import mujoco.mjx.third_party.mujoco_warp as mjwarp
+from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
+import warp as wp
+
+
+_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 _collision_shim(
+ # Model
+ nworld: int,
+ block_dim: mjwp_types.BlockDim,
+ geom_aabb: wp.array2d(dtype=wp.vec3),
+ geom_condim: wp.array(dtype=int),
+ geom_dataid: wp.array(dtype=int),
+ geom_friction: wp.array2d(dtype=wp.vec3),
+ geom_gap: wp.array2d(dtype=float),
+ geom_margin: wp.array2d(dtype=float),
+ geom_pair_type_count: tuple[int, ...],
+ geom_plugin_index: wp.array(dtype=int),
+ geom_pos: wp.array2d(dtype=wp.vec3),
+ geom_priority: wp.array(dtype=int),
+ geom_quat: wp.array2d(dtype=wp.quat),
+ geom_rbound: wp.array2d(dtype=float),
+ geom_size: wp.array2d(dtype=wp.vec3),
+ geom_solimp: wp.array2d(dtype=mjwp_types.vec5),
+ geom_solmix: wp.array2d(dtype=float),
+ geom_solref: wp.array2d(dtype=wp.vec2),
+ geom_type: wp.array(dtype=int),
+ geompair2hfgeompair: wp.array(dtype=int),
+ has_sdf_geom: bool,
+ hfield_adr: wp.array(dtype=int),
+ hfield_data: wp.array(dtype=float),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ mesh_graph: wp.array(dtype=int),
+ mesh_graphadr: wp.array(dtype=int),
+ mesh_polyadr: wp.array(dtype=int),
+ mesh_polymap: wp.array(dtype=int),
+ mesh_polymapadr: wp.array(dtype=int),
+ mesh_polymapnum: wp.array(dtype=int),
+ mesh_polynormal: wp.array(dtype=wp.vec3),
+ mesh_polynum: wp.array(dtype=int),
+ mesh_polyvert: wp.array(dtype=int),
+ mesh_polyvertadr: wp.array(dtype=int),
+ mesh_polyvertnum: wp.array(dtype=int),
+ mesh_vert: wp.array(dtype=wp.vec3),
+ mesh_vertadr: wp.array(dtype=int),
+ mesh_vertnum: wp.array(dtype=int),
+ ngeom: int,
+ nhfield: int,
+ nxn_geom_pair_filtered: wp.array(dtype=wp.vec2i),
+ nxn_pairid: wp.array(dtype=int),
+ nxn_pairid_filtered: wp.array(dtype=int),
+ pair_dim: wp.array(dtype=int),
+ pair_friction: wp.array2d(dtype=mjwp_types.vec5),
+ pair_gap: wp.array2d(dtype=float),
+ pair_margin: wp.array2d(dtype=float),
+ pair_solimp: wp.array2d(dtype=mjwp_types.vec5),
+ pair_solref: wp.array2d(dtype=wp.vec2),
+ pair_solreffriction: wp.array2d(dtype=wp.vec2),
+ plugin: wp.array(dtype=int),
+ plugin_attr: wp.array(dtype=wp.vec3f),
+ opt__broadphase: int,
+ opt__broadphase_filter: int,
+ opt__disableflags: int,
+ opt__epa_iterations: int,
+ opt__gjk_iterations: int,
+ opt__graph_conditional: bool,
+ opt__sdf_initpoints: int,
+ opt__sdf_iterations: int,
+ # Data
+ nconmax: int,
+ collision_hftri_index: wp.array(dtype=int),
+ collision_pair: wp.array(dtype=wp.vec2i),
+ collision_pairid: wp.array(dtype=int),
+ collision_worldid: wp.array(dtype=int),
+ epa_face: wp.array2d(dtype=wp.vec3i),
+ epa_horizon: wp.array2d(dtype=int),
+ epa_index: wp.array2d(dtype=int),
+ epa_map: wp.array2d(dtype=int),
+ epa_norm2: wp.array2d(dtype=float),
+ epa_pr: wp.array2d(dtype=wp.vec3),
+ epa_vert: wp.array2d(dtype=wp.vec3),
+ epa_vert1: wp.array2d(dtype=wp.vec3),
+ epa_vert2: wp.array2d(dtype=wp.vec3),
+ epa_vert_index1: wp.array2d(dtype=int),
+ epa_vert_index2: wp.array2d(dtype=int),
+ geom_xmat: wp.array2d(dtype=wp.mat33),
+ geom_xpos: wp.array2d(dtype=wp.vec3),
+ ncollision: wp.array(dtype=int),
+ ncon: wp.array(dtype=int),
+ ncon_hfield: wp.array2d(dtype=int),
+ sap_cumulative_sum: wp.array2d(dtype=int),
+ sap_projection_lower: wp.array3d(dtype=float),
+ sap_projection_upper: wp.array2d(dtype=float),
+ sap_range: wp.array2d(dtype=int),
+ sap_segment_index: wp.array2d(dtype=int),
+ sap_sort_index: wp.array3d(dtype=int),
+ contact__dim: wp.array(dtype=int),
+ contact__dist: wp.array(dtype=float),
+ contact__frame: wp.array(dtype=wp.mat33),
+ contact__friction: wp.array(dtype=mjwp_types.vec5),
+ contact__geom: wp.array(dtype=wp.vec2i),
+ contact__includemargin: wp.array(dtype=float),
+ contact__pos: wp.array(dtype=wp.vec3),
+ contact__solimp: wp.array(dtype=mjwp_types.vec5),
+ contact__solref: wp.array(dtype=wp.vec2),
+ contact__solreffriction: wp.array(dtype=wp.vec2),
+ contact__worldid: wp.array(dtype=int),
+):
+ _m.stat = _s
+ _m.opt = _o
+ _d.efc = _e
+ _d.contact = _c
+ _m.block_dim = block_dim
+ _m.geom_aabb = geom_aabb
+ _m.geom_condim = geom_condim
+ _m.geom_dataid = geom_dataid
+ _m.geom_friction = geom_friction
+ _m.geom_gap = geom_gap
+ _m.geom_margin = geom_margin
+ _m.geom_pair_type_count = geom_pair_type_count
+ _m.geom_plugin_index = geom_plugin_index
+ _m.geom_pos = geom_pos
+ _m.geom_priority = geom_priority
+ _m.geom_quat = geom_quat
+ _m.geom_rbound = geom_rbound
+ _m.geom_size = geom_size
+ _m.geom_solimp = geom_solimp
+ _m.geom_solmix = geom_solmix
+ _m.geom_solref = geom_solref
+ _m.geom_type = geom_type
+ _m.geompair2hfgeompair = geompair2hfgeompair
+ _m.has_sdf_geom = has_sdf_geom
+ _m.hfield_adr = hfield_adr
+ _m.hfield_data = hfield_data
+ _m.hfield_ncol = hfield_ncol
+ _m.hfield_nrow = hfield_nrow
+ _m.hfield_size = hfield_size
+ _m.mesh_graph = mesh_graph
+ _m.mesh_graphadr = mesh_graphadr
+ _m.mesh_polyadr = mesh_polyadr
+ _m.mesh_polymap = mesh_polymap
+ _m.mesh_polymapadr = mesh_polymapadr
+ _m.mesh_polymapnum = mesh_polymapnum
+ _m.mesh_polynormal = mesh_polynormal
+ _m.mesh_polynum = mesh_polynum
+ _m.mesh_polyvert = mesh_polyvert
+ _m.mesh_polyvertadr = mesh_polyvertadr
+ _m.mesh_polyvertnum = mesh_polyvertnum
+ _m.mesh_vert = mesh_vert
+ _m.mesh_vertadr = mesh_vertadr
+ _m.mesh_vertnum = mesh_vertnum
+ _m.ngeom = ngeom
+ _m.nhfield = nhfield
+ _m.nxn_geom_pair_filtered = nxn_geom_pair_filtered
+ _m.nxn_pairid = nxn_pairid
+ _m.nxn_pairid_filtered = nxn_pairid_filtered
+ _m.opt.broadphase = opt__broadphase
+ _m.opt.broadphase_filter = opt__broadphase_filter
+ _m.opt.disableflags = opt__disableflags
+ _m.opt.epa_iterations = opt__epa_iterations
+ _m.opt.gjk_iterations = opt__gjk_iterations
+ _m.opt.graph_conditional = opt__graph_conditional
+ _m.opt.sdf_initpoints = opt__sdf_initpoints
+ _m.opt.sdf_iterations = opt__sdf_iterations
+ _m.pair_dim = pair_dim
+ _m.pair_friction = pair_friction
+ _m.pair_gap = pair_gap
+ _m.pair_margin = pair_margin
+ _m.pair_solimp = pair_solimp
+ _m.pair_solref = pair_solref
+ _m.pair_solreffriction = pair_solreffriction
+ _m.plugin = plugin
+ _m.plugin_attr = plugin_attr
+ _d.collision_hftri_index = collision_hftri_index
+ _d.collision_pair = collision_pair
+ _d.collision_pairid = collision_pairid
+ _d.collision_worldid = collision_worldid
+ _d.contact.dim = contact__dim
+ _d.contact.dist = contact__dist
+ _d.contact.frame = contact__frame
+ _d.contact.friction = contact__friction
+ _d.contact.geom = contact__geom
+ _d.contact.includemargin = contact__includemargin
+ _d.contact.pos = contact__pos
+ _d.contact.solimp = contact__solimp
+ _d.contact.solref = contact__solref
+ _d.contact.solreffriction = contact__solreffriction
+ _d.contact.worldid = contact__worldid
+ _d.epa_face = epa_face
+ _d.epa_horizon = epa_horizon
+ _d.epa_index = epa_index
+ _d.epa_map = epa_map
+ _d.epa_norm2 = epa_norm2
+ _d.epa_pr = epa_pr
+ _d.epa_vert = epa_vert
+ _d.epa_vert1 = epa_vert1
+ _d.epa_vert2 = epa_vert2
+ _d.epa_vert_index1 = epa_vert_index1
+ _d.epa_vert_index2 = epa_vert_index2
+ _d.geom_xmat = geom_xmat
+ _d.geom_xpos = geom_xpos
+ _d.ncollision = ncollision
+ _d.ncon = ncon
+ _d.ncon_hfield = ncon_hfield
+ _d.nconmax = nconmax
+ _d.sap_cumulative_sum = sap_cumulative_sum
+ _d.sap_projection_lower = sap_projection_lower
+ _d.sap_projection_upper = sap_projection_upper
+ _d.sap_range = sap_range
+ _d.sap_segment_index = sap_segment_index
+ _d.sap_sort_index = sap_sort_index
+ _d.nworld = nworld
+ mjwarp.collision(_m, _d)
+
+
+def _collision_jax_impl(m: types.Model, d: types.Data):
+ output_dims = {
+ 'collision_hftri_index': d._impl.collision_hftri_index.shape,
+ 'collision_pair': d._impl.collision_pair.shape,
+ 'collision_pairid': d._impl.collision_pairid.shape,
+ 'collision_worldid': d._impl.collision_worldid.shape,
+ 'epa_face': d._impl.epa_face.shape,
+ 'epa_horizon': d._impl.epa_horizon.shape,
+ 'epa_index': d._impl.epa_index.shape,
+ 'epa_map': d._impl.epa_map.shape,
+ 'epa_norm2': d._impl.epa_norm2.shape,
+ 'epa_pr': d._impl.epa_pr.shape,
+ 'epa_vert': d._impl.epa_vert.shape,
+ 'epa_vert1': d._impl.epa_vert1.shape,
+ 'epa_vert2': d._impl.epa_vert2.shape,
+ 'epa_vert_index1': d._impl.epa_vert_index1.shape,
+ 'epa_vert_index2': d._impl.epa_vert_index2.shape,
+ 'geom_xmat': d.geom_xmat.shape,
+ 'geom_xpos': d.geom_xpos.shape,
+ 'ncollision': d._impl.ncollision.shape,
+ 'ncon': d._impl.ncon.shape,
+ 'ncon_hfield': d._impl.ncon_hfield.shape,
+ 'sap_cumulative_sum': d._impl.sap_cumulative_sum.shape,
+ 'sap_projection_lower': d._impl.sap_projection_lower.shape,
+ 'sap_projection_upper': d._impl.sap_projection_upper.shape,
+ 'sap_range': d._impl.sap_range.shape,
+ 'sap_segment_index': d._impl.sap_segment_index.shape,
+ 'sap_sort_index': d._impl.sap_sort_index.shape,
+ 'contact__dim': d._impl.contact__dim.shape,
+ 'contact__dist': d._impl.contact__dist.shape,
+ 'contact__frame': d._impl.contact__frame.shape,
+ 'contact__friction': d._impl.contact__friction.shape,
+ 'contact__geom': d._impl.contact__geom.shape,
+ 'contact__includemargin': d._impl.contact__includemargin.shape,
+ 'contact__pos': d._impl.contact__pos.shape,
+ 'contact__solimp': d._impl.contact__solimp.shape,
+ 'contact__solref': d._impl.contact__solref.shape,
+ 'contact__solreffriction': d._impl.contact__solreffriction.shape,
+ 'contact__worldid': d._impl.contact__worldid.shape,
+ }
+ jf = ffi.jax_callable_variadic_tuple(
+ _collision_shim,
+ num_outputs=37,
+ output_dims=output_dims,
+ vmap_method=None,
+ graph_compatible=True,
+ in_out_argnames={
+ 'collision_hftri_index',
+ 'collision_pair',
+ 'collision_pairid',
+ 'collision_worldid',
+ 'epa_face',
+ 'epa_horizon',
+ 'epa_index',
+ 'epa_map',
+ 'epa_norm2',
+ 'epa_pr',
+ 'epa_vert',
+ 'epa_vert1',
+ 'epa_vert2',
+ 'epa_vert_index1',
+ 'epa_vert_index2',
+ 'geom_xmat',
+ 'geom_xpos',
+ 'ncollision',
+ 'ncon',
+ 'ncon_hfield',
+ 'sap_cumulative_sum',
+ 'sap_projection_lower',
+ 'sap_projection_upper',
+ 'sap_range',
+ 'sap_segment_index',
+ 'sap_sort_index',
+ 'contact__dim',
+ 'contact__dist',
+ 'contact__frame',
+ 'contact__friction',
+ 'contact__geom',
+ 'contact__includemargin',
+ 'contact__pos',
+ 'contact__solimp',
+ 'contact__solref',
+ 'contact__solreffriction',
+ 'contact__worldid',
+ },
+ )
+ out = jf(
+ d.qpos.shape[0],
+ m._impl.block_dim,
+ m.geom_aabb,
+ m.geom_condim,
+ m.geom_dataid,
+ m.geom_friction,
+ m.geom_gap,
+ m.geom_margin,
+ m._impl.geom_pair_type_count,
+ m._impl.geom_plugin_index,
+ m.geom_pos,
+ m.geom_priority,
+ m.geom_quat,
+ m.geom_rbound,
+ m.geom_size,
+ m.geom_solimp,
+ m.geom_solmix,
+ m.geom_solref,
+ m.geom_type,
+ m._impl.geompair2hfgeompair,
+ m._impl.has_sdf_geom,
+ m.hfield_adr,
+ m.hfield_data,
+ m.hfield_ncol,
+ m.hfield_nrow,
+ m.hfield_size,
+ m.mesh_graph,
+ m.mesh_graphadr,
+ m._impl.mesh_polyadr,
+ m._impl.mesh_polymap,
+ m._impl.mesh_polymapadr,
+ m._impl.mesh_polymapnum,
+ m._impl.mesh_polynormal,
+ m._impl.mesh_polynum,
+ m._impl.mesh_polyvert,
+ m._impl.mesh_polyvertadr,
+ m._impl.mesh_polyvertnum,
+ m.mesh_vert,
+ m.mesh_vertadr,
+ m.mesh_vertnum,
+ m.ngeom,
+ m.nhfield,
+ m._impl.nxn_geom_pair_filtered,
+ m._impl.nxn_pairid,
+ m._impl.nxn_pairid_filtered,
+ m.pair_dim,
+ m.pair_friction,
+ m.pair_gap,
+ m.pair_margin,
+ m.pair_solimp,
+ m.pair_solref,
+ m.pair_solreffriction,
+ m._impl.plugin,
+ m._impl.plugin_attr,
+ m.opt._impl.broadphase,
+ m.opt._impl.broadphase_filter,
+ m.opt.disableflags,
+ m.opt._impl.epa_iterations,
+ m.opt._impl.gjk_iterations,
+ m.opt._impl.graph_conditional,
+ m.opt._impl.sdf_initpoints,
+ m.opt._impl.sdf_iterations,
+ d._impl.nconmax,
+ d._impl.collision_hftri_index,
+ d._impl.collision_pair,
+ d._impl.collision_pairid,
+ d._impl.collision_worldid,
+ d._impl.epa_face,
+ d._impl.epa_horizon,
+ d._impl.epa_index,
+ d._impl.epa_map,
+ d._impl.epa_norm2,
+ d._impl.epa_pr,
+ d._impl.epa_vert,
+ d._impl.epa_vert1,
+ d._impl.epa_vert2,
+ d._impl.epa_vert_index1,
+ d._impl.epa_vert_index2,
+ d.geom_xmat,
+ d.geom_xpos,
+ d._impl.ncollision,
+ d._impl.ncon,
+ d._impl.ncon_hfield,
+ d._impl.sap_cumulative_sum,
+ d._impl.sap_projection_lower,
+ d._impl.sap_projection_upper,
+ d._impl.sap_range,
+ d._impl.sap_segment_index,
+ d._impl.sap_sort_index,
+ d._impl.contact__dim,
+ d._impl.contact__dist,
+ d._impl.contact__frame,
+ d._impl.contact__friction,
+ d._impl.contact__geom,
+ d._impl.contact__includemargin,
+ d._impl.contact__pos,
+ d._impl.contact__solimp,
+ d._impl.contact__solref,
+ d._impl.contact__solreffriction,
+ d._impl.contact__worldid,
+ )
+ d = d.tree_replace({
+ '_impl.collision_hftri_index': out[0],
+ '_impl.collision_pair': out[1],
+ '_impl.collision_pairid': out[2],
+ '_impl.collision_worldid': out[3],
+ '_impl.epa_face': out[4],
+ '_impl.epa_horizon': out[5],
+ '_impl.epa_index': out[6],
+ '_impl.epa_map': out[7],
+ '_impl.epa_norm2': out[8],
+ '_impl.epa_pr': out[9],
+ '_impl.epa_vert': out[10],
+ '_impl.epa_vert1': out[11],
+ '_impl.epa_vert2': out[12],
+ '_impl.epa_vert_index1': out[13],
+ '_impl.epa_vert_index2': out[14],
+ 'geom_xmat': out[15],
+ 'geom_xpos': out[16],
+ '_impl.ncollision': out[17],
+ '_impl.ncon': out[18],
+ '_impl.ncon_hfield': out[19],
+ '_impl.sap_cumulative_sum': out[20],
+ '_impl.sap_projection_lower': out[21],
+ '_impl.sap_projection_upper': out[22],
+ '_impl.sap_range': out[23],
+ '_impl.sap_segment_index': out[24],
+ '_impl.sap_sort_index': out[25],
+ '_impl.contact__dim': out[26],
+ '_impl.contact__dist': out[27],
+ '_impl.contact__frame': out[28],
+ '_impl.contact__friction': out[29],
+ '_impl.contact__geom': out[30],
+ '_impl.contact__includemargin': out[31],
+ '_impl.contact__pos': out[32],
+ '_impl.contact__solimp': out[33],
+ '_impl.contact__solref': out[34],
+ '_impl.contact__solreffriction': out[35],
+ '_impl.contact__worldid': out[36],
+ })
+ return d
+
+
+@jax.custom_batching.custom_vmap
+@ffi.marshal_jax_warp_callable
+def collision(m: types.Model, d: types.Data):
+ return _collision_jax_impl(m, d)
+
+
+@collision.def_vmap
+@ffi.marshal_custom_vmap
+def collision_vmap(unused_axis_size, is_batched, m, d):
+ d = collision(m, d)
+ return d, is_batched[1]
diff --git a/mjx/mujoco/mjx/warp/collision_driver_test.py b/mjx/mujoco/mjx/warp/collision_driver_test.py
new file mode 100644
index 00000000..e69227f9
--- /dev/null
+++ b/mjx/mujoco/mjx/warp/collision_driver_test.py
@@ -0,0 +1,107 @@
+# Copyright 2025 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 collision driver."""
+
+from absl.testing import absltest
+import jax
+import mujoco
+from mujoco import mjx
+from mujoco.mjx._src import io
+import mujoco.mjx.warp as mjxw
+from mujoco.mjx.warp import test_util as tu
+from mujoco.mjx.warp import warp as wp # pylint: disable=g-importing-member
+import numpy as np
+
+
+try:
+ from mujoco.mjx.warp import collision_driver # pylint: disable=g-import-not-at-top
+ from mujoco.mjx.warp import smooth # pylint: disable=g-import-not-at-top
+except ImportError:
+ collision_driver = None
+ smooth = None
+
+
+class CollisionTest(absltest.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ if mjxw.WARP_INSTALLED:
+ wp.clear_kernel_cache()
+ np.random.seed(0)
+
+ _SPHERE_SPHERE = """
+
+
+
+
+
+
+
+
+
+
+
+
+ """
+
+ def test_collision_nested_vmap(self):
+ """Tests collision with batched data."""
+ if not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp not installed.')
+ if not io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device available.')
+
+ m = mujoco.MjModel.from_xml_string(self._SPHERE_SPHERE)
+ d = mujoco.MjData(m)
+ mx = mjx.put_model(m, impl='warp')
+
+ def make_data(rng):
+ dx = mjx.make_data(m, impl='warp')
+ _, key = jax.random.split(rng)
+ qpos = jax.random.uniform(key, (m.nq,), minval=-0.01, maxval=0.01)
+ return dx.replace(
+ qpos=qpos,
+ )
+
+ rng = jax.random.split(jax.random.PRNGKey(0), 8)
+ rng = rng.reshape((2, 4, -1))
+ dx_batch = jax.vmap(jax.vmap(make_data))(rng)
+
+ dx_batch = jax.jit(
+ jax.vmap(
+ jax.vmap(smooth.kinematics, in_axes=(None, 0)), in_axes=(None, 0)
+ )
+ )(mx, dx_batch)
+ dx_batch = jax.jit(
+ jax.vmap(
+ jax.vmap(collision_driver.collision, in_axes=(None, 0)),
+ in_axes=(None, 0),
+ )
+ )(mx, dx_batch)
+
+ for i in range(2):
+ for j in range(4):
+ dx = dx_batch[i, j]
+
+ d.qpos[:] = dx.qpos
+ mujoco.mj_forward(m, d)
+
+ if not d.contact.pos.shape[0]:
+ continue
+ tu.assert_contact_eq(d, dx, worldid=i * 4 + j)
+
+
+if __name__ == '__main__':
+ absltest.main()
diff --git a/mjx/mujoco/mjx/warp/ffi.py b/mjx/mujoco/mjx/warp/ffi.py
new file mode 100644
index 00000000..e6db9595
--- /dev/null
+++ b/mjx/mujoco/mjx/warp/ffi.py
@@ -0,0 +1,350 @@
+# Copyright 2025 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.
+# ==============================================================================
+"""FFI helper functions for MJX."""
+
+import dataclasses
+import functools
+import inspect
+import typing
+from typing import Any, Callable, Optional, Sequence, Tuple, Union
+
+import jax
+from jax import numpy as jp
+from mujoco.mjx.warp import types as mjx_warp_types
+import numpy as np
+import warp as wp
+from mujoco.mjx.third_party.warp.jax_experimental import ffi
+
+
+def flatten_signature(signature: inspect.Signature, args: Tuple[Any, ...]):
+ """Flattens a tuple/dataclass signature."""
+
+ def expand_parameter(parameter, arg_iter):
+ p = parameter
+ if p.kind != inspect.Parameter.POSITIONAL_OR_KEYWORD:
+ raise ValueError(f'Unsupported parameter kind: {p.kind}')
+
+ try:
+ arg = next(arg_iter)
+ except StopIteration:
+ # We ran out input arguments. Let us keep output parameters as is.
+ assert typing.get_origin(p.annotation) != tuple, p.annotation
+ return [p]
+
+ # If it is a tuple, we need to duplicate the parameter for each element.
+ if isinstance(arg, tuple):
+ assert typing.get_origin(p.annotation) == tuple, p.annotation
+ type_args = typing.get_args(p.annotation)
+ if len(type_args) != 2 or type_args[1] != ...:
+ raise NotImplementedError(
+ f'Unsupported tuple argument: {type_args} '
+ '(currently, only Tuple[t, ...] is supported).'
+ )
+ type_ = type_args[0]
+ types = [type_]
+ if dataclasses.is_dataclass(type_):
+ # If the tuple element is a dataclass, we need to recycle the
+ # types in the same order as they appear in the dataclass.
+ fields = list(type_.__dataclass_fields__.values())
+ types = [f.type for f in fields]
+ return [
+ inspect.Parameter(
+ f'{p.name}__{i}',
+ p.kind,
+ default=p.default,
+ annotation=types[i % len(types)],
+ )
+ for i in range(len(arg) * len(types))
+ ]
+ elif dataclasses.is_dataclass(arg):
+ assert dataclasses.is_dataclass(p.annotation), p.annotation
+ fields = list(arg.__dataclass_fields__.values())
+ return [
+ inspect.Parameter(
+ f'{p.name}__{fields[i].name}',
+ p.kind,
+ default=p.default,
+ annotation=fields[i].type,
+ )
+ for i in range(len(fields))
+ ]
+
+ assert typing.get_origin(p.annotation) != tuple, p.annotation
+ return [p]
+
+ parameters = []
+ arg_iter = iter(args)
+ for p in signature.parameters.values():
+ parameters.extend(expand_parameter(p, arg_iter))
+
+ return inspect.Signature(
+ parameters=parameters, return_annotation=signature.return_annotation
+ )
+
+
+def jax_callable_variadic_tuple(
+ func: Callable, # pylint: disable=g-bare-generic
+ num_outputs: int = 1,
+ graph_compatible: bool = True,
+ vmap_method: Optional[str] = None,
+ output_dims: Optional[dict[str, tuple[int, ...]]] = None,
+ in_out_argnames: Optional[Sequence[str]] = None,
+):
+ """Wraps a JAX callable to support variadic tuples and dataclasses."""
+
+ def callable_wrapper(*args, **kwargs):
+ def func_wrapper(*flat_args, **kwargs):
+ unflat_args = jax.tree.unflatten(in_tree, flat_args)
+ return func(*unflat_args, **kwargs)
+
+ # Provide a flattened signature for the Warp callable machinery.
+ func_wrapper.__signature__ = flatten_signature(
+ inspect.signature(func), args
+ )
+ my_callable = ffi.jax_callable(
+ func_wrapper,
+ num_outputs=num_outputs,
+ graph_compatible=graph_compatible,
+ vmap_method=vmap_method,
+ output_dims=output_dims,
+ in_out_argnames=in_out_argnames,
+ )
+
+ flat_args, in_tree = jax.tree.flatten(args)
+ return my_callable(*flat_args, **kwargs)
+
+ return callable_wrapper
+
+
+def _format_arg(arg: Any, name: str, annotation: Any, verbose: bool):
+ """Formats a single argument for warp."""
+ typ_args = typing.get_args(annotation)
+ annotation_origin = typing.get_origin(annotation)
+
+ # Handle variadic tuples.
+ if annotation_origin == tuple and len(typ_args) == 2 and typ_args[1] == ...:
+ return tuple(
+ _format_arg(arg[i], name + f'_{i}', typ_args[0], verbose)
+ for i in range(len(arg))
+ )
+
+ if not isinstance(annotation, wp.types.array):
+ if verbose:
+ print(f'Skipping {name}: {arg}')
+ return arg
+
+ expected_ndim = annotation.ndim
+ if arg.ndim != expected_ndim:
+ raise AssertionError(
+ f'Arg ndim {arg.ndim} does not match expected ndim {expected_ndim}.'
+ )
+
+ # Add stride 0 to first axis in case the underlying argument should be
+ # batched.
+ # NB: the outer marshalling does an "expand_dims" on Model fields.
+ is_batch_field = mjx_warp_types.BATCH_DIM['Model'].get(name, False)
+ if arg.shape[0] == 1 and is_batch_field:
+ old_strides = arg.strides
+ arg.strides = (0,) + arg.strides[1:]
+ if verbose:
+ print(
+ f'Leading batch dim of 1, adding stride: {name} {old_strides} =>'
+ f' {arg.strides}'
+ )
+ return arg
+
+ if verbose:
+ print(f'Did nothing: {name}: {arg.shape}')
+ return arg
+
+
+def format_args_for_warp(func, verbose=False):
+ @functools.wraps(func)
+ def wrapper(*args):
+ args = list(args)
+ annotations = func.__annotations__
+ assert len(args) == len(annotations)
+ for i, (name, annotation) in enumerate(annotations.items()):
+ args[i] = _format_arg(args[i], name, annotation, verbose)
+ return func(*args)
+
+ return wrapper
+
+
+def _get_mapping_from_tree_path(
+ path: jax.tree_util.KeyPath,
+ mapping: dict[str, int],
+) -> Optional[int]:
+ """Gets the mapped value from a tree path."""
+ if not isinstance(path, tuple):
+ raise NotImplementedError(
+ f'Parsing for jax tree path {path} not implemented.'
+ )
+
+ if any(isinstance(p, jax.tree_util.SequenceKey) for p in path):
+ # get the path up to the first sequence key, we assume variadic sequences
+ is_seq_key = [isinstance(p, jax.tree_util.SequenceKey) for p in path]
+ path = path[: is_seq_key.index(True)]
+
+ assert all(isinstance(p, jax.tree_util.GetAttrKey) for p in path)
+ path = [p for p in path if p.name != '_impl']
+ attr = '__'.join(p.name for p in path)
+
+ # None if the MJX public field is not present in the MJX-Warp mapping.
+ return mapping.get(attr)
+
+
+def _expand_dim_from_path(
+ path: jax.tree_util.KeyPath, leaf: Any, ndim_map: dict[str, int]
+) -> Any:
+ """Expands the dimension of a leaf node based on the ndim_map."""
+ ndim = _get_mapping_from_tree_path(path, ndim_map)
+ if ndim is None or ndim < 0:
+ return leaf
+ if ndim > leaf.ndim:
+ leaf = jp.expand_dims(leaf, axis=np.arange(ndim - leaf.ndim))
+ if ndim != leaf.ndim:
+ raise AssertionError(
+ f'Leaf node ndim ({leaf.ndim}) and expected ndim ({ndim}) do not match'
+ f' for path {path}.'
+ )
+ return leaf
+
+
+def _squeeze_dim(leaf_expanded: Any, leaf: Any) -> Any:
+ if leaf_expanded.ndim < leaf.ndim:
+ raise AssertionError(
+ f'Expanded leaf ndim {leaf_expanded.ndim} is smaller than original leaf'
+ f' ndim {leaf.ndim}'
+ )
+ if leaf_expanded.ndim > leaf.ndim:
+ return jp.squeeze(leaf_expanded, np.arange(leaf_expanded.ndim - leaf.ndim))
+ return leaf_expanded
+
+
+def marshal_jax_warp_callable(func):
+ """Marshal fields into a MuJoCo Warp function."""
+
+ @functools.wraps(func)
+ def wrapper(m, d):
+ # Expand dims for Warp implicit vmap before calling into the FFI wrapped
+ # function.
+ m_expanded = jax.tree.map_with_path(
+ lambda path, x: _expand_dim_from_path(
+ path, x, mjx_warp_types.NDIM['Model']
+ ),
+ m,
+ )
+ d_expanded = jax.tree.map_with_path(
+ lambda path, x: _expand_dim_from_path(
+ path, x, mjx_warp_types.NDIM['Data']
+ ),
+ d,
+ )
+ d_expanded_result = func(m_expanded, d_expanded)
+ d_result = jax.tree.map(_squeeze_dim, d_expanded_result, d)
+ return d_result
+
+ return wrapper
+
+
+def _flatten_batch_dim(
+ path: jax.tree_util.KeyPath, leaf: Any, ndim_map: dict[str, int]
+) -> Any:
+ ndim = _get_mapping_from_tree_path(path, ndim_map)
+ if ndim is None or ndim < 0:
+ return leaf
+ if ndim < leaf.ndim:
+ assert leaf.ndim - ndim == 1
+ batch_dim = np.prod(leaf.shape[: leaf.ndim - ndim + 1])
+ leaf = jp.reshape(leaf, (batch_dim,) + leaf.shape[leaf.ndim - ndim + 1 :])
+ return leaf
+
+
+def _unflatten_batch_dim(leaf_squeezed: Any, leaf: Any) -> Any:
+ if leaf_squeezed.ndim > leaf.ndim:
+ raise AssertionError(
+ f'Squeezed leaf ndim {leaf_squeezed.ndim} is greater than original leaf'
+ f' ndim {leaf.ndim}'
+ )
+ if leaf_squeezed.ndim < leaf.ndim:
+ return leaf_squeezed.reshape(leaf.shape)
+ return leaf_squeezed
+
+
+def _maybe_broadcast_to(
+ path: jax.tree_util.KeyPath,
+ leaf: Any,
+ is_batched: bool,
+ axis_size: Union[int, tuple[int, ...]],
+ cls_str: str,
+) -> Any:
+ """Broadcasts fields that are used in MuJoCo Warp."""
+ ndim = _get_mapping_from_tree_path(path, mjx_warp_types.NDIM[cls_str])
+ needs_batch_dim = _get_mapping_from_tree_path(
+ path, mjx_warp_types.BATCH_DIM[cls_str]
+ )
+ needs_batch_dim = bool(needs_batch_dim) and (ndim is not None and ndim > 0)
+ if needs_batch_dim and not is_batched:
+ leaf = jp.broadcast_to(leaf, (axis_size,) + leaf.shape)
+ return leaf
+
+
+def marshal_custom_vmap(vmap_func):
+ """Marshal fields for a custom vmap into an MuJoCo Warp function."""
+
+ @functools.wraps(vmap_func)
+ def wrapper(axis_size, is_batched, m, d):
+ # Vmappable data fields may not have been broadcasted if vmap_func is called
+ # within a vmap trace. Since data fields are read/write in warp, we need to
+ # explicitly broadcast them here.
+ d_broadcast = jax.tree.map_with_path(
+ lambda path, x, is_b: _maybe_broadcast_to(
+ path, x, is_b, axis_size, 'Data'
+ ),
+ d, is_batched[1], # fmt: skip
+ )
+ # Flatten batch dims into the first axis if the vmap was nested.
+ m_flat = jax.tree.map_with_path(
+ lambda path, x: _flatten_batch_dim(
+ path, x, mjx_warp_types.NDIM['Model']
+ ),
+ m,
+ )
+ d_broadcast_flat = jax.tree.map_with_path(
+ lambda path, x: _flatten_batch_dim(
+ path, x, mjx_warp_types.NDIM['Data']
+ ),
+ d_broadcast,
+ )
+ d_broadcast_flat_result, out_batched = vmap_func(
+ axis_size, is_batched, m_flat, d_broadcast_flat
+ )
+ # Explicitly mark MuJoCo Warp data fields as batched after vmapping is done.
+ out_batched = jax.tree.map_with_path(
+ # NB: if a field is not in MuJoCo Warp, we let JAX do its magic.
+ lambda path, x: _get_mapping_from_tree_path(
+ path, mjx_warp_types.BATCH_DIM['Data']
+ )
+ or x,
+ out_batched,
+ )
+ # Unflatten batch dimensions but keep the broadcasting.
+ d_result = jax.tree.map(
+ _unflatten_batch_dim, d_broadcast_flat_result, d_broadcast
+ )
+ return d_result, out_batched
+
+ return wrapper
diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py
new file mode 100644
index 00000000..8238c24d
--- /dev/null
+++ b/mjx/mujoco/mjx/warp/forward.py
@@ -0,0 +1,4123 @@
+# Copyright 2025 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.
+# ==============================================================================
+
+"""DO NOT EDIT. This file is auto-generated."""
+import dataclasses
+import jax
+from mujoco.mjx._src import types
+from mujoco.mjx.warp import ffi
+import mujoco.mjx.third_party.mujoco_warp as mjwarp
+from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
+import warp as wp
+
+
+_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 _forward_shim(
+ # Model
+ nworld: int,
+ M_rowadr: wp.array(dtype=int),
+ M_rownnz: wp.array(dtype=int),
+ actuator_acc0: wp.array(dtype=float),
+ actuator_actadr: wp.array(dtype=int),
+ actuator_actearly: wp.array(dtype=bool),
+ actuator_actlimited: wp.array(dtype=bool),
+ actuator_actnum: wp.array(dtype=int),
+ actuator_actrange: wp.array2d(dtype=wp.vec2),
+ actuator_biasprm: wp.array2d(dtype=mjwp_types.vec10f),
+ actuator_biastype: wp.array(dtype=int),
+ actuator_cranklength: wp.array(dtype=float),
+ actuator_ctrllimited: wp.array(dtype=bool),
+ actuator_ctrlrange: wp.array2d(dtype=wp.vec2),
+ actuator_dynprm: wp.array2d(dtype=mjwp_types.vec10f),
+ actuator_dyntype: wp.array(dtype=int),
+ actuator_forcelimited: wp.array(dtype=bool),
+ actuator_forcerange: wp.array2d(dtype=wp.vec2),
+ actuator_gainprm: wp.array2d(dtype=mjwp_types.vec10f),
+ actuator_gaintype: wp.array(dtype=int),
+ actuator_gear: wp.array2d(dtype=wp.spatial_vector),
+ actuator_lengthrange: wp.array(dtype=wp.vec2),
+ actuator_trnid: wp.array(dtype=wp.vec2i),
+ actuator_trntype: wp.array(dtype=int),
+ actuator_trntype_body_adr: wp.array(dtype=int),
+ block_dim: mjwp_types.BlockDim,
+ body_dofadr: wp.array(dtype=int),
+ body_dofnum: wp.array(dtype=int),
+ body_gravcomp: wp.array2d(dtype=float),
+ body_inertia: wp.array2d(dtype=wp.vec3),
+ body_invweight0: wp.array2d(dtype=wp.vec2),
+ body_ipos: wp.array2d(dtype=wp.vec3),
+ body_iquat: wp.array2d(dtype=wp.quat),
+ body_jntadr: wp.array(dtype=int),
+ body_jntnum: wp.array(dtype=int),
+ body_mass: wp.array2d(dtype=float),
+ body_parentid: wp.array(dtype=int),
+ body_pos: wp.array2d(dtype=wp.vec3),
+ body_quat: wp.array2d(dtype=wp.quat),
+ body_rootid: wp.array(dtype=int),
+ body_subtreemass: wp.array2d(dtype=float),
+ body_tree: tuple[wp.array(dtype=int), ...],
+ body_weldid: wp.array(dtype=int),
+ cam_bodyid: wp.array(dtype=int),
+ cam_fovy: wp.array(dtype=float),
+ cam_intrinsic: wp.array(dtype=wp.vec4),
+ cam_mat0: wp.array2d(dtype=wp.mat33),
+ cam_mode: wp.array(dtype=int),
+ cam_pos: wp.array2d(dtype=wp.vec3),
+ cam_pos0: wp.array2d(dtype=wp.vec3),
+ cam_poscom0: wp.array2d(dtype=wp.vec3),
+ cam_quat: wp.array2d(dtype=wp.quat),
+ cam_resolution: wp.array(dtype=wp.vec2i),
+ cam_sensorsize: wp.array(dtype=wp.vec2),
+ cam_targetbodyid: wp.array(dtype=int),
+ condim_max: int,
+ dof_Madr: wp.array(dtype=int),
+ dof_armature: wp.array2d(dtype=float),
+ dof_bodyid: wp.array(dtype=int),
+ dof_damping: wp.array2d(dtype=float),
+ dof_frictionloss: wp.array2d(dtype=float),
+ dof_invweight0: wp.array2d(dtype=float),
+ dof_jntid: wp.array(dtype=int),
+ dof_parentid: wp.array(dtype=int),
+ dof_solimp: wp.array2d(dtype=mjwp_types.vec5),
+ dof_solref: wp.array2d(dtype=wp.vec2),
+ dof_tri_col: wp.array(dtype=int),
+ dof_tri_row: wp.array(dtype=int),
+ eq_connect_adr: wp.array(dtype=int),
+ eq_data: wp.array2d(dtype=mjwp_types.vec11),
+ eq_jnt_adr: wp.array(dtype=int),
+ eq_obj1id: wp.array(dtype=int),
+ eq_obj2id: wp.array(dtype=int),
+ eq_objtype: wp.array(dtype=int),
+ eq_solimp: wp.array2d(dtype=mjwp_types.vec5),
+ eq_solref: wp.array2d(dtype=wp.vec2),
+ eq_ten_adr: wp.array(dtype=int),
+ eq_wld_adr: wp.array(dtype=int),
+ flex_bending: wp.array(dtype=wp.mat44f),
+ flex_damping: wp.array(dtype=float),
+ flex_dim: wp.array(dtype=int),
+ flex_edge: wp.array(dtype=wp.vec2i),
+ flex_edgeadr: wp.array(dtype=int),
+ flex_edgeflap: wp.array(dtype=wp.vec2i),
+ flex_elem: wp.array(dtype=int),
+ flex_elemedge: wp.array(dtype=int),
+ flex_elemedgeadr: wp.array(dtype=int),
+ flex_stiffness: wp.array(dtype=float),
+ flex_vertadr: wp.array(dtype=int),
+ flex_vertbodyid: wp.array(dtype=int),
+ flexedge_length0: wp.array(dtype=float),
+ geom_aabb: wp.array2d(dtype=wp.vec3),
+ geom_bodyid: wp.array(dtype=int),
+ geom_condim: wp.array(dtype=int),
+ geom_dataid: wp.array(dtype=int),
+ geom_friction: wp.array2d(dtype=wp.vec3),
+ geom_gap: wp.array2d(dtype=float),
+ geom_group: wp.array(dtype=int),
+ geom_margin: wp.array2d(dtype=float),
+ geom_matid: wp.array2d(dtype=int),
+ geom_pair_type_count: tuple[int, ...],
+ geom_plugin_index: wp.array(dtype=int),
+ geom_pos: wp.array2d(dtype=wp.vec3),
+ geom_priority: wp.array(dtype=int),
+ geom_quat: wp.array2d(dtype=wp.quat),
+ geom_rbound: wp.array2d(dtype=float),
+ geom_rgba: wp.array2d(dtype=wp.vec4),
+ geom_size: wp.array2d(dtype=wp.vec3),
+ geom_solimp: wp.array2d(dtype=mjwp_types.vec5),
+ geom_solmix: wp.array2d(dtype=float),
+ geom_solref: wp.array2d(dtype=wp.vec2),
+ geom_type: wp.array(dtype=int),
+ geompair2hfgeompair: wp.array(dtype=int),
+ has_sdf_geom: bool,
+ hfield_adr: wp.array(dtype=int),
+ hfield_data: wp.array(dtype=float),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ jnt_actfrclimited: wp.array(dtype=bool),
+ jnt_actfrcrange: wp.array2d(dtype=wp.vec2),
+ jnt_actgravcomp: wp.array(dtype=int),
+ jnt_axis: wp.array2d(dtype=wp.vec3),
+ jnt_bodyid: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ jnt_limited_ball_adr: wp.array(dtype=int),
+ jnt_limited_slide_hinge_adr: wp.array(dtype=int),
+ jnt_margin: wp.array2d(dtype=float),
+ jnt_pos: wp.array2d(dtype=wp.vec3),
+ jnt_qposadr: wp.array(dtype=int),
+ jnt_range: wp.array2d(dtype=wp.vec2),
+ jnt_solimp: wp.array2d(dtype=mjwp_types.vec5),
+ jnt_solref: wp.array2d(dtype=wp.vec2),
+ jnt_stiffness: wp.array2d(dtype=float),
+ jnt_type: wp.array(dtype=int),
+ light_bodyid: wp.array(dtype=int),
+ light_dir: wp.array2d(dtype=wp.vec3),
+ light_dir0: wp.array2d(dtype=wp.vec3),
+ light_mode: wp.array(dtype=int),
+ light_pos: wp.array2d(dtype=wp.vec3),
+ light_pos0: wp.array2d(dtype=wp.vec3),
+ light_poscom0: wp.array2d(dtype=wp.vec3),
+ light_targetbodyid: wp.array(dtype=int),
+ mapM2M: wp.array(dtype=int),
+ mat_rgba: wp.array2d(dtype=wp.vec4),
+ mesh_face: wp.array(dtype=wp.vec3i),
+ mesh_faceadr: wp.array(dtype=int),
+ mesh_graph: wp.array(dtype=int),
+ mesh_graphadr: wp.array(dtype=int),
+ mesh_polyadr: wp.array(dtype=int),
+ mesh_polymap: wp.array(dtype=int),
+ mesh_polymapadr: wp.array(dtype=int),
+ mesh_polymapnum: wp.array(dtype=int),
+ mesh_polynormal: wp.array(dtype=wp.vec3),
+ mesh_polynum: wp.array(dtype=int),
+ mesh_polyvert: wp.array(dtype=int),
+ mesh_polyvertadr: wp.array(dtype=int),
+ mesh_polyvertnum: wp.array(dtype=int),
+ mesh_vert: wp.array(dtype=wp.vec3),
+ mesh_vertadr: wp.array(dtype=int),
+ mesh_vertnum: wp.array(dtype=int),
+ mocap_bodyid: wp.array(dtype=int),
+ nC: int,
+ na: int,
+ nbody: int,
+ ncam: int,
+ neq: int,
+ nflexedge: int,
+ nflexelem: int,
+ nflexvert: int,
+ ngeom: int,
+ ngravcomp: int,
+ nhfield: int,
+ njnt: int,
+ nlight: int,
+ nlsp: int,
+ nmeshface: int,
+ nmocap: int,
+ nsite: int,
+ ntendon: int,
+ nu: int,
+ nv: int,
+ nxn_geom_pair_filtered: wp.array(dtype=wp.vec2i),
+ nxn_pairid: wp.array(dtype=int),
+ nxn_pairid_filtered: wp.array(dtype=int),
+ pair_dim: wp.array(dtype=int),
+ pair_friction: wp.array2d(dtype=mjwp_types.vec5),
+ pair_gap: wp.array2d(dtype=float),
+ pair_margin: wp.array2d(dtype=float),
+ pair_solimp: wp.array2d(dtype=mjwp_types.vec5),
+ pair_solref: wp.array2d(dtype=wp.vec2),
+ pair_solreffriction: wp.array2d(dtype=wp.vec2),
+ plugin: wp.array(dtype=int),
+ plugin_attr: wp.array(dtype=wp.vec3f),
+ qLD_updates: tuple[wp.array(dtype=wp.vec3i), ...],
+ qM_fullm_i: wp.array(dtype=int),
+ qM_fullm_j: wp.array(dtype=int),
+ qM_madr_ij: wp.array(dtype=int),
+ qM_mulm_i: wp.array(dtype=int),
+ qM_mulm_j: wp.array(dtype=int),
+ qM_tiles: tuple[mjwp_types.TileSet, ...],
+ qpos0: wp.array2d(dtype=float),
+ qpos_spring: wp.array2d(dtype=float),
+ rangefinder_sensor_adr: wp.array(dtype=int),
+ sensor_acc_adr: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_cutoff: wp.array(dtype=float),
+ sensor_datatype: wp.array(dtype=int),
+ sensor_e_kinetic: bool,
+ sensor_e_potential: bool,
+ sensor_limitfrc_adr: wp.array(dtype=int),
+ sensor_limitpos_adr: wp.array(dtype=int),
+ sensor_limitvel_adr: wp.array(dtype=int),
+ sensor_objid: wp.array(dtype=int),
+ sensor_objtype: wp.array(dtype=int),
+ sensor_pos_adr: wp.array(dtype=int),
+ sensor_rangefinder_adr: wp.array(dtype=int),
+ sensor_rangefinder_bodyid: wp.array(dtype=int),
+ sensor_refid: wp.array(dtype=int),
+ sensor_reftype: wp.array(dtype=int),
+ sensor_rne_postconstraint: bool,
+ sensor_subtree_vel: bool,
+ sensor_tendonactfrc_adr: wp.array(dtype=int),
+ sensor_touch_adr: wp.array(dtype=int),
+ sensor_type: wp.array(dtype=int),
+ sensor_vel_adr: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ site_pos: wp.array2d(dtype=wp.vec3),
+ site_quat: wp.array2d(dtype=wp.quat),
+ site_size: wp.array(dtype=wp.vec3),
+ site_type: wp.array(dtype=int),
+ subtree_mass: wp.array2d(dtype=float),
+ tendon_actfrclimited: wp.array(dtype=bool),
+ tendon_actfrcrange: wp.array2d(dtype=wp.vec2),
+ tendon_adr: wp.array(dtype=int),
+ tendon_armature: wp.array2d(dtype=float),
+ tendon_damping: wp.array2d(dtype=float),
+ tendon_frictionloss: wp.array2d(dtype=float),
+ tendon_geom_adr: wp.array(dtype=int),
+ tendon_invweight0: wp.array2d(dtype=float),
+ tendon_jnt_adr: wp.array(dtype=int),
+ tendon_length0: wp.array2d(dtype=float),
+ tendon_lengthspring: wp.array2d(dtype=wp.vec2),
+ tendon_limited_adr: wp.array(dtype=int),
+ tendon_margin: wp.array2d(dtype=float),
+ tendon_num: wp.array(dtype=int),
+ tendon_range: wp.array2d(dtype=wp.vec2),
+ tendon_site_pair_adr: wp.array(dtype=int),
+ tendon_solimp_fri: wp.array2d(dtype=mjwp_types.vec5),
+ tendon_solimp_lim: wp.array2d(dtype=mjwp_types.vec5),
+ tendon_solref_fri: wp.array2d(dtype=wp.vec2),
+ tendon_solref_lim: wp.array2d(dtype=wp.vec2),
+ tendon_stiffness: wp.array2d(dtype=float),
+ 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),
+ opt__broadphase: int,
+ opt__broadphase_filter: int,
+ opt__cone: int,
+ opt__density: wp.array(dtype=float),
+ opt__disableflags: int,
+ opt__enableflags: int,
+ opt__epa_iterations: int,
+ opt__gjk_iterations: int,
+ opt__graph_conditional: bool,
+ opt__gravity: wp.array(dtype=wp.vec3),
+ opt__has_fluid: bool,
+ opt__impratio: wp.array(dtype=float),
+ opt__is_sparse: bool,
+ opt__iterations: int,
+ opt__ls_iterations: int,
+ opt__ls_parallel: bool,
+ opt__ls_tolerance: wp.array(dtype=float),
+ opt__magnetic: wp.array(dtype=wp.vec3),
+ opt__run_collision_detection: bool,
+ opt__sdf_initpoints: int,
+ opt__sdf_iterations: int,
+ opt__solver: int,
+ opt__timestep: wp.array(dtype=float),
+ opt__tolerance: wp.array(dtype=float),
+ opt__viscosity: wp.array(dtype=float),
+ opt__wind: wp.array(dtype=wp.vec3),
+ stat__meaninertia: float,
+ # Data
+ nconmax: int,
+ njmax: int,
+ act: wp.array2d(dtype=float),
+ act_dot: wp.array2d(dtype=float),
+ actuator_force: wp.array2d(dtype=float),
+ actuator_length: wp.array2d(dtype=float),
+ actuator_moment: wp.array3d(dtype=float),
+ actuator_trntype_body_ncon: wp.array2d(dtype=int),
+ actuator_velocity: wp.array2d(dtype=float),
+ cacc: wp.array2d(dtype=wp.spatial_vector),
+ cam_xmat: wp.array2d(dtype=wp.mat33),
+ cam_xpos: wp.array2d(dtype=wp.vec3),
+ cdof: wp.array2d(dtype=wp.spatial_vector),
+ cdof_dot: wp.array2d(dtype=wp.spatial_vector),
+ cfrc_ext: wp.array2d(dtype=wp.spatial_vector),
+ cfrc_int: wp.array2d(dtype=wp.spatial_vector),
+ cinert: wp.array2d(dtype=mjwp_types.vec10),
+ collision_hftri_index: wp.array(dtype=int),
+ collision_pair: wp.array(dtype=wp.vec2i),
+ collision_pairid: wp.array(dtype=int),
+ collision_worldid: wp.array(dtype=int),
+ crb: wp.array2d(dtype=mjwp_types.vec10),
+ ctrl: wp.array2d(dtype=float),
+ cvel: wp.array2d(dtype=wp.spatial_vector),
+ energy: wp.array(dtype=wp.vec2),
+ epa_face: wp.array2d(dtype=wp.vec3i),
+ epa_horizon: wp.array2d(dtype=int),
+ epa_index: wp.array2d(dtype=int),
+ epa_map: wp.array2d(dtype=int),
+ epa_norm2: wp.array2d(dtype=float),
+ epa_pr: wp.array2d(dtype=wp.vec3),
+ epa_vert: wp.array2d(dtype=wp.vec3),
+ epa_vert1: wp.array2d(dtype=wp.vec3),
+ epa_vert2: wp.array2d(dtype=wp.vec3),
+ epa_vert_index1: wp.array2d(dtype=int),
+ epa_vert_index2: wp.array2d(dtype=int),
+ eq_active: wp.array2d(dtype=bool),
+ flexedge_length: wp.array2d(dtype=float),
+ flexedge_velocity: wp.array2d(dtype=float),
+ flexvert_xpos: wp.array2d(dtype=wp.vec3),
+ fluid_applied: wp.array2d(dtype=wp.spatial_vector),
+ geom_skip: wp.array(dtype=bool),
+ geom_xmat: wp.array2d(dtype=wp.mat33),
+ geom_xpos: wp.array2d(dtype=wp.vec3),
+ light_xdir: wp.array2d(dtype=wp.vec3),
+ light_xpos: wp.array2d(dtype=wp.vec3),
+ mocap_pos: wp.array2d(dtype=wp.vec3),
+ mocap_quat: wp.array2d(dtype=wp.quat),
+ ncollision: wp.array(dtype=int),
+ ncon: wp.array(dtype=int),
+ ncon_hfield: wp.array2d(dtype=int),
+ ne: wp.array(dtype=int),
+ ne_connect: wp.array(dtype=int),
+ ne_jnt: wp.array(dtype=int),
+ ne_ten: wp.array(dtype=int),
+ ne_weld: wp.array(dtype=int),
+ nefc: wp.array(dtype=int),
+ nf: wp.array(dtype=int),
+ nl: wp.array(dtype=int),
+ nsolving: wp.array(dtype=int),
+ qLD: wp.array3d(dtype=float),
+ qLDiagInv: wp.array2d(dtype=float),
+ qM: wp.array3d(dtype=float),
+ qacc: wp.array2d(dtype=float),
+ qacc_smooth: wp.array2d(dtype=float),
+ qacc_warmstart: wp.array2d(dtype=float),
+ qfrc_actuator: wp.array2d(dtype=float),
+ qfrc_applied: wp.array2d(dtype=float),
+ qfrc_bias: wp.array2d(dtype=float),
+ qfrc_constraint: wp.array2d(dtype=float),
+ qfrc_damper: wp.array2d(dtype=float),
+ qfrc_fluid: wp.array2d(dtype=float),
+ qfrc_gravcomp: wp.array2d(dtype=float),
+ qfrc_passive: wp.array2d(dtype=float),
+ qfrc_smooth: wp.array2d(dtype=float),
+ qfrc_spring: wp.array2d(dtype=float),
+ qpos: wp.array2d(dtype=float),
+ qvel: wp.array2d(dtype=float),
+ sap_cumulative_sum: wp.array2d(dtype=int),
+ sap_projection_lower: wp.array3d(dtype=float),
+ sap_projection_upper: wp.array2d(dtype=float),
+ sap_range: wp.array2d(dtype=int),
+ sap_segment_index: wp.array2d(dtype=int),
+ sap_sort_index: wp.array3d(dtype=int),
+ sensor_rangefinder_dist: wp.array2d(dtype=float),
+ sensor_rangefinder_geomid: wp.array2d(dtype=int),
+ sensor_rangefinder_pnt: wp.array2d(dtype=wp.vec3),
+ sensor_rangefinder_vec: wp.array2d(dtype=wp.vec3),
+ sensordata: wp.array2d(dtype=float),
+ site_xmat: wp.array2d(dtype=wp.mat33),
+ site_xpos: wp.array2d(dtype=wp.vec3),
+ solver_niter: wp.array(dtype=int),
+ subtree_angmom: wp.array2d(dtype=wp.vec3),
+ subtree_bodyvel: wp.array2d(dtype=wp.spatial_vector),
+ subtree_com: wp.array2d(dtype=wp.vec3),
+ subtree_linvel: wp.array2d(dtype=wp.vec3),
+ ten_J: wp.array3d(dtype=float),
+ ten_Jdot: wp.array3d(dtype=float),
+ ten_actfrc: wp.array2d(dtype=float),
+ ten_bias_coef: wp.array2d(dtype=float),
+ ten_length: wp.array2d(dtype=float),
+ ten_velocity: wp.array2d(dtype=float),
+ ten_wrapadr: wp.array2d(dtype=int),
+ ten_wrapnum: wp.array2d(dtype=int),
+ time: wp.array(dtype=float),
+ wrap_geom_xpos: wp.array2d(dtype=wp.spatial_vector),
+ wrap_obj: wp.array2d(dtype=wp.vec2i),
+ wrap_xpos: wp.array2d(dtype=wp.spatial_vector),
+ xanchor: wp.array2d(dtype=wp.vec3),
+ xaxis: wp.array2d(dtype=wp.vec3),
+ xfrc_applied: wp.array2d(dtype=wp.spatial_vector),
+ ximat: wp.array2d(dtype=wp.mat33),
+ xipos: wp.array2d(dtype=wp.vec3),
+ xmat: wp.array2d(dtype=wp.mat33),
+ xpos: wp.array2d(dtype=wp.vec3),
+ xquat: wp.array2d(dtype=wp.quat),
+ contact__dim: wp.array(dtype=int),
+ contact__dist: wp.array(dtype=float),
+ contact__efc_address: wp.array2d(dtype=int),
+ contact__frame: wp.array(dtype=wp.mat33),
+ contact__friction: wp.array(dtype=mjwp_types.vec5),
+ contact__geom: wp.array(dtype=wp.vec2i),
+ contact__includemargin: wp.array(dtype=float),
+ contact__pos: wp.array(dtype=wp.vec3),
+ contact__solimp: wp.array(dtype=mjwp_types.vec5),
+ contact__solref: wp.array(dtype=wp.vec2),
+ contact__solreffriction: wp.array(dtype=wp.vec2),
+ contact__worldid: wp.array(dtype=int),
+ efc__D: wp.array2d(dtype=float),
+ efc__J: wp.array3d(dtype=float),
+ efc__Jaref: wp.array2d(dtype=float),
+ efc__Ma: wp.array2d(dtype=float),
+ efc__Mgrad: wp.array2d(dtype=float),
+ efc__active: wp.array2d(dtype=bool),
+ efc__alpha: wp.array(dtype=float),
+ efc__aref: wp.array2d(dtype=float),
+ efc__beta: wp.array(dtype=float),
+ efc__beta_den: wp.array(dtype=float),
+ efc__beta_num: wp.array(dtype=float),
+ efc__cholesky_L_tmp: wp.array3d(dtype=float),
+ efc__cholesky_y_tmp: wp.array2d(dtype=float),
+ efc__condim: wp.array2d(dtype=int),
+ efc__cost: wp.array(dtype=float),
+ efc__cost_candidate: wp.array2d(dtype=float),
+ efc__done: wp.array(dtype=bool),
+ efc__force: wp.array2d(dtype=float),
+ efc__frictionloss: wp.array2d(dtype=float),
+ efc__gauss: wp.array(dtype=float),
+ efc__grad: wp.array2d(dtype=float),
+ efc__grad_dot: wp.array(dtype=float),
+ efc__gtol: wp.array(dtype=float),
+ efc__h: wp.array3d(dtype=float),
+ efc__hi: wp.array(dtype=wp.vec3),
+ efc__hi_alpha: wp.array(dtype=float),
+ efc__hi_next: wp.array(dtype=wp.vec3),
+ efc__hi_next_alpha: wp.array(dtype=float),
+ efc__id: wp.array2d(dtype=int),
+ efc__jv: wp.array2d(dtype=float),
+ efc__lo: wp.array(dtype=wp.vec3),
+ efc__lo_alpha: wp.array(dtype=float),
+ efc__lo_next: wp.array(dtype=wp.vec3),
+ efc__lo_next_alpha: wp.array(dtype=float),
+ efc__ls_done: wp.array(dtype=bool),
+ efc__margin: wp.array2d(dtype=float),
+ efc__mid: wp.array(dtype=wp.vec3),
+ efc__mid_alpha: wp.array(dtype=float),
+ efc__mv: wp.array2d(dtype=float),
+ efc__p0: wp.array(dtype=wp.vec3),
+ efc__pos: wp.array2d(dtype=float),
+ efc__prev_Mgrad: wp.array2d(dtype=float),
+ efc__prev_cost: wp.array(dtype=float),
+ efc__prev_grad: wp.array2d(dtype=float),
+ efc__quad: wp.array2d(dtype=wp.vec3),
+ efc__quad_gauss: wp.array(dtype=wp.vec3),
+ efc__search: wp.array2d(dtype=float),
+ efc__search_dot: wp.array(dtype=float),
+ efc__type: wp.array2d(dtype=int),
+ efc__u: wp.array(dtype=mjwp_types.vec6),
+ efc__uu: wp.array(dtype=float),
+ efc__uv: wp.array(dtype=float),
+ efc__vel: wp.array2d(dtype=float),
+ efc__vv: wp.array(dtype=float),
+):
+ _m.stat = _s
+ _m.opt = _o
+ _d.efc = _e
+ _d.contact = _c
+ _m.M_rowadr = M_rowadr
+ _m.M_rownnz = M_rownnz
+ _m.actuator_acc0 = actuator_acc0
+ _m.actuator_actadr = actuator_actadr
+ _m.actuator_actearly = actuator_actearly
+ _m.actuator_actlimited = actuator_actlimited
+ _m.actuator_actnum = actuator_actnum
+ _m.actuator_actrange = actuator_actrange
+ _m.actuator_biasprm = actuator_biasprm
+ _m.actuator_biastype = actuator_biastype
+ _m.actuator_cranklength = actuator_cranklength
+ _m.actuator_ctrllimited = actuator_ctrllimited
+ _m.actuator_ctrlrange = actuator_ctrlrange
+ _m.actuator_dynprm = actuator_dynprm
+ _m.actuator_dyntype = actuator_dyntype
+ _m.actuator_forcelimited = actuator_forcelimited
+ _m.actuator_forcerange = actuator_forcerange
+ _m.actuator_gainprm = actuator_gainprm
+ _m.actuator_gaintype = actuator_gaintype
+ _m.actuator_gear = actuator_gear
+ _m.actuator_lengthrange = actuator_lengthrange
+ _m.actuator_trnid = actuator_trnid
+ _m.actuator_trntype = actuator_trntype
+ _m.actuator_trntype_body_adr = actuator_trntype_body_adr
+ _m.block_dim = block_dim
+ _m.body_dofadr = body_dofadr
+ _m.body_dofnum = body_dofnum
+ _m.body_gravcomp = body_gravcomp
+ _m.body_inertia = body_inertia
+ _m.body_invweight0 = body_invweight0
+ _m.body_ipos = body_ipos
+ _m.body_iquat = body_iquat
+ _m.body_jntadr = body_jntadr
+ _m.body_jntnum = body_jntnum
+ _m.body_mass = body_mass
+ _m.body_parentid = body_parentid
+ _m.body_pos = body_pos
+ _m.body_quat = body_quat
+ _m.body_rootid = body_rootid
+ _m.body_subtreemass = body_subtreemass
+ _m.body_tree = body_tree
+ _m.body_weldid = body_weldid
+ _m.cam_bodyid = cam_bodyid
+ _m.cam_fovy = cam_fovy
+ _m.cam_intrinsic = cam_intrinsic
+ _m.cam_mat0 = cam_mat0
+ _m.cam_mode = cam_mode
+ _m.cam_pos = cam_pos
+ _m.cam_pos0 = cam_pos0
+ _m.cam_poscom0 = cam_poscom0
+ _m.cam_quat = cam_quat
+ _m.cam_resolution = cam_resolution
+ _m.cam_sensorsize = cam_sensorsize
+ _m.cam_targetbodyid = cam_targetbodyid
+ _m.condim_max = condim_max
+ _m.dof_Madr = dof_Madr
+ _m.dof_armature = dof_armature
+ _m.dof_bodyid = dof_bodyid
+ _m.dof_damping = dof_damping
+ _m.dof_frictionloss = dof_frictionloss
+ _m.dof_invweight0 = dof_invweight0
+ _m.dof_jntid = dof_jntid
+ _m.dof_parentid = dof_parentid
+ _m.dof_solimp = dof_solimp
+ _m.dof_solref = dof_solref
+ _m.dof_tri_col = dof_tri_col
+ _m.dof_tri_row = dof_tri_row
+ _m.eq_connect_adr = eq_connect_adr
+ _m.eq_data = eq_data
+ _m.eq_jnt_adr = eq_jnt_adr
+ _m.eq_obj1id = eq_obj1id
+ _m.eq_obj2id = eq_obj2id
+ _m.eq_objtype = eq_objtype
+ _m.eq_solimp = eq_solimp
+ _m.eq_solref = eq_solref
+ _m.eq_ten_adr = eq_ten_adr
+ _m.eq_wld_adr = eq_wld_adr
+ _m.flex_bending = flex_bending
+ _m.flex_damping = flex_damping
+ _m.flex_dim = flex_dim
+ _m.flex_edge = flex_edge
+ _m.flex_edgeadr = flex_edgeadr
+ _m.flex_edgeflap = flex_edgeflap
+ _m.flex_elem = flex_elem
+ _m.flex_elemedge = flex_elemedge
+ _m.flex_elemedgeadr = flex_elemedgeadr
+ _m.flex_stiffness = flex_stiffness
+ _m.flex_vertadr = flex_vertadr
+ _m.flex_vertbodyid = flex_vertbodyid
+ _m.flexedge_length0 = flexedge_length0
+ _m.geom_aabb = geom_aabb
+ _m.geom_bodyid = geom_bodyid
+ _m.geom_condim = geom_condim
+ _m.geom_dataid = geom_dataid
+ _m.geom_friction = geom_friction
+ _m.geom_gap = geom_gap
+ _m.geom_group = geom_group
+ _m.geom_margin = geom_margin
+ _m.geom_matid = geom_matid
+ _m.geom_pair_type_count = geom_pair_type_count
+ _m.geom_plugin_index = geom_plugin_index
+ _m.geom_pos = geom_pos
+ _m.geom_priority = geom_priority
+ _m.geom_quat = geom_quat
+ _m.geom_rbound = geom_rbound
+ _m.geom_rgba = geom_rgba
+ _m.geom_size = geom_size
+ _m.geom_solimp = geom_solimp
+ _m.geom_solmix = geom_solmix
+ _m.geom_solref = geom_solref
+ _m.geom_type = geom_type
+ _m.geompair2hfgeompair = geompair2hfgeompair
+ _m.has_sdf_geom = has_sdf_geom
+ _m.hfield_adr = hfield_adr
+ _m.hfield_data = hfield_data
+ _m.hfield_ncol = hfield_ncol
+ _m.hfield_nrow = hfield_nrow
+ _m.hfield_size = hfield_size
+ _m.jnt_actfrclimited = jnt_actfrclimited
+ _m.jnt_actfrcrange = jnt_actfrcrange
+ _m.jnt_actgravcomp = jnt_actgravcomp
+ _m.jnt_axis = jnt_axis
+ _m.jnt_bodyid = jnt_bodyid
+ _m.jnt_dofadr = jnt_dofadr
+ _m.jnt_limited_ball_adr = jnt_limited_ball_adr
+ _m.jnt_limited_slide_hinge_adr = jnt_limited_slide_hinge_adr
+ _m.jnt_margin = jnt_margin
+ _m.jnt_pos = jnt_pos
+ _m.jnt_qposadr = jnt_qposadr
+ _m.jnt_range = jnt_range
+ _m.jnt_solimp = jnt_solimp
+ _m.jnt_solref = jnt_solref
+ _m.jnt_stiffness = jnt_stiffness
+ _m.jnt_type = jnt_type
+ _m.light_bodyid = light_bodyid
+ _m.light_dir = light_dir
+ _m.light_dir0 = light_dir0
+ _m.light_mode = light_mode
+ _m.light_pos = light_pos
+ _m.light_pos0 = light_pos0
+ _m.light_poscom0 = light_poscom0
+ _m.light_targetbodyid = light_targetbodyid
+ _m.mapM2M = mapM2M
+ _m.mat_rgba = mat_rgba
+ _m.mesh_face = mesh_face
+ _m.mesh_faceadr = mesh_faceadr
+ _m.mesh_graph = mesh_graph
+ _m.mesh_graphadr = mesh_graphadr
+ _m.mesh_polyadr = mesh_polyadr
+ _m.mesh_polymap = mesh_polymap
+ _m.mesh_polymapadr = mesh_polymapadr
+ _m.mesh_polymapnum = mesh_polymapnum
+ _m.mesh_polynormal = mesh_polynormal
+ _m.mesh_polynum = mesh_polynum
+ _m.mesh_polyvert = mesh_polyvert
+ _m.mesh_polyvertadr = mesh_polyvertadr
+ _m.mesh_polyvertnum = mesh_polyvertnum
+ _m.mesh_vert = mesh_vert
+ _m.mesh_vertadr = mesh_vertadr
+ _m.mesh_vertnum = mesh_vertnum
+ _m.mocap_bodyid = mocap_bodyid
+ _m.nC = nC
+ _m.na = na
+ _m.nbody = nbody
+ _m.ncam = ncam
+ _m.neq = neq
+ _m.nflexedge = nflexedge
+ _m.nflexelem = nflexelem
+ _m.nflexvert = nflexvert
+ _m.ngeom = ngeom
+ _m.ngravcomp = ngravcomp
+ _m.nhfield = nhfield
+ _m.njnt = njnt
+ _m.nlight = nlight
+ _m.nlsp = nlsp
+ _m.nmeshface = nmeshface
+ _m.nmocap = nmocap
+ _m.nsite = nsite
+ _m.ntendon = ntendon
+ _m.nu = nu
+ _m.nv = nv
+ _m.nxn_geom_pair_filtered = nxn_geom_pair_filtered
+ _m.nxn_pairid = nxn_pairid
+ _m.nxn_pairid_filtered = nxn_pairid_filtered
+ _m.opt.broadphase = opt__broadphase
+ _m.opt.broadphase_filter = opt__broadphase_filter
+ _m.opt.cone = opt__cone
+ _m.opt.density = opt__density
+ _m.opt.disableflags = opt__disableflags
+ _m.opt.enableflags = opt__enableflags
+ _m.opt.epa_iterations = opt__epa_iterations
+ _m.opt.gjk_iterations = opt__gjk_iterations
+ _m.opt.graph_conditional = opt__graph_conditional
+ _m.opt.gravity = opt__gravity
+ _m.opt.has_fluid = opt__has_fluid
+ _m.opt.impratio = opt__impratio
+ _m.opt.is_sparse = opt__is_sparse
+ _m.opt.iterations = opt__iterations
+ _m.opt.ls_iterations = opt__ls_iterations
+ _m.opt.ls_parallel = opt__ls_parallel
+ _m.opt.ls_tolerance = opt__ls_tolerance
+ _m.opt.magnetic = opt__magnetic
+ _m.opt.run_collision_detection = opt__run_collision_detection
+ _m.opt.sdf_initpoints = opt__sdf_initpoints
+ _m.opt.sdf_iterations = opt__sdf_iterations
+ _m.opt.solver = opt__solver
+ _m.opt.timestep = opt__timestep
+ _m.opt.tolerance = opt__tolerance
+ _m.opt.viscosity = opt__viscosity
+ _m.opt.wind = opt__wind
+ _m.pair_dim = pair_dim
+ _m.pair_friction = pair_friction
+ _m.pair_gap = pair_gap
+ _m.pair_margin = pair_margin
+ _m.pair_solimp = pair_solimp
+ _m.pair_solref = pair_solref
+ _m.pair_solreffriction = pair_solreffriction
+ _m.plugin = plugin
+ _m.plugin_attr = plugin_attr
+ _m.qLD_updates = qLD_updates
+ _m.qM_fullm_i = qM_fullm_i
+ _m.qM_fullm_j = qM_fullm_j
+ _m.qM_madr_ij = qM_madr_ij
+ _m.qM_mulm_i = qM_mulm_i
+ _m.qM_mulm_j = qM_mulm_j
+ _m.qM_tiles = qM_tiles
+ _m.qpos0 = qpos0
+ _m.qpos_spring = qpos_spring
+ _m.rangefinder_sensor_adr = rangefinder_sensor_adr
+ _m.sensor_acc_adr = sensor_acc_adr
+ _m.sensor_adr = sensor_adr
+ _m.sensor_cutoff = sensor_cutoff
+ _m.sensor_datatype = sensor_datatype
+ _m.sensor_e_kinetic = sensor_e_kinetic
+ _m.sensor_e_potential = sensor_e_potential
+ _m.sensor_limitfrc_adr = sensor_limitfrc_adr
+ _m.sensor_limitpos_adr = sensor_limitpos_adr
+ _m.sensor_limitvel_adr = sensor_limitvel_adr
+ _m.sensor_objid = sensor_objid
+ _m.sensor_objtype = sensor_objtype
+ _m.sensor_pos_adr = sensor_pos_adr
+ _m.sensor_rangefinder_adr = sensor_rangefinder_adr
+ _m.sensor_rangefinder_bodyid = sensor_rangefinder_bodyid
+ _m.sensor_refid = sensor_refid
+ _m.sensor_reftype = sensor_reftype
+ _m.sensor_rne_postconstraint = sensor_rne_postconstraint
+ _m.sensor_subtree_vel = sensor_subtree_vel
+ _m.sensor_tendonactfrc_adr = sensor_tendonactfrc_adr
+ _m.sensor_touch_adr = sensor_touch_adr
+ _m.sensor_type = sensor_type
+ _m.sensor_vel_adr = sensor_vel_adr
+ _m.site_bodyid = site_bodyid
+ _m.site_pos = site_pos
+ _m.site_quat = site_quat
+ _m.site_size = site_size
+ _m.site_type = site_type
+ _m.stat.meaninertia = stat__meaninertia
+ _m.subtree_mass = subtree_mass
+ _m.tendon_actfrclimited = tendon_actfrclimited
+ _m.tendon_actfrcrange = tendon_actfrcrange
+ _m.tendon_adr = tendon_adr
+ _m.tendon_armature = tendon_armature
+ _m.tendon_damping = tendon_damping
+ _m.tendon_frictionloss = tendon_frictionloss
+ _m.tendon_geom_adr = tendon_geom_adr
+ _m.tendon_invweight0 = tendon_invweight0
+ _m.tendon_jnt_adr = tendon_jnt_adr
+ _m.tendon_length0 = tendon_length0
+ _m.tendon_lengthspring = tendon_lengthspring
+ _m.tendon_limited_adr = tendon_limited_adr
+ _m.tendon_margin = tendon_margin
+ _m.tendon_num = tendon_num
+ _m.tendon_range = tendon_range
+ _m.tendon_site_pair_adr = tendon_site_pair_adr
+ _m.tendon_solimp_fri = tendon_solimp_fri
+ _m.tendon_solimp_lim = tendon_solimp_lim
+ _m.tendon_solref_fri = tendon_solref_fri
+ _m.tendon_solref_lim = tendon_solref_lim
+ _m.tendon_stiffness = tendon_stiffness
+ _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.act = act
+ _d.act_dot = act_dot
+ _d.actuator_force = actuator_force
+ _d.actuator_length = actuator_length
+ _d.actuator_moment = actuator_moment
+ _d.actuator_trntype_body_ncon = actuator_trntype_body_ncon
+ _d.actuator_velocity = actuator_velocity
+ _d.cacc = cacc
+ _d.cam_xmat = cam_xmat
+ _d.cam_xpos = cam_xpos
+ _d.cdof = cdof
+ _d.cdof_dot = cdof_dot
+ _d.cfrc_ext = cfrc_ext
+ _d.cfrc_int = cfrc_int
+ _d.cinert = cinert
+ _d.collision_hftri_index = collision_hftri_index
+ _d.collision_pair = collision_pair
+ _d.collision_pairid = collision_pairid
+ _d.collision_worldid = collision_worldid
+ _d.contact.dim = contact__dim
+ _d.contact.dist = contact__dist
+ _d.contact.efc_address = contact__efc_address
+ _d.contact.frame = contact__frame
+ _d.contact.friction = contact__friction
+ _d.contact.geom = contact__geom
+ _d.contact.includemargin = contact__includemargin
+ _d.contact.pos = contact__pos
+ _d.contact.solimp = contact__solimp
+ _d.contact.solref = contact__solref
+ _d.contact.solreffriction = contact__solreffriction
+ _d.contact.worldid = contact__worldid
+ _d.crb = crb
+ _d.ctrl = ctrl
+ _d.cvel = cvel
+ _d.efc.D = efc__D
+ _d.efc.J = efc__J
+ _d.efc.Jaref = efc__Jaref
+ _d.efc.Ma = efc__Ma
+ _d.efc.Mgrad = efc__Mgrad
+ _d.efc.active = efc__active
+ _d.efc.alpha = efc__alpha
+ _d.efc.aref = efc__aref
+ _d.efc.beta = efc__beta
+ _d.efc.beta_den = efc__beta_den
+ _d.efc.beta_num = efc__beta_num
+ _d.efc.cholesky_L_tmp = efc__cholesky_L_tmp
+ _d.efc.cholesky_y_tmp = efc__cholesky_y_tmp
+ _d.efc.condim = efc__condim
+ _d.efc.cost = efc__cost
+ _d.efc.cost_candidate = efc__cost_candidate
+ _d.efc.done = efc__done
+ _d.efc.force = efc__force
+ _d.efc.frictionloss = efc__frictionloss
+ _d.efc.gauss = efc__gauss
+ _d.efc.grad = efc__grad
+ _d.efc.grad_dot = efc__grad_dot
+ _d.efc.gtol = efc__gtol
+ _d.efc.h = efc__h
+ _d.efc.hi = efc__hi
+ _d.efc.hi_alpha = efc__hi_alpha
+ _d.efc.hi_next = efc__hi_next
+ _d.efc.hi_next_alpha = efc__hi_next_alpha
+ _d.efc.id = efc__id
+ _d.efc.jv = efc__jv
+ _d.efc.lo = efc__lo
+ _d.efc.lo_alpha = efc__lo_alpha
+ _d.efc.lo_next = efc__lo_next
+ _d.efc.lo_next_alpha = efc__lo_next_alpha
+ _d.efc.ls_done = efc__ls_done
+ _d.efc.margin = efc__margin
+ _d.efc.mid = efc__mid
+ _d.efc.mid_alpha = efc__mid_alpha
+ _d.efc.mv = efc__mv
+ _d.efc.p0 = efc__p0
+ _d.efc.pos = efc__pos
+ _d.efc.prev_Mgrad = efc__prev_Mgrad
+ _d.efc.prev_cost = efc__prev_cost
+ _d.efc.prev_grad = efc__prev_grad
+ _d.efc.quad = efc__quad
+ _d.efc.quad_gauss = efc__quad_gauss
+ _d.efc.search = efc__search
+ _d.efc.search_dot = efc__search_dot
+ _d.efc.type = efc__type
+ _d.efc.u = efc__u
+ _d.efc.uu = efc__uu
+ _d.efc.uv = efc__uv
+ _d.efc.vel = efc__vel
+ _d.efc.vv = efc__vv
+ _d.energy = energy
+ _d.epa_face = epa_face
+ _d.epa_horizon = epa_horizon
+ _d.epa_index = epa_index
+ _d.epa_map = epa_map
+ _d.epa_norm2 = epa_norm2
+ _d.epa_pr = epa_pr
+ _d.epa_vert = epa_vert
+ _d.epa_vert1 = epa_vert1
+ _d.epa_vert2 = epa_vert2
+ _d.epa_vert_index1 = epa_vert_index1
+ _d.epa_vert_index2 = epa_vert_index2
+ _d.eq_active = eq_active
+ _d.flexedge_length = flexedge_length
+ _d.flexedge_velocity = flexedge_velocity
+ _d.flexvert_xpos = flexvert_xpos
+ _d.fluid_applied = fluid_applied
+ _d.geom_skip = geom_skip
+ _d.geom_xmat = geom_xmat
+ _d.geom_xpos = geom_xpos
+ _d.light_xdir = light_xdir
+ _d.light_xpos = light_xpos
+ _d.mocap_pos = mocap_pos
+ _d.mocap_quat = mocap_quat
+ _d.ncollision = ncollision
+ _d.ncon = ncon
+ _d.ncon_hfield = ncon_hfield
+ _d.nconmax = nconmax
+ _d.ne = ne
+ _d.ne_connect = ne_connect
+ _d.ne_jnt = ne_jnt
+ _d.ne_ten = ne_ten
+ _d.ne_weld = ne_weld
+ _d.nefc = nefc
+ _d.nf = nf
+ _d.njmax = njmax
+ _d.nl = nl
+ _d.nsolving = nsolving
+ _d.qLD = qLD
+ _d.qLDiagInv = qLDiagInv
+ _d.qM = qM
+ _d.qacc = qacc
+ _d.qacc_smooth = qacc_smooth
+ _d.qacc_warmstart = qacc_warmstart
+ _d.qfrc_actuator = qfrc_actuator
+ _d.qfrc_applied = qfrc_applied
+ _d.qfrc_bias = qfrc_bias
+ _d.qfrc_constraint = qfrc_constraint
+ _d.qfrc_damper = qfrc_damper
+ _d.qfrc_fluid = qfrc_fluid
+ _d.qfrc_gravcomp = qfrc_gravcomp
+ _d.qfrc_passive = qfrc_passive
+ _d.qfrc_smooth = qfrc_smooth
+ _d.qfrc_spring = qfrc_spring
+ _d.qpos = qpos
+ _d.qvel = qvel
+ _d.sap_cumulative_sum = sap_cumulative_sum
+ _d.sap_projection_lower = sap_projection_lower
+ _d.sap_projection_upper = sap_projection_upper
+ _d.sap_range = sap_range
+ _d.sap_segment_index = sap_segment_index
+ _d.sap_sort_index = sap_sort_index
+ _d.sensor_rangefinder_dist = sensor_rangefinder_dist
+ _d.sensor_rangefinder_geomid = sensor_rangefinder_geomid
+ _d.sensor_rangefinder_pnt = sensor_rangefinder_pnt
+ _d.sensor_rangefinder_vec = sensor_rangefinder_vec
+ _d.sensordata = sensordata
+ _d.site_xmat = site_xmat
+ _d.site_xpos = site_xpos
+ _d.solver_niter = solver_niter
+ _d.subtree_angmom = subtree_angmom
+ _d.subtree_bodyvel = subtree_bodyvel
+ _d.subtree_com = subtree_com
+ _d.subtree_linvel = subtree_linvel
+ _d.ten_J = ten_J
+ _d.ten_Jdot = ten_Jdot
+ _d.ten_actfrc = ten_actfrc
+ _d.ten_bias_coef = ten_bias_coef
+ _d.ten_length = ten_length
+ _d.ten_velocity = ten_velocity
+ _d.ten_wrapadr = ten_wrapadr
+ _d.ten_wrapnum = ten_wrapnum
+ _d.time = time
+ _d.wrap_geom_xpos = wrap_geom_xpos
+ _d.wrap_obj = wrap_obj
+ _d.wrap_xpos = wrap_xpos
+ _d.xanchor = xanchor
+ _d.xaxis = xaxis
+ _d.xfrc_applied = xfrc_applied
+ _d.ximat = ximat
+ _d.xipos = xipos
+ _d.xmat = xmat
+ _d.xpos = xpos
+ _d.xquat = xquat
+ _d.nworld = nworld
+ mjwarp.forward(_m, _d)
+
+
+def _forward_jax_impl(m: types.Model, d: types.Data):
+ output_dims = {
+ 'act': d.act.shape,
+ 'act_dot': d.act_dot.shape,
+ 'actuator_force': d.actuator_force.shape,
+ 'actuator_length': d._impl.actuator_length.shape,
+ 'actuator_moment': d._impl.actuator_moment.shape,
+ 'actuator_trntype_body_ncon': d._impl.actuator_trntype_body_ncon.shape,
+ 'actuator_velocity': d._impl.actuator_velocity.shape,
+ 'cacc': d._impl.cacc.shape,
+ 'cam_xmat': d.cam_xmat.shape,
+ 'cam_xpos': d.cam_xpos.shape,
+ 'cdof': d._impl.cdof.shape,
+ 'cdof_dot': d._impl.cdof_dot.shape,
+ 'cfrc_ext': d._impl.cfrc_ext.shape,
+ 'cfrc_int': d._impl.cfrc_int.shape,
+ 'cinert': d._impl.cinert.shape,
+ 'collision_hftri_index': d._impl.collision_hftri_index.shape,
+ 'collision_pair': d._impl.collision_pair.shape,
+ 'collision_pairid': d._impl.collision_pairid.shape,
+ 'collision_worldid': d._impl.collision_worldid.shape,
+ 'crb': d._impl.crb.shape,
+ 'ctrl': d.ctrl.shape,
+ 'cvel': d.cvel.shape,
+ 'energy': d._impl.energy.shape,
+ 'epa_face': d._impl.epa_face.shape,
+ 'epa_horizon': d._impl.epa_horizon.shape,
+ 'epa_index': d._impl.epa_index.shape,
+ 'epa_map': d._impl.epa_map.shape,
+ 'epa_norm2': d._impl.epa_norm2.shape,
+ 'epa_pr': d._impl.epa_pr.shape,
+ 'epa_vert': d._impl.epa_vert.shape,
+ 'epa_vert1': d._impl.epa_vert1.shape,
+ 'epa_vert2': d._impl.epa_vert2.shape,
+ 'epa_vert_index1': d._impl.epa_vert_index1.shape,
+ 'epa_vert_index2': d._impl.epa_vert_index2.shape,
+ 'eq_active': d.eq_active.shape,
+ 'flexedge_length': d._impl.flexedge_length.shape,
+ 'flexedge_velocity': d._impl.flexedge_velocity.shape,
+ 'flexvert_xpos': d._impl.flexvert_xpos.shape,
+ 'fluid_applied': d._impl.fluid_applied.shape,
+ 'geom_skip': d._impl.geom_skip.shape,
+ 'geom_xmat': d.geom_xmat.shape,
+ 'geom_xpos': d.geom_xpos.shape,
+ 'light_xdir': d._impl.light_xdir.shape,
+ 'light_xpos': d._impl.light_xpos.shape,
+ 'mocap_pos': d.mocap_pos.shape,
+ 'mocap_quat': d.mocap_quat.shape,
+ 'ncollision': d._impl.ncollision.shape,
+ 'ncon': d._impl.ncon.shape,
+ 'ncon_hfield': d._impl.ncon_hfield.shape,
+ 'ne': d._impl.ne.shape,
+ 'ne_connect': d._impl.ne_connect.shape,
+ 'ne_jnt': d._impl.ne_jnt.shape,
+ 'ne_ten': d._impl.ne_ten.shape,
+ 'ne_weld': d._impl.ne_weld.shape,
+ 'nefc': d._impl.nefc.shape,
+ 'nf': d._impl.nf.shape,
+ 'nl': d._impl.nl.shape,
+ 'nsolving': d._impl.nsolving.shape,
+ 'qLD': d._impl.qLD.shape,
+ 'qLDiagInv': d._impl.qLDiagInv.shape,
+ 'qM': d._impl.qM.shape,
+ 'qacc': d.qacc.shape,
+ 'qacc_smooth': d.qacc_smooth.shape,
+ 'qacc_warmstart': d.qacc_warmstart.shape,
+ 'qfrc_actuator': d.qfrc_actuator.shape,
+ 'qfrc_applied': d.qfrc_applied.shape,
+ 'qfrc_bias': d.qfrc_bias.shape,
+ 'qfrc_constraint': d.qfrc_constraint.shape,
+ 'qfrc_damper': d._impl.qfrc_damper.shape,
+ 'qfrc_fluid': d.qfrc_fluid.shape,
+ 'qfrc_gravcomp': d.qfrc_gravcomp.shape,
+ 'qfrc_passive': d.qfrc_passive.shape,
+ 'qfrc_smooth': d.qfrc_smooth.shape,
+ 'qfrc_spring': d._impl.qfrc_spring.shape,
+ 'qpos': d.qpos.shape,
+ 'qvel': d.qvel.shape,
+ 'sap_cumulative_sum': d._impl.sap_cumulative_sum.shape,
+ 'sap_projection_lower': d._impl.sap_projection_lower.shape,
+ 'sap_projection_upper': d._impl.sap_projection_upper.shape,
+ 'sap_range': d._impl.sap_range.shape,
+ 'sap_segment_index': d._impl.sap_segment_index.shape,
+ 'sap_sort_index': d._impl.sap_sort_index.shape,
+ 'sensor_rangefinder_dist': d._impl.sensor_rangefinder_dist.shape,
+ 'sensor_rangefinder_geomid': d._impl.sensor_rangefinder_geomid.shape,
+ 'sensor_rangefinder_pnt': d._impl.sensor_rangefinder_pnt.shape,
+ 'sensor_rangefinder_vec': d._impl.sensor_rangefinder_vec.shape,
+ 'sensordata': d.sensordata.shape,
+ 'site_xmat': d.site_xmat.shape,
+ 'site_xpos': d.site_xpos.shape,
+ 'solver_niter': d._impl.solver_niter.shape,
+ 'subtree_angmom': d._impl.subtree_angmom.shape,
+ 'subtree_bodyvel': d._impl.subtree_bodyvel.shape,
+ 'subtree_com': d.subtree_com.shape,
+ 'subtree_linvel': d._impl.subtree_linvel.shape,
+ 'ten_J': d._impl.ten_J.shape,
+ '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._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,
+ 'time': d.time.shape,
+ 'wrap_geom_xpos': d._impl.wrap_geom_xpos.shape,
+ 'wrap_obj': d._impl.wrap_obj.shape,
+ 'wrap_xpos': d._impl.wrap_xpos.shape,
+ 'xanchor': d.xanchor.shape,
+ 'xaxis': d.xaxis.shape,
+ 'xfrc_applied': d.xfrc_applied.shape,
+ 'ximat': d.ximat.shape,
+ 'xipos': d.xipos.shape,
+ 'xmat': d.xmat.shape,
+ 'xpos': d.xpos.shape,
+ 'xquat': d.xquat.shape,
+ 'contact__dim': d._impl.contact__dim.shape,
+ 'contact__dist': d._impl.contact__dist.shape,
+ 'contact__efc_address': d._impl.contact__efc_address.shape,
+ 'contact__frame': d._impl.contact__frame.shape,
+ 'contact__friction': d._impl.contact__friction.shape,
+ 'contact__geom': d._impl.contact__geom.shape,
+ 'contact__includemargin': d._impl.contact__includemargin.shape,
+ 'contact__pos': d._impl.contact__pos.shape,
+ 'contact__solimp': d._impl.contact__solimp.shape,
+ 'contact__solref': d._impl.contact__solref.shape,
+ 'contact__solreffriction': d._impl.contact__solreffriction.shape,
+ 'contact__worldid': d._impl.contact__worldid.shape,
+ 'efc__D': d._impl.efc__D.shape,
+ 'efc__J': d._impl.efc__J.shape,
+ 'efc__Jaref': d._impl.efc__Jaref.shape,
+ 'efc__Ma': d._impl.efc__Ma.shape,
+ 'efc__Mgrad': d._impl.efc__Mgrad.shape,
+ 'efc__active': d._impl.efc__active.shape,
+ 'efc__alpha': d._impl.efc__alpha.shape,
+ 'efc__aref': d._impl.efc__aref.shape,
+ 'efc__beta': d._impl.efc__beta.shape,
+ 'efc__beta_den': d._impl.efc__beta_den.shape,
+ 'efc__beta_num': d._impl.efc__beta_num.shape,
+ 'efc__cholesky_L_tmp': d._impl.efc__cholesky_L_tmp.shape,
+ 'efc__cholesky_y_tmp': d._impl.efc__cholesky_y_tmp.shape,
+ 'efc__condim': d._impl.efc__condim.shape,
+ 'efc__cost': d._impl.efc__cost.shape,
+ 'efc__cost_candidate': d._impl.efc__cost_candidate.shape,
+ 'efc__done': d._impl.efc__done.shape,
+ 'efc__force': d._impl.efc__force.shape,
+ 'efc__frictionloss': d._impl.efc__frictionloss.shape,
+ 'efc__gauss': d._impl.efc__gauss.shape,
+ 'efc__grad': d._impl.efc__grad.shape,
+ 'efc__grad_dot': d._impl.efc__grad_dot.shape,
+ 'efc__gtol': d._impl.efc__gtol.shape,
+ 'efc__h': d._impl.efc__h.shape,
+ 'efc__hi': d._impl.efc__hi.shape,
+ 'efc__hi_alpha': d._impl.efc__hi_alpha.shape,
+ 'efc__hi_next': d._impl.efc__hi_next.shape,
+ 'efc__hi_next_alpha': d._impl.efc__hi_next_alpha.shape,
+ 'efc__id': d._impl.efc__id.shape,
+ 'efc__jv': d._impl.efc__jv.shape,
+ 'efc__lo': d._impl.efc__lo.shape,
+ 'efc__lo_alpha': d._impl.efc__lo_alpha.shape,
+ 'efc__lo_next': d._impl.efc__lo_next.shape,
+ 'efc__lo_next_alpha': d._impl.efc__lo_next_alpha.shape,
+ 'efc__ls_done': d._impl.efc__ls_done.shape,
+ 'efc__margin': d._impl.efc__margin.shape,
+ 'efc__mid': d._impl.efc__mid.shape,
+ 'efc__mid_alpha': d._impl.efc__mid_alpha.shape,
+ 'efc__mv': d._impl.efc__mv.shape,
+ 'efc__p0': d._impl.efc__p0.shape,
+ 'efc__pos': d._impl.efc__pos.shape,
+ 'efc__prev_Mgrad': d._impl.efc__prev_Mgrad.shape,
+ 'efc__prev_cost': d._impl.efc__prev_cost.shape,
+ 'efc__prev_grad': d._impl.efc__prev_grad.shape,
+ 'efc__quad': d._impl.efc__quad.shape,
+ 'efc__quad_gauss': d._impl.efc__quad_gauss.shape,
+ 'efc__search': d._impl.efc__search.shape,
+ 'efc__search_dot': d._impl.efc__search_dot.shape,
+ 'efc__type': d._impl.efc__type.shape,
+ 'efc__u': d._impl.efc__u.shape,
+ 'efc__uu': d._impl.efc__uu.shape,
+ 'efc__uv': d._impl.efc__uv.shape,
+ 'efc__vel': d._impl.efc__vel.shape,
+ 'efc__vv': d._impl.efc__vv.shape,
+ }
+ jf = ffi.jax_callable_variadic_tuple(
+ _forward_shim,
+ num_outputs=180,
+ output_dims=output_dims,
+ vmap_method=None,
+ graph_compatible=True,
+ in_out_argnames={
+ 'act',
+ 'act_dot',
+ 'actuator_force',
+ 'actuator_length',
+ 'actuator_moment',
+ 'actuator_trntype_body_ncon',
+ 'actuator_velocity',
+ 'cacc',
+ 'cam_xmat',
+ 'cam_xpos',
+ 'cdof',
+ 'cdof_dot',
+ 'cfrc_ext',
+ 'cfrc_int',
+ 'cinert',
+ 'collision_hftri_index',
+ 'collision_pair',
+ 'collision_pairid',
+ 'collision_worldid',
+ 'crb',
+ 'ctrl',
+ 'cvel',
+ 'energy',
+ 'epa_face',
+ 'epa_horizon',
+ 'epa_index',
+ 'epa_map',
+ 'epa_norm2',
+ 'epa_pr',
+ 'epa_vert',
+ 'epa_vert1',
+ 'epa_vert2',
+ 'epa_vert_index1',
+ 'epa_vert_index2',
+ 'eq_active',
+ 'flexedge_length',
+ 'flexedge_velocity',
+ 'flexvert_xpos',
+ 'fluid_applied',
+ 'geom_skip',
+ 'geom_xmat',
+ 'geom_xpos',
+ 'light_xdir',
+ 'light_xpos',
+ 'mocap_pos',
+ 'mocap_quat',
+ 'ncollision',
+ 'ncon',
+ 'ncon_hfield',
+ 'ne',
+ 'ne_connect',
+ 'ne_jnt',
+ 'ne_ten',
+ 'ne_weld',
+ 'nefc',
+ 'nf',
+ 'nl',
+ 'nsolving',
+ 'qLD',
+ 'qLDiagInv',
+ 'qM',
+ 'qacc',
+ 'qacc_smooth',
+ 'qacc_warmstart',
+ 'qfrc_actuator',
+ 'qfrc_applied',
+ 'qfrc_bias',
+ 'qfrc_constraint',
+ 'qfrc_damper',
+ 'qfrc_fluid',
+ 'qfrc_gravcomp',
+ 'qfrc_passive',
+ 'qfrc_smooth',
+ 'qfrc_spring',
+ 'qpos',
+ 'qvel',
+ 'sap_cumulative_sum',
+ 'sap_projection_lower',
+ 'sap_projection_upper',
+ 'sap_range',
+ 'sap_segment_index',
+ 'sap_sort_index',
+ 'sensor_rangefinder_dist',
+ 'sensor_rangefinder_geomid',
+ 'sensor_rangefinder_pnt',
+ 'sensor_rangefinder_vec',
+ 'sensordata',
+ 'site_xmat',
+ 'site_xpos',
+ 'solver_niter',
+ 'subtree_angmom',
+ 'subtree_bodyvel',
+ 'subtree_com',
+ 'subtree_linvel',
+ 'ten_J',
+ 'ten_Jdot',
+ 'ten_actfrc',
+ 'ten_bias_coef',
+ 'ten_length',
+ 'ten_velocity',
+ 'ten_wrapadr',
+ 'ten_wrapnum',
+ 'time',
+ 'wrap_geom_xpos',
+ 'wrap_obj',
+ 'wrap_xpos',
+ 'xanchor',
+ 'xaxis',
+ 'xfrc_applied',
+ 'ximat',
+ 'xipos',
+ 'xmat',
+ 'xpos',
+ 'xquat',
+ 'contact__dim',
+ 'contact__dist',
+ 'contact__efc_address',
+ 'contact__frame',
+ 'contact__friction',
+ 'contact__geom',
+ 'contact__includemargin',
+ 'contact__pos',
+ 'contact__solimp',
+ 'contact__solref',
+ 'contact__solreffriction',
+ 'contact__worldid',
+ 'efc__D',
+ 'efc__J',
+ 'efc__Jaref',
+ 'efc__Ma',
+ 'efc__Mgrad',
+ 'efc__active',
+ 'efc__alpha',
+ 'efc__aref',
+ 'efc__beta',
+ 'efc__beta_den',
+ 'efc__beta_num',
+ 'efc__cholesky_L_tmp',
+ 'efc__cholesky_y_tmp',
+ 'efc__condim',
+ 'efc__cost',
+ 'efc__cost_candidate',
+ 'efc__done',
+ 'efc__force',
+ 'efc__frictionloss',
+ 'efc__gauss',
+ 'efc__grad',
+ 'efc__grad_dot',
+ 'efc__gtol',
+ 'efc__h',
+ 'efc__hi',
+ 'efc__hi_alpha',
+ 'efc__hi_next',
+ 'efc__hi_next_alpha',
+ 'efc__id',
+ 'efc__jv',
+ 'efc__lo',
+ 'efc__lo_alpha',
+ 'efc__lo_next',
+ 'efc__lo_next_alpha',
+ 'efc__ls_done',
+ 'efc__margin',
+ 'efc__mid',
+ 'efc__mid_alpha',
+ 'efc__mv',
+ 'efc__p0',
+ 'efc__pos',
+ 'efc__prev_Mgrad',
+ 'efc__prev_cost',
+ 'efc__prev_grad',
+ 'efc__quad',
+ 'efc__quad_gauss',
+ 'efc__search',
+ 'efc__search_dot',
+ 'efc__type',
+ 'efc__u',
+ 'efc__uu',
+ 'efc__uv',
+ 'efc__vel',
+ 'efc__vv',
+ },
+ )
+ out = jf(
+ d.qpos.shape[0],
+ m._impl.M_rowadr,
+ m._impl.M_rownnz,
+ m.actuator_acc0,
+ m.actuator_actadr,
+ m.actuator_actearly,
+ m.actuator_actlimited,
+ m.actuator_actnum,
+ m.actuator_actrange,
+ m.actuator_biasprm,
+ m.actuator_biastype,
+ m.actuator_cranklength,
+ m.actuator_ctrllimited,
+ m.actuator_ctrlrange,
+ m.actuator_dynprm,
+ m.actuator_dyntype,
+ m.actuator_forcelimited,
+ m.actuator_forcerange,
+ m.actuator_gainprm,
+ m.actuator_gaintype,
+ m.actuator_gear,
+ m.actuator_lengthrange,
+ m.actuator_trnid,
+ m.actuator_trntype,
+ m._impl.actuator_trntype_body_adr,
+ m._impl.block_dim,
+ m.body_dofadr,
+ m.body_dofnum,
+ m.body_gravcomp,
+ m.body_inertia,
+ m.body_invweight0,
+ m.body_ipos,
+ m.body_iquat,
+ m.body_jntadr,
+ m.body_jntnum,
+ m.body_mass,
+ m.body_parentid,
+ m.body_pos,
+ m.body_quat,
+ m.body_rootid,
+ m.body_subtreemass,
+ m._impl.body_tree,
+ m.body_weldid,
+ m.cam_bodyid,
+ m.cam_fovy,
+ m.cam_intrinsic,
+ m.cam_mat0,
+ m.cam_mode,
+ m.cam_pos,
+ m.cam_pos0,
+ m.cam_poscom0,
+ m.cam_quat,
+ m.cam_resolution,
+ m.cam_sensorsize,
+ m.cam_targetbodyid,
+ m._impl.condim_max,
+ m.dof_Madr,
+ m.dof_armature,
+ m.dof_bodyid,
+ m.dof_damping,
+ m.dof_frictionloss,
+ m.dof_invweight0,
+ m.dof_jntid,
+ m.dof_parentid,
+ m.dof_solimp,
+ m.dof_solref,
+ m._impl.dof_tri_col,
+ m._impl.dof_tri_row,
+ m._impl.eq_connect_adr,
+ m.eq_data,
+ m._impl.eq_jnt_adr,
+ m.eq_obj1id,
+ m.eq_obj2id,
+ m.eq_objtype,
+ m.eq_solimp,
+ m.eq_solref,
+ m._impl.eq_ten_adr,
+ m._impl.eq_wld_adr,
+ m._impl.flex_bending,
+ m._impl.flex_damping,
+ m._impl.flex_dim,
+ m._impl.flex_edge,
+ m._impl.flex_edgeadr,
+ m._impl.flex_edgeflap,
+ m._impl.flex_elem,
+ m._impl.flex_elemedge,
+ m._impl.flex_elemedgeadr,
+ m._impl.flex_stiffness,
+ m._impl.flex_vertadr,
+ m._impl.flex_vertbodyid,
+ m._impl.flexedge_length0,
+ m.geom_aabb,
+ m.geom_bodyid,
+ m.geom_condim,
+ m.geom_dataid,
+ m.geom_friction,
+ m.geom_gap,
+ m.geom_group,
+ m.geom_margin,
+ m.geom_matid,
+ m._impl.geom_pair_type_count,
+ m._impl.geom_plugin_index,
+ m.geom_pos,
+ m.geom_priority,
+ m.geom_quat,
+ m.geom_rbound,
+ m.geom_rgba,
+ m.geom_size,
+ m.geom_solimp,
+ m.geom_solmix,
+ m.geom_solref,
+ m.geom_type,
+ m._impl.geompair2hfgeompair,
+ m._impl.has_sdf_geom,
+ m.hfield_adr,
+ m.hfield_data,
+ m.hfield_ncol,
+ m.hfield_nrow,
+ m.hfield_size,
+ m.jnt_actfrclimited,
+ m.jnt_actfrcrange,
+ m.jnt_actgravcomp,
+ m.jnt_axis,
+ m.jnt_bodyid,
+ m.jnt_dofadr,
+ m._impl.jnt_limited_ball_adr,
+ m._impl.jnt_limited_slide_hinge_adr,
+ m.jnt_margin,
+ m.jnt_pos,
+ m.jnt_qposadr,
+ m.jnt_range,
+ m.jnt_solimp,
+ m.jnt_solref,
+ m.jnt_stiffness,
+ m.jnt_type,
+ m._impl.light_bodyid,
+ m.light_dir,
+ m.light_dir0,
+ m.light_mode,
+ m.light_pos,
+ m.light_pos0,
+ m.light_poscom0,
+ m._impl.light_targetbodyid,
+ m._impl.mapM2M,
+ m.mat_rgba,
+ m.mesh_face,
+ m.mesh_faceadr,
+ m.mesh_graph,
+ m.mesh_graphadr,
+ m._impl.mesh_polyadr,
+ m._impl.mesh_polymap,
+ m._impl.mesh_polymapadr,
+ m._impl.mesh_polymapnum,
+ m._impl.mesh_polynormal,
+ m._impl.mesh_polynum,
+ m._impl.mesh_polyvert,
+ m._impl.mesh_polyvertadr,
+ m._impl.mesh_polyvertnum,
+ m.mesh_vert,
+ m.mesh_vertadr,
+ m.mesh_vertnum,
+ m._impl.mocap_bodyid,
+ m.nC,
+ m.na,
+ m.nbody,
+ m.ncam,
+ m.neq,
+ m._impl.nflexedge,
+ m._impl.nflexelem,
+ m._impl.nflexvert,
+ m.ngeom,
+ m.ngravcomp,
+ m.nhfield,
+ m.njnt,
+ m.nlight,
+ m._impl.nlsp,
+ m.nmeshface,
+ m.nmocap,
+ m.nsite,
+ m.ntendon,
+ m.nu,
+ m.nv,
+ m._impl.nxn_geom_pair_filtered,
+ m._impl.nxn_pairid,
+ m._impl.nxn_pairid_filtered,
+ m.pair_dim,
+ m.pair_friction,
+ m.pair_gap,
+ m.pair_margin,
+ m.pair_solimp,
+ m.pair_solref,
+ m.pair_solreffriction,
+ m._impl.plugin,
+ m._impl.plugin_attr,
+ m._impl.qLD_updates,
+ m._impl.qM_fullm_i,
+ m._impl.qM_fullm_j,
+ m._impl.qM_madr_ij,
+ m._impl.qM_mulm_i,
+ m._impl.qM_mulm_j,
+ m._impl.qM_tiles,
+ m.qpos0,
+ m.qpos_spring,
+ m._impl.rangefinder_sensor_adr,
+ m._impl.sensor_acc_adr,
+ m.sensor_adr,
+ m.sensor_cutoff,
+ m.sensor_datatype,
+ m._impl.sensor_e_kinetic,
+ m._impl.sensor_e_potential,
+ m._impl.sensor_limitfrc_adr,
+ m._impl.sensor_limitpos_adr,
+ m._impl.sensor_limitvel_adr,
+ m.sensor_objid,
+ m.sensor_objtype,
+ m._impl.sensor_pos_adr,
+ m._impl.sensor_rangefinder_adr,
+ m._impl.sensor_rangefinder_bodyid,
+ m.sensor_refid,
+ m.sensor_reftype,
+ m._impl.sensor_rne_postconstraint,
+ m._impl.sensor_subtree_vel,
+ m._impl.sensor_tendonactfrc_adr,
+ m._impl.sensor_touch_adr,
+ m.sensor_type,
+ m._impl.sensor_vel_adr,
+ m.site_bodyid,
+ m.site_pos,
+ m.site_quat,
+ m.site_size,
+ m.site_type,
+ m._impl.subtree_mass,
+ m.tendon_actfrclimited,
+ m.tendon_actfrcrange,
+ m.tendon_adr,
+ m.tendon_armature,
+ m.tendon_damping,
+ m.tendon_frictionloss,
+ m._impl.tendon_geom_adr,
+ m.tendon_invweight0,
+ m._impl.tendon_jnt_adr,
+ m.tendon_length0,
+ m.tendon_lengthspring,
+ m._impl.tendon_limited_adr,
+ m.tendon_margin,
+ m.tendon_num,
+ m.tendon_range,
+ m._impl.tendon_site_pair_adr,
+ m.tendon_solimp_fri,
+ m.tendon_solimp_lim,
+ m.tendon_solref_fri,
+ m.tendon_solref_lim,
+ m.tendon_stiffness,
+ 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,
+ m.opt._impl.broadphase,
+ m.opt._impl.broadphase_filter,
+ m.opt.cone,
+ m.opt.density,
+ m.opt.disableflags,
+ m.opt.enableflags,
+ m.opt._impl.epa_iterations,
+ m.opt._impl.gjk_iterations,
+ m.opt._impl.graph_conditional,
+ m.opt.gravity,
+ m.opt._impl.has_fluid,
+ m.opt.impratio,
+ m.opt._impl.is_sparse,
+ m.opt.iterations,
+ m.opt.ls_iterations,
+ m.opt._impl.ls_parallel,
+ m.opt.ls_tolerance,
+ m.opt.magnetic,
+ m.opt._impl.run_collision_detection,
+ m.opt._impl.sdf_initpoints,
+ m.opt._impl.sdf_iterations,
+ m.opt.solver,
+ m.opt.timestep,
+ m.opt.tolerance,
+ m.opt.viscosity,
+ m.opt.wind,
+ m.stat.meaninertia,
+ d._impl.nconmax,
+ d._impl.njmax,
+ d.act,
+ d.act_dot,
+ d.actuator_force,
+ d._impl.actuator_length,
+ d._impl.actuator_moment,
+ d._impl.actuator_trntype_body_ncon,
+ d._impl.actuator_velocity,
+ d._impl.cacc,
+ d.cam_xmat,
+ d.cam_xpos,
+ d._impl.cdof,
+ d._impl.cdof_dot,
+ d._impl.cfrc_ext,
+ d._impl.cfrc_int,
+ d._impl.cinert,
+ d._impl.collision_hftri_index,
+ d._impl.collision_pair,
+ d._impl.collision_pairid,
+ d._impl.collision_worldid,
+ d._impl.crb,
+ d.ctrl,
+ d.cvel,
+ d._impl.energy,
+ d._impl.epa_face,
+ d._impl.epa_horizon,
+ d._impl.epa_index,
+ d._impl.epa_map,
+ d._impl.epa_norm2,
+ d._impl.epa_pr,
+ d._impl.epa_vert,
+ d._impl.epa_vert1,
+ d._impl.epa_vert2,
+ d._impl.epa_vert_index1,
+ d._impl.epa_vert_index2,
+ d.eq_active,
+ d._impl.flexedge_length,
+ d._impl.flexedge_velocity,
+ d._impl.flexvert_xpos,
+ d._impl.fluid_applied,
+ d._impl.geom_skip,
+ d.geom_xmat,
+ d.geom_xpos,
+ d._impl.light_xdir,
+ d._impl.light_xpos,
+ d.mocap_pos,
+ d.mocap_quat,
+ d._impl.ncollision,
+ d._impl.ncon,
+ d._impl.ncon_hfield,
+ d._impl.ne,
+ d._impl.ne_connect,
+ d._impl.ne_jnt,
+ d._impl.ne_ten,
+ d._impl.ne_weld,
+ d._impl.nefc,
+ d._impl.nf,
+ d._impl.nl,
+ d._impl.nsolving,
+ d._impl.qLD,
+ d._impl.qLDiagInv,
+ d._impl.qM,
+ d.qacc,
+ d.qacc_smooth,
+ d.qacc_warmstart,
+ d.qfrc_actuator,
+ d.qfrc_applied,
+ d.qfrc_bias,
+ d.qfrc_constraint,
+ d._impl.qfrc_damper,
+ d.qfrc_fluid,
+ d.qfrc_gravcomp,
+ d.qfrc_passive,
+ d.qfrc_smooth,
+ d._impl.qfrc_spring,
+ d.qpos,
+ d.qvel,
+ d._impl.sap_cumulative_sum,
+ d._impl.sap_projection_lower,
+ d._impl.sap_projection_upper,
+ d._impl.sap_range,
+ d._impl.sap_segment_index,
+ d._impl.sap_sort_index,
+ d._impl.sensor_rangefinder_dist,
+ d._impl.sensor_rangefinder_geomid,
+ d._impl.sensor_rangefinder_pnt,
+ d._impl.sensor_rangefinder_vec,
+ d.sensordata,
+ d.site_xmat,
+ d.site_xpos,
+ d._impl.solver_niter,
+ d._impl.subtree_angmom,
+ d._impl.subtree_bodyvel,
+ d.subtree_com,
+ d._impl.subtree_linvel,
+ d._impl.ten_J,
+ d._impl.ten_Jdot,
+ d._impl.ten_actfrc,
+ d._impl.ten_bias_coef,
+ d._impl.ten_length,
+ d._impl.ten_velocity,
+ d._impl.ten_wrapadr,
+ d._impl.ten_wrapnum,
+ d.time,
+ d._impl.wrap_geom_xpos,
+ d._impl.wrap_obj,
+ d._impl.wrap_xpos,
+ d.xanchor,
+ d.xaxis,
+ d.xfrc_applied,
+ d.ximat,
+ d.xipos,
+ d.xmat,
+ d.xpos,
+ d.xquat,
+ d._impl.contact__dim,
+ d._impl.contact__dist,
+ d._impl.contact__efc_address,
+ d._impl.contact__frame,
+ d._impl.contact__friction,
+ d._impl.contact__geom,
+ d._impl.contact__includemargin,
+ d._impl.contact__pos,
+ d._impl.contact__solimp,
+ d._impl.contact__solref,
+ d._impl.contact__solreffriction,
+ d._impl.contact__worldid,
+ d._impl.efc__D,
+ d._impl.efc__J,
+ d._impl.efc__Jaref,
+ d._impl.efc__Ma,
+ d._impl.efc__Mgrad,
+ d._impl.efc__active,
+ d._impl.efc__alpha,
+ d._impl.efc__aref,
+ d._impl.efc__beta,
+ d._impl.efc__beta_den,
+ d._impl.efc__beta_num,
+ d._impl.efc__cholesky_L_tmp,
+ d._impl.efc__cholesky_y_tmp,
+ d._impl.efc__condim,
+ d._impl.efc__cost,
+ d._impl.efc__cost_candidate,
+ d._impl.efc__done,
+ d._impl.efc__force,
+ d._impl.efc__frictionloss,
+ d._impl.efc__gauss,
+ d._impl.efc__grad,
+ d._impl.efc__grad_dot,
+ d._impl.efc__gtol,
+ d._impl.efc__h,
+ d._impl.efc__hi,
+ d._impl.efc__hi_alpha,
+ d._impl.efc__hi_next,
+ d._impl.efc__hi_next_alpha,
+ d._impl.efc__id,
+ d._impl.efc__jv,
+ d._impl.efc__lo,
+ d._impl.efc__lo_alpha,
+ d._impl.efc__lo_next,
+ d._impl.efc__lo_next_alpha,
+ d._impl.efc__ls_done,
+ d._impl.efc__margin,
+ d._impl.efc__mid,
+ d._impl.efc__mid_alpha,
+ d._impl.efc__mv,
+ d._impl.efc__p0,
+ d._impl.efc__pos,
+ d._impl.efc__prev_Mgrad,
+ d._impl.efc__prev_cost,
+ d._impl.efc__prev_grad,
+ d._impl.efc__quad,
+ d._impl.efc__quad_gauss,
+ d._impl.efc__search,
+ d._impl.efc__search_dot,
+ d._impl.efc__type,
+ d._impl.efc__u,
+ d._impl.efc__uu,
+ d._impl.efc__uv,
+ d._impl.efc__vel,
+ d._impl.efc__vv,
+ )
+ d = d.tree_replace({
+ 'act': out[0],
+ 'act_dot': out[1],
+ 'actuator_force': out[2],
+ '_impl.actuator_length': out[3],
+ '_impl.actuator_moment': out[4],
+ '_impl.actuator_trntype_body_ncon': out[5],
+ '_impl.actuator_velocity': out[6],
+ '_impl.cacc': out[7],
+ 'cam_xmat': out[8],
+ 'cam_xpos': out[9],
+ '_impl.cdof': out[10],
+ '_impl.cdof_dot': out[11],
+ '_impl.cfrc_ext': out[12],
+ '_impl.cfrc_int': out[13],
+ '_impl.cinert': out[14],
+ '_impl.collision_hftri_index': out[15],
+ '_impl.collision_pair': out[16],
+ '_impl.collision_pairid': out[17],
+ '_impl.collision_worldid': out[18],
+ '_impl.crb': out[19],
+ 'ctrl': out[20],
+ 'cvel': out[21],
+ '_impl.energy': out[22],
+ '_impl.epa_face': out[23],
+ '_impl.epa_horizon': out[24],
+ '_impl.epa_index': out[25],
+ '_impl.epa_map': out[26],
+ '_impl.epa_norm2': out[27],
+ '_impl.epa_pr': out[28],
+ '_impl.epa_vert': out[29],
+ '_impl.epa_vert1': out[30],
+ '_impl.epa_vert2': out[31],
+ '_impl.epa_vert_index1': out[32],
+ '_impl.epa_vert_index2': out[33],
+ 'eq_active': out[34],
+ '_impl.flexedge_length': out[35],
+ '_impl.flexedge_velocity': out[36],
+ '_impl.flexvert_xpos': out[37],
+ '_impl.fluid_applied': out[38],
+ '_impl.geom_skip': out[39],
+ 'geom_xmat': out[40],
+ 'geom_xpos': out[41],
+ '_impl.light_xdir': out[42],
+ '_impl.light_xpos': out[43],
+ 'mocap_pos': out[44],
+ 'mocap_quat': out[45],
+ '_impl.ncollision': out[46],
+ '_impl.ncon': out[47],
+ '_impl.ncon_hfield': out[48],
+ '_impl.ne': out[49],
+ '_impl.ne_connect': out[50],
+ '_impl.ne_jnt': out[51],
+ '_impl.ne_ten': out[52],
+ '_impl.ne_weld': out[53],
+ '_impl.nefc': out[54],
+ '_impl.nf': out[55],
+ '_impl.nl': out[56],
+ '_impl.nsolving': out[57],
+ '_impl.qLD': out[58],
+ '_impl.qLDiagInv': out[59],
+ '_impl.qM': out[60],
+ 'qacc': out[61],
+ 'qacc_smooth': out[62],
+ 'qacc_warmstart': out[63],
+ 'qfrc_actuator': out[64],
+ 'qfrc_applied': out[65],
+ 'qfrc_bias': out[66],
+ 'qfrc_constraint': out[67],
+ '_impl.qfrc_damper': out[68],
+ 'qfrc_fluid': out[69],
+ 'qfrc_gravcomp': out[70],
+ 'qfrc_passive': out[71],
+ 'qfrc_smooth': out[72],
+ '_impl.qfrc_spring': out[73],
+ 'qpos': out[74],
+ 'qvel': out[75],
+ '_impl.sap_cumulative_sum': out[76],
+ '_impl.sap_projection_lower': out[77],
+ '_impl.sap_projection_upper': out[78],
+ '_impl.sap_range': out[79],
+ '_impl.sap_segment_index': out[80],
+ '_impl.sap_sort_index': out[81],
+ '_impl.sensor_rangefinder_dist': out[82],
+ '_impl.sensor_rangefinder_geomid': out[83],
+ '_impl.sensor_rangefinder_pnt': out[84],
+ '_impl.sensor_rangefinder_vec': out[85],
+ 'sensordata': out[86],
+ 'site_xmat': out[87],
+ 'site_xpos': out[88],
+ '_impl.solver_niter': out[89],
+ '_impl.subtree_angmom': out[90],
+ '_impl.subtree_bodyvel': out[91],
+ 'subtree_com': out[92],
+ '_impl.subtree_linvel': out[93],
+ '_impl.ten_J': out[94],
+ '_impl.ten_Jdot': out[95],
+ '_impl.ten_actfrc': out[96],
+ '_impl.ten_bias_coef': out[97],
+ '_impl.ten_length': out[98],
+ '_impl.ten_velocity': out[99],
+ '_impl.ten_wrapadr': out[100],
+ '_impl.ten_wrapnum': out[101],
+ 'time': out[102],
+ '_impl.wrap_geom_xpos': out[103],
+ '_impl.wrap_obj': out[104],
+ '_impl.wrap_xpos': out[105],
+ 'xanchor': out[106],
+ 'xaxis': out[107],
+ 'xfrc_applied': out[108],
+ 'ximat': out[109],
+ 'xipos': out[110],
+ 'xmat': out[111],
+ 'xpos': out[112],
+ 'xquat': out[113],
+ '_impl.contact__dim': out[114],
+ '_impl.contact__dist': out[115],
+ '_impl.contact__efc_address': out[116],
+ '_impl.contact__frame': out[117],
+ '_impl.contact__friction': out[118],
+ '_impl.contact__geom': out[119],
+ '_impl.contact__includemargin': out[120],
+ '_impl.contact__pos': out[121],
+ '_impl.contact__solimp': out[122],
+ '_impl.contact__solref': out[123],
+ '_impl.contact__solreffriction': out[124],
+ '_impl.contact__worldid': out[125],
+ '_impl.efc__D': out[126],
+ '_impl.efc__J': out[127],
+ '_impl.efc__Jaref': out[128],
+ '_impl.efc__Ma': out[129],
+ '_impl.efc__Mgrad': out[130],
+ '_impl.efc__active': out[131],
+ '_impl.efc__alpha': out[132],
+ '_impl.efc__aref': out[133],
+ '_impl.efc__beta': out[134],
+ '_impl.efc__beta_den': out[135],
+ '_impl.efc__beta_num': out[136],
+ '_impl.efc__cholesky_L_tmp': out[137],
+ '_impl.efc__cholesky_y_tmp': out[138],
+ '_impl.efc__condim': out[139],
+ '_impl.efc__cost': out[140],
+ '_impl.efc__cost_candidate': out[141],
+ '_impl.efc__done': out[142],
+ '_impl.efc__force': out[143],
+ '_impl.efc__frictionloss': out[144],
+ '_impl.efc__gauss': out[145],
+ '_impl.efc__grad': out[146],
+ '_impl.efc__grad_dot': out[147],
+ '_impl.efc__gtol': out[148],
+ '_impl.efc__h': out[149],
+ '_impl.efc__hi': out[150],
+ '_impl.efc__hi_alpha': out[151],
+ '_impl.efc__hi_next': out[152],
+ '_impl.efc__hi_next_alpha': out[153],
+ '_impl.efc__id': out[154],
+ '_impl.efc__jv': out[155],
+ '_impl.efc__lo': out[156],
+ '_impl.efc__lo_alpha': out[157],
+ '_impl.efc__lo_next': out[158],
+ '_impl.efc__lo_next_alpha': out[159],
+ '_impl.efc__ls_done': out[160],
+ '_impl.efc__margin': out[161],
+ '_impl.efc__mid': out[162],
+ '_impl.efc__mid_alpha': out[163],
+ '_impl.efc__mv': out[164],
+ '_impl.efc__p0': out[165],
+ '_impl.efc__pos': out[166],
+ '_impl.efc__prev_Mgrad': out[167],
+ '_impl.efc__prev_cost': out[168],
+ '_impl.efc__prev_grad': out[169],
+ '_impl.efc__quad': out[170],
+ '_impl.efc__quad_gauss': out[171],
+ '_impl.efc__search': out[172],
+ '_impl.efc__search_dot': out[173],
+ '_impl.efc__type': out[174],
+ '_impl.efc__u': out[175],
+ '_impl.efc__uu': out[176],
+ '_impl.efc__uv': out[177],
+ '_impl.efc__vel': out[178],
+ '_impl.efc__vv': out[179],
+ })
+ return d
+
+
+@jax.custom_batching.custom_vmap
+@ffi.marshal_jax_warp_callable
+def forward(m: types.Model, d: types.Data):
+ return _forward_jax_impl(m, d)
+
+
+@forward.def_vmap
+@ffi.marshal_custom_vmap
+def forward_vmap(unused_axis_size, is_batched, m, d):
+ d = forward(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 _step_shim(
+ # Model
+ nworld: int,
+ M_rowadr: wp.array(dtype=int),
+ M_rownnz: wp.array(dtype=int),
+ actuator_acc0: wp.array(dtype=float),
+ actuator_actadr: wp.array(dtype=int),
+ actuator_actearly: wp.array(dtype=bool),
+ actuator_actlimited: wp.array(dtype=bool),
+ actuator_actnum: wp.array(dtype=int),
+ actuator_actrange: wp.array2d(dtype=wp.vec2),
+ actuator_affine_bias_gain: bool,
+ actuator_biasprm: wp.array2d(dtype=mjwp_types.vec10f),
+ actuator_biastype: wp.array(dtype=int),
+ actuator_cranklength: wp.array(dtype=float),
+ actuator_ctrllimited: wp.array(dtype=bool),
+ actuator_ctrlrange: wp.array2d(dtype=wp.vec2),
+ actuator_dynprm: wp.array2d(dtype=mjwp_types.vec10f),
+ actuator_dyntype: wp.array(dtype=int),
+ actuator_forcelimited: wp.array(dtype=bool),
+ actuator_forcerange: wp.array2d(dtype=wp.vec2),
+ actuator_gainprm: wp.array2d(dtype=mjwp_types.vec10f),
+ actuator_gaintype: wp.array(dtype=int),
+ actuator_gear: wp.array2d(dtype=wp.spatial_vector),
+ actuator_lengthrange: wp.array(dtype=wp.vec2),
+ actuator_trnid: wp.array(dtype=wp.vec2i),
+ actuator_trntype: wp.array(dtype=int),
+ actuator_trntype_body_adr: wp.array(dtype=int),
+ block_dim: mjwp_types.BlockDim,
+ body_dofadr: wp.array(dtype=int),
+ body_dofnum: wp.array(dtype=int),
+ body_gravcomp: wp.array2d(dtype=float),
+ body_inertia: wp.array2d(dtype=wp.vec3),
+ body_invweight0: wp.array2d(dtype=wp.vec2),
+ body_ipos: wp.array2d(dtype=wp.vec3),
+ body_iquat: wp.array2d(dtype=wp.quat),
+ body_jntadr: wp.array(dtype=int),
+ body_jntnum: wp.array(dtype=int),
+ body_mass: wp.array2d(dtype=float),
+ body_parentid: wp.array(dtype=int),
+ body_pos: wp.array2d(dtype=wp.vec3),
+ body_quat: wp.array2d(dtype=wp.quat),
+ body_rootid: wp.array(dtype=int),
+ body_subtreemass: wp.array2d(dtype=float),
+ body_tree: tuple[wp.array(dtype=int), ...],
+ body_weldid: wp.array(dtype=int),
+ cam_bodyid: wp.array(dtype=int),
+ cam_fovy: wp.array(dtype=float),
+ cam_intrinsic: wp.array(dtype=wp.vec4),
+ cam_mat0: wp.array2d(dtype=wp.mat33),
+ cam_mode: wp.array(dtype=int),
+ cam_pos: wp.array2d(dtype=wp.vec3),
+ cam_pos0: wp.array2d(dtype=wp.vec3),
+ cam_poscom0: wp.array2d(dtype=wp.vec3),
+ cam_quat: wp.array2d(dtype=wp.quat),
+ cam_resolution: wp.array(dtype=wp.vec2i),
+ cam_sensorsize: wp.array(dtype=wp.vec2),
+ cam_targetbodyid: wp.array(dtype=int),
+ condim_max: int,
+ dof_Madr: wp.array(dtype=int),
+ dof_armature: wp.array2d(dtype=float),
+ dof_bodyid: wp.array(dtype=int),
+ dof_damping: wp.array2d(dtype=float),
+ dof_frictionloss: wp.array2d(dtype=float),
+ dof_invweight0: wp.array2d(dtype=float),
+ dof_jntid: wp.array(dtype=int),
+ dof_parentid: wp.array(dtype=int),
+ dof_solimp: wp.array2d(dtype=mjwp_types.vec5),
+ dof_solref: wp.array2d(dtype=wp.vec2),
+ dof_tri_col: wp.array(dtype=int),
+ dof_tri_row: wp.array(dtype=int),
+ eq_connect_adr: wp.array(dtype=int),
+ eq_data: wp.array2d(dtype=mjwp_types.vec11),
+ eq_jnt_adr: wp.array(dtype=int),
+ eq_obj1id: wp.array(dtype=int),
+ eq_obj2id: wp.array(dtype=int),
+ eq_objtype: wp.array(dtype=int),
+ eq_solimp: wp.array2d(dtype=mjwp_types.vec5),
+ eq_solref: wp.array2d(dtype=wp.vec2),
+ eq_ten_adr: wp.array(dtype=int),
+ eq_wld_adr: wp.array(dtype=int),
+ flex_bending: wp.array(dtype=wp.mat44f),
+ flex_damping: wp.array(dtype=float),
+ flex_dim: wp.array(dtype=int),
+ flex_edge: wp.array(dtype=wp.vec2i),
+ flex_edgeadr: wp.array(dtype=int),
+ flex_edgeflap: wp.array(dtype=wp.vec2i),
+ flex_elem: wp.array(dtype=int),
+ flex_elemedge: wp.array(dtype=int),
+ flex_elemedgeadr: wp.array(dtype=int),
+ flex_stiffness: wp.array(dtype=float),
+ flex_vertadr: wp.array(dtype=int),
+ flex_vertbodyid: wp.array(dtype=int),
+ flexedge_length0: wp.array(dtype=float),
+ geom_aabb: wp.array2d(dtype=wp.vec3),
+ geom_bodyid: wp.array(dtype=int),
+ geom_condim: wp.array(dtype=int),
+ geom_dataid: wp.array(dtype=int),
+ geom_friction: wp.array2d(dtype=wp.vec3),
+ geom_gap: wp.array2d(dtype=float),
+ geom_group: wp.array(dtype=int),
+ geom_margin: wp.array2d(dtype=float),
+ geom_matid: wp.array2d(dtype=int),
+ geom_pair_type_count: tuple[int, ...],
+ geom_plugin_index: wp.array(dtype=int),
+ geom_pos: wp.array2d(dtype=wp.vec3),
+ geom_priority: wp.array(dtype=int),
+ geom_quat: wp.array2d(dtype=wp.quat),
+ geom_rbound: wp.array2d(dtype=float),
+ geom_rgba: wp.array2d(dtype=wp.vec4),
+ geom_size: wp.array2d(dtype=wp.vec3),
+ geom_solimp: wp.array2d(dtype=mjwp_types.vec5),
+ geom_solmix: wp.array2d(dtype=float),
+ geom_solref: wp.array2d(dtype=wp.vec2),
+ geom_type: wp.array(dtype=int),
+ geompair2hfgeompair: wp.array(dtype=int),
+ has_sdf_geom: bool,
+ hfield_adr: wp.array(dtype=int),
+ hfield_data: wp.array(dtype=float),
+ hfield_ncol: wp.array(dtype=int),
+ hfield_nrow: wp.array(dtype=int),
+ hfield_size: wp.array(dtype=wp.vec4),
+ jnt_actfrclimited: wp.array(dtype=bool),
+ jnt_actfrcrange: wp.array2d(dtype=wp.vec2),
+ jnt_actgravcomp: wp.array(dtype=int),
+ jnt_axis: wp.array2d(dtype=wp.vec3),
+ jnt_bodyid: wp.array(dtype=int),
+ jnt_dofadr: wp.array(dtype=int),
+ jnt_limited_ball_adr: wp.array(dtype=int),
+ jnt_limited_slide_hinge_adr: wp.array(dtype=int),
+ jnt_margin: wp.array2d(dtype=float),
+ jnt_pos: wp.array2d(dtype=wp.vec3),
+ jnt_qposadr: wp.array(dtype=int),
+ jnt_range: wp.array2d(dtype=wp.vec2),
+ jnt_solimp: wp.array2d(dtype=mjwp_types.vec5),
+ jnt_solref: wp.array2d(dtype=wp.vec2),
+ jnt_stiffness: wp.array2d(dtype=float),
+ jnt_type: wp.array(dtype=int),
+ light_bodyid: wp.array(dtype=int),
+ light_dir: wp.array2d(dtype=wp.vec3),
+ light_dir0: wp.array2d(dtype=wp.vec3),
+ light_mode: wp.array(dtype=int),
+ light_pos: wp.array2d(dtype=wp.vec3),
+ light_pos0: wp.array2d(dtype=wp.vec3),
+ light_poscom0: wp.array2d(dtype=wp.vec3),
+ light_targetbodyid: wp.array(dtype=int),
+ mapM2M: wp.array(dtype=int),
+ mat_rgba: wp.array2d(dtype=wp.vec4),
+ mesh_face: wp.array(dtype=wp.vec3i),
+ mesh_faceadr: wp.array(dtype=int),
+ mesh_graph: wp.array(dtype=int),
+ mesh_graphadr: wp.array(dtype=int),
+ mesh_polyadr: wp.array(dtype=int),
+ mesh_polymap: wp.array(dtype=int),
+ mesh_polymapadr: wp.array(dtype=int),
+ mesh_polymapnum: wp.array(dtype=int),
+ mesh_polynormal: wp.array(dtype=wp.vec3),
+ mesh_polynum: wp.array(dtype=int),
+ mesh_polyvert: wp.array(dtype=int),
+ mesh_polyvertadr: wp.array(dtype=int),
+ mesh_polyvertnum: wp.array(dtype=int),
+ mesh_vert: wp.array(dtype=wp.vec3),
+ mesh_vertadr: wp.array(dtype=int),
+ mesh_vertnum: wp.array(dtype=int),
+ mocap_bodyid: wp.array(dtype=int),
+ nC: int,
+ na: int,
+ nbody: int,
+ ncam: int,
+ neq: int,
+ nflexedge: int,
+ nflexelem: int,
+ nflexvert: int,
+ ngeom: int,
+ ngravcomp: int,
+ nhfield: int,
+ njnt: int,
+ nlight: int,
+ nlsp: int,
+ nmeshface: int,
+ nmocap: int,
+ nsite: int,
+ ntendon: int,
+ nu: int,
+ nv: int,
+ nxn_geom_pair_filtered: wp.array(dtype=wp.vec2i),
+ nxn_pairid: wp.array(dtype=int),
+ nxn_pairid_filtered: wp.array(dtype=int),
+ pair_dim: wp.array(dtype=int),
+ pair_friction: wp.array2d(dtype=mjwp_types.vec5),
+ pair_gap: wp.array2d(dtype=float),
+ pair_margin: wp.array2d(dtype=float),
+ pair_solimp: wp.array2d(dtype=mjwp_types.vec5),
+ pair_solref: wp.array2d(dtype=wp.vec2),
+ pair_solreffriction: wp.array2d(dtype=wp.vec2),
+ plugin: wp.array(dtype=int),
+ plugin_attr: wp.array(dtype=wp.vec3f),
+ qLD_updates: tuple[wp.array(dtype=wp.vec3i), ...],
+ qM_fullm_i: wp.array(dtype=int),
+ qM_fullm_j: wp.array(dtype=int),
+ qM_madr_ij: wp.array(dtype=int),
+ qM_mulm_i: wp.array(dtype=int),
+ qM_mulm_j: wp.array(dtype=int),
+ qM_tiles: tuple[mjwp_types.TileSet, ...],
+ qpos0: wp.array2d(dtype=float),
+ qpos_spring: wp.array2d(dtype=float),
+ rangefinder_sensor_adr: wp.array(dtype=int),
+ sensor_acc_adr: wp.array(dtype=int),
+ sensor_adr: wp.array(dtype=int),
+ sensor_cutoff: wp.array(dtype=float),
+ sensor_datatype: wp.array(dtype=int),
+ sensor_e_kinetic: bool,
+ sensor_e_potential: bool,
+ sensor_limitfrc_adr: wp.array(dtype=int),
+ sensor_limitpos_adr: wp.array(dtype=int),
+ sensor_limitvel_adr: wp.array(dtype=int),
+ sensor_objid: wp.array(dtype=int),
+ sensor_objtype: wp.array(dtype=int),
+ sensor_pos_adr: wp.array(dtype=int),
+ sensor_rangefinder_adr: wp.array(dtype=int),
+ sensor_rangefinder_bodyid: wp.array(dtype=int),
+ sensor_refid: wp.array(dtype=int),
+ sensor_reftype: wp.array(dtype=int),
+ sensor_rne_postconstraint: bool,
+ sensor_subtree_vel: bool,
+ sensor_tendonactfrc_adr: wp.array(dtype=int),
+ sensor_touch_adr: wp.array(dtype=int),
+ sensor_type: wp.array(dtype=int),
+ sensor_vel_adr: wp.array(dtype=int),
+ site_bodyid: wp.array(dtype=int),
+ site_pos: wp.array2d(dtype=wp.vec3),
+ site_quat: wp.array2d(dtype=wp.quat),
+ site_size: wp.array(dtype=wp.vec3),
+ site_type: wp.array(dtype=int),
+ subtree_mass: wp.array2d(dtype=float),
+ tendon_actfrclimited: wp.array(dtype=bool),
+ tendon_actfrcrange: wp.array2d(dtype=wp.vec2),
+ tendon_adr: wp.array(dtype=int),
+ tendon_armature: wp.array2d(dtype=float),
+ tendon_damping: wp.array2d(dtype=float),
+ tendon_frictionloss: wp.array2d(dtype=float),
+ tendon_geom_adr: wp.array(dtype=int),
+ tendon_invweight0: wp.array2d(dtype=float),
+ tendon_jnt_adr: wp.array(dtype=int),
+ tendon_length0: wp.array2d(dtype=float),
+ tendon_lengthspring: wp.array2d(dtype=wp.vec2),
+ tendon_limited_adr: wp.array(dtype=int),
+ tendon_margin: wp.array2d(dtype=float),
+ tendon_num: wp.array(dtype=int),
+ tendon_range: wp.array2d(dtype=wp.vec2),
+ tendon_site_pair_adr: wp.array(dtype=int),
+ tendon_solimp_fri: wp.array2d(dtype=mjwp_types.vec5),
+ tendon_solimp_lim: wp.array2d(dtype=mjwp_types.vec5),
+ tendon_solref_fri: wp.array2d(dtype=wp.vec2),
+ tendon_solref_lim: wp.array2d(dtype=wp.vec2),
+ tendon_stiffness: wp.array2d(dtype=float),
+ 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),
+ opt__broadphase: int,
+ opt__broadphase_filter: int,
+ opt__cone: int,
+ opt__density: wp.array(dtype=float),
+ opt__disableflags: int,
+ opt__enableflags: int,
+ opt__epa_iterations: int,
+ opt__gjk_iterations: int,
+ opt__graph_conditional: bool,
+ opt__gravity: wp.array(dtype=wp.vec3),
+ opt__has_fluid: bool,
+ opt__impratio: wp.array(dtype=float),
+ opt__integrator: int,
+ opt__is_sparse: bool,
+ opt__iterations: int,
+ opt__ls_iterations: int,
+ opt__ls_parallel: bool,
+ opt__ls_tolerance: wp.array(dtype=float),
+ opt__magnetic: wp.array(dtype=wp.vec3),
+ opt__run_collision_detection: bool,
+ opt__sdf_initpoints: int,
+ opt__sdf_iterations: int,
+ opt__solver: int,
+ opt__timestep: wp.array(dtype=float),
+ opt__tolerance: wp.array(dtype=float),
+ opt__viscosity: wp.array(dtype=float),
+ opt__wind: wp.array(dtype=wp.vec3),
+ stat__meaninertia: float,
+ # Data
+ nconmax: int,
+ njmax: int,
+ act: wp.array2d(dtype=float),
+ act_dot: wp.array2d(dtype=float),
+ act_dot_rk: wp.array2d(dtype=float),
+ act_t0: wp.array2d(dtype=float),
+ actuator_force: wp.array2d(dtype=float),
+ actuator_length: wp.array2d(dtype=float),
+ actuator_moment: wp.array3d(dtype=float),
+ actuator_trntype_body_ncon: wp.array2d(dtype=int),
+ actuator_velocity: wp.array2d(dtype=float),
+ cacc: wp.array2d(dtype=wp.spatial_vector),
+ cam_xmat: wp.array2d(dtype=wp.mat33),
+ cam_xpos: wp.array2d(dtype=wp.vec3),
+ cdof: wp.array2d(dtype=wp.spatial_vector),
+ cdof_dot: wp.array2d(dtype=wp.spatial_vector),
+ cfrc_ext: wp.array2d(dtype=wp.spatial_vector),
+ cfrc_int: wp.array2d(dtype=wp.spatial_vector),
+ cinert: wp.array2d(dtype=mjwp_types.vec10),
+ collision_hftri_index: wp.array(dtype=int),
+ collision_pair: wp.array(dtype=wp.vec2i),
+ collision_pairid: wp.array(dtype=int),
+ collision_worldid: wp.array(dtype=int),
+ crb: wp.array2d(dtype=mjwp_types.vec10),
+ ctrl: wp.array2d(dtype=float),
+ cvel: wp.array2d(dtype=wp.spatial_vector),
+ energy: wp.array(dtype=wp.vec2),
+ epa_face: wp.array2d(dtype=wp.vec3i),
+ epa_horizon: wp.array2d(dtype=int),
+ epa_index: wp.array2d(dtype=int),
+ epa_map: wp.array2d(dtype=int),
+ epa_norm2: wp.array2d(dtype=float),
+ epa_pr: wp.array2d(dtype=wp.vec3),
+ epa_vert: wp.array2d(dtype=wp.vec3),
+ epa_vert1: wp.array2d(dtype=wp.vec3),
+ epa_vert2: wp.array2d(dtype=wp.vec3),
+ epa_vert_index1: wp.array2d(dtype=int),
+ epa_vert_index2: wp.array2d(dtype=int),
+ eq_active: wp.array2d(dtype=bool),
+ flexedge_length: wp.array2d(dtype=float),
+ flexedge_velocity: wp.array2d(dtype=float),
+ flexvert_xpos: wp.array2d(dtype=wp.vec3),
+ fluid_applied: wp.array2d(dtype=wp.spatial_vector),
+ geom_skip: wp.array(dtype=bool),
+ geom_xmat: wp.array2d(dtype=wp.mat33),
+ geom_xpos: wp.array2d(dtype=wp.vec3),
+ inverse_mul_m_skip: wp.array(dtype=bool),
+ light_xdir: wp.array2d(dtype=wp.vec3),
+ light_xpos: wp.array2d(dtype=wp.vec3),
+ mocap_pos: wp.array2d(dtype=wp.vec3),
+ mocap_quat: wp.array2d(dtype=wp.quat),
+ ncollision: wp.array(dtype=int),
+ ncon: wp.array(dtype=int),
+ ncon_hfield: wp.array2d(dtype=int),
+ ne: wp.array(dtype=int),
+ ne_connect: wp.array(dtype=int),
+ ne_jnt: wp.array(dtype=int),
+ ne_ten: wp.array(dtype=int),
+ ne_weld: wp.array(dtype=int),
+ nefc: wp.array(dtype=int),
+ nf: wp.array(dtype=int),
+ nl: wp.array(dtype=int),
+ nsolving: wp.array(dtype=int),
+ qLD: wp.array3d(dtype=float),
+ qLD_integration: wp.array3d(dtype=float),
+ qLDiagInv: wp.array2d(dtype=float),
+ qLDiagInv_integration: wp.array2d(dtype=float),
+ qM: wp.array3d(dtype=float),
+ qM_integration: wp.array3d(dtype=float),
+ qacc: wp.array2d(dtype=float),
+ qacc_integration: wp.array2d(dtype=float),
+ qacc_rk: wp.array2d(dtype=float),
+ qacc_smooth: wp.array2d(dtype=float),
+ qacc_warmstart: wp.array2d(dtype=float),
+ qfrc_actuator: wp.array2d(dtype=float),
+ qfrc_applied: wp.array2d(dtype=float),
+ qfrc_bias: wp.array2d(dtype=float),
+ qfrc_constraint: wp.array2d(dtype=float),
+ qfrc_damper: wp.array2d(dtype=float),
+ qfrc_fluid: wp.array2d(dtype=float),
+ qfrc_gravcomp: wp.array2d(dtype=float),
+ qfrc_integration: wp.array2d(dtype=float),
+ qfrc_passive: wp.array2d(dtype=float),
+ qfrc_smooth: wp.array2d(dtype=float),
+ qfrc_spring: wp.array2d(dtype=float),
+ qpos: wp.array2d(dtype=float),
+ qpos_t0: wp.array2d(dtype=float),
+ qvel: wp.array2d(dtype=float),
+ qvel_rk: wp.array2d(dtype=float),
+ qvel_t0: wp.array2d(dtype=float),
+ sap_cumulative_sum: wp.array2d(dtype=int),
+ sap_projection_lower: wp.array3d(dtype=float),
+ sap_projection_upper: wp.array2d(dtype=float),
+ sap_range: wp.array2d(dtype=int),
+ sap_segment_index: wp.array2d(dtype=int),
+ sap_sort_index: wp.array3d(dtype=int),
+ sensor_rangefinder_dist: wp.array2d(dtype=float),
+ sensor_rangefinder_geomid: wp.array2d(dtype=int),
+ sensor_rangefinder_pnt: wp.array2d(dtype=wp.vec3),
+ sensor_rangefinder_vec: wp.array2d(dtype=wp.vec3),
+ sensordata: wp.array2d(dtype=float),
+ site_xmat: wp.array2d(dtype=wp.mat33),
+ site_xpos: wp.array2d(dtype=wp.vec3),
+ solver_niter: wp.array(dtype=int),
+ subtree_angmom: wp.array2d(dtype=wp.vec3),
+ subtree_bodyvel: wp.array2d(dtype=wp.spatial_vector),
+ subtree_com: wp.array2d(dtype=wp.vec3),
+ subtree_linvel: wp.array2d(dtype=wp.vec3),
+ ten_J: wp.array3d(dtype=float),
+ ten_Jdot: wp.array3d(dtype=float),
+ ten_actfrc: wp.array2d(dtype=float),
+ ten_bias_coef: wp.array2d(dtype=float),
+ ten_length: wp.array2d(dtype=float),
+ ten_velocity: wp.array2d(dtype=float),
+ ten_wrapadr: wp.array2d(dtype=int),
+ ten_wrapnum: wp.array2d(dtype=int),
+ time: wp.array(dtype=float),
+ wrap_geom_xpos: wp.array2d(dtype=wp.spatial_vector),
+ wrap_obj: wp.array2d(dtype=wp.vec2i),
+ wrap_xpos: wp.array2d(dtype=wp.spatial_vector),
+ xanchor: wp.array2d(dtype=wp.vec3),
+ xaxis: wp.array2d(dtype=wp.vec3),
+ xfrc_applied: wp.array2d(dtype=wp.spatial_vector),
+ ximat: wp.array2d(dtype=wp.mat33),
+ xipos: wp.array2d(dtype=wp.vec3),
+ xmat: wp.array2d(dtype=wp.mat33),
+ xpos: wp.array2d(dtype=wp.vec3),
+ xquat: wp.array2d(dtype=wp.quat),
+ contact__dim: wp.array(dtype=int),
+ contact__dist: wp.array(dtype=float),
+ contact__efc_address: wp.array2d(dtype=int),
+ contact__frame: wp.array(dtype=wp.mat33),
+ contact__friction: wp.array(dtype=mjwp_types.vec5),
+ contact__geom: wp.array(dtype=wp.vec2i),
+ contact__includemargin: wp.array(dtype=float),
+ contact__pos: wp.array(dtype=wp.vec3),
+ contact__solimp: wp.array(dtype=mjwp_types.vec5),
+ contact__solref: wp.array(dtype=wp.vec2),
+ contact__solreffriction: wp.array(dtype=wp.vec2),
+ contact__worldid: wp.array(dtype=int),
+ efc__D: wp.array2d(dtype=float),
+ efc__J: wp.array3d(dtype=float),
+ efc__Jaref: wp.array2d(dtype=float),
+ efc__Ma: wp.array2d(dtype=float),
+ efc__Mgrad: wp.array2d(dtype=float),
+ efc__active: wp.array2d(dtype=bool),
+ efc__alpha: wp.array(dtype=float),
+ efc__aref: wp.array2d(dtype=float),
+ efc__beta: wp.array(dtype=float),
+ efc__beta_den: wp.array(dtype=float),
+ efc__beta_num: wp.array(dtype=float),
+ efc__cholesky_L_tmp: wp.array3d(dtype=float),
+ efc__cholesky_y_tmp: wp.array2d(dtype=float),
+ efc__condim: wp.array2d(dtype=int),
+ efc__cost: wp.array(dtype=float),
+ efc__cost_candidate: wp.array2d(dtype=float),
+ efc__done: wp.array(dtype=bool),
+ efc__force: wp.array2d(dtype=float),
+ efc__frictionloss: wp.array2d(dtype=float),
+ efc__gauss: wp.array(dtype=float),
+ efc__grad: wp.array2d(dtype=float),
+ efc__grad_dot: wp.array(dtype=float),
+ efc__gtol: wp.array(dtype=float),
+ efc__h: wp.array3d(dtype=float),
+ efc__hi: wp.array(dtype=wp.vec3),
+ efc__hi_alpha: wp.array(dtype=float),
+ efc__hi_next: wp.array(dtype=wp.vec3),
+ efc__hi_next_alpha: wp.array(dtype=float),
+ efc__id: wp.array2d(dtype=int),
+ efc__jv: wp.array2d(dtype=float),
+ efc__lo: wp.array(dtype=wp.vec3),
+ efc__lo_alpha: wp.array(dtype=float),
+ efc__lo_next: wp.array(dtype=wp.vec3),
+ efc__lo_next_alpha: wp.array(dtype=float),
+ efc__ls_done: wp.array(dtype=bool),
+ efc__margin: wp.array2d(dtype=float),
+ efc__mid: wp.array(dtype=wp.vec3),
+ efc__mid_alpha: wp.array(dtype=float),
+ efc__mv: wp.array2d(dtype=float),
+ efc__p0: wp.array(dtype=wp.vec3),
+ efc__pos: wp.array2d(dtype=float),
+ efc__prev_Mgrad: wp.array2d(dtype=float),
+ efc__prev_cost: wp.array(dtype=float),
+ efc__prev_grad: wp.array2d(dtype=float),
+ efc__quad: wp.array2d(dtype=wp.vec3),
+ efc__quad_gauss: wp.array(dtype=wp.vec3),
+ efc__search: wp.array2d(dtype=float),
+ efc__search_dot: wp.array(dtype=float),
+ efc__type: wp.array2d(dtype=int),
+ efc__u: wp.array(dtype=mjwp_types.vec6),
+ efc__uu: wp.array(dtype=float),
+ efc__uv: wp.array(dtype=float),
+ efc__vel: wp.array2d(dtype=float),
+ efc__vv: wp.array(dtype=float),
+):
+ _m.stat = _s
+ _m.opt = _o
+ _d.efc = _e
+ _d.contact = _c
+ _m.M_rowadr = M_rowadr
+ _m.M_rownnz = M_rownnz
+ _m.actuator_acc0 = actuator_acc0
+ _m.actuator_actadr = actuator_actadr
+ _m.actuator_actearly = actuator_actearly
+ _m.actuator_actlimited = actuator_actlimited
+ _m.actuator_actnum = actuator_actnum
+ _m.actuator_actrange = actuator_actrange
+ _m.actuator_affine_bias_gain = actuator_affine_bias_gain
+ _m.actuator_biasprm = actuator_biasprm
+ _m.actuator_biastype = actuator_biastype
+ _m.actuator_cranklength = actuator_cranklength
+ _m.actuator_ctrllimited = actuator_ctrllimited
+ _m.actuator_ctrlrange = actuator_ctrlrange
+ _m.actuator_dynprm = actuator_dynprm
+ _m.actuator_dyntype = actuator_dyntype
+ _m.actuator_forcelimited = actuator_forcelimited
+ _m.actuator_forcerange = actuator_forcerange
+ _m.actuator_gainprm = actuator_gainprm
+ _m.actuator_gaintype = actuator_gaintype
+ _m.actuator_gear = actuator_gear
+ _m.actuator_lengthrange = actuator_lengthrange
+ _m.actuator_trnid = actuator_trnid
+ _m.actuator_trntype = actuator_trntype
+ _m.actuator_trntype_body_adr = actuator_trntype_body_adr
+ _m.block_dim = block_dim
+ _m.body_dofadr = body_dofadr
+ _m.body_dofnum = body_dofnum
+ _m.body_gravcomp = body_gravcomp
+ _m.body_inertia = body_inertia
+ _m.body_invweight0 = body_invweight0
+ _m.body_ipos = body_ipos
+ _m.body_iquat = body_iquat
+ _m.body_jntadr = body_jntadr
+ _m.body_jntnum = body_jntnum
+ _m.body_mass = body_mass
+ _m.body_parentid = body_parentid
+ _m.body_pos = body_pos
+ _m.body_quat = body_quat
+ _m.body_rootid = body_rootid
+ _m.body_subtreemass = body_subtreemass
+ _m.body_tree = body_tree
+ _m.body_weldid = body_weldid
+ _m.cam_bodyid = cam_bodyid
+ _m.cam_fovy = cam_fovy
+ _m.cam_intrinsic = cam_intrinsic
+ _m.cam_mat0 = cam_mat0
+ _m.cam_mode = cam_mode
+ _m.cam_pos = cam_pos
+ _m.cam_pos0 = cam_pos0
+ _m.cam_poscom0 = cam_poscom0
+ _m.cam_quat = cam_quat
+ _m.cam_resolution = cam_resolution
+ _m.cam_sensorsize = cam_sensorsize
+ _m.cam_targetbodyid = cam_targetbodyid
+ _m.condim_max = condim_max
+ _m.dof_Madr = dof_Madr
+ _m.dof_armature = dof_armature
+ _m.dof_bodyid = dof_bodyid
+ _m.dof_damping = dof_damping
+ _m.dof_frictionloss = dof_frictionloss
+ _m.dof_invweight0 = dof_invweight0
+ _m.dof_jntid = dof_jntid
+ _m.dof_parentid = dof_parentid
+ _m.dof_solimp = dof_solimp
+ _m.dof_solref = dof_solref
+ _m.dof_tri_col = dof_tri_col
+ _m.dof_tri_row = dof_tri_row
+ _m.eq_connect_adr = eq_connect_adr
+ _m.eq_data = eq_data
+ _m.eq_jnt_adr = eq_jnt_adr
+ _m.eq_obj1id = eq_obj1id
+ _m.eq_obj2id = eq_obj2id
+ _m.eq_objtype = eq_objtype
+ _m.eq_solimp = eq_solimp
+ _m.eq_solref = eq_solref
+ _m.eq_ten_adr = eq_ten_adr
+ _m.eq_wld_adr = eq_wld_adr
+ _m.flex_bending = flex_bending
+ _m.flex_damping = flex_damping
+ _m.flex_dim = flex_dim
+ _m.flex_edge = flex_edge
+ _m.flex_edgeadr = flex_edgeadr
+ _m.flex_edgeflap = flex_edgeflap
+ _m.flex_elem = flex_elem
+ _m.flex_elemedge = flex_elemedge
+ _m.flex_elemedgeadr = flex_elemedgeadr
+ _m.flex_stiffness = flex_stiffness
+ _m.flex_vertadr = flex_vertadr
+ _m.flex_vertbodyid = flex_vertbodyid
+ _m.flexedge_length0 = flexedge_length0
+ _m.geom_aabb = geom_aabb
+ _m.geom_bodyid = geom_bodyid
+ _m.geom_condim = geom_condim
+ _m.geom_dataid = geom_dataid
+ _m.geom_friction = geom_friction
+ _m.geom_gap = geom_gap
+ _m.geom_group = geom_group
+ _m.geom_margin = geom_margin
+ _m.geom_matid = geom_matid
+ _m.geom_pair_type_count = geom_pair_type_count
+ _m.geom_plugin_index = geom_plugin_index
+ _m.geom_pos = geom_pos
+ _m.geom_priority = geom_priority
+ _m.geom_quat = geom_quat
+ _m.geom_rbound = geom_rbound
+ _m.geom_rgba = geom_rgba
+ _m.geom_size = geom_size
+ _m.geom_solimp = geom_solimp
+ _m.geom_solmix = geom_solmix
+ _m.geom_solref = geom_solref
+ _m.geom_type = geom_type
+ _m.geompair2hfgeompair = geompair2hfgeompair
+ _m.has_sdf_geom = has_sdf_geom
+ _m.hfield_adr = hfield_adr
+ _m.hfield_data = hfield_data
+ _m.hfield_ncol = hfield_ncol
+ _m.hfield_nrow = hfield_nrow
+ _m.hfield_size = hfield_size
+ _m.jnt_actfrclimited = jnt_actfrclimited
+ _m.jnt_actfrcrange = jnt_actfrcrange
+ _m.jnt_actgravcomp = jnt_actgravcomp
+ _m.jnt_axis = jnt_axis
+ _m.jnt_bodyid = jnt_bodyid
+ _m.jnt_dofadr = jnt_dofadr
+ _m.jnt_limited_ball_adr = jnt_limited_ball_adr
+ _m.jnt_limited_slide_hinge_adr = jnt_limited_slide_hinge_adr
+ _m.jnt_margin = jnt_margin
+ _m.jnt_pos = jnt_pos
+ _m.jnt_qposadr = jnt_qposadr
+ _m.jnt_range = jnt_range
+ _m.jnt_solimp = jnt_solimp
+ _m.jnt_solref = jnt_solref
+ _m.jnt_stiffness = jnt_stiffness
+ _m.jnt_type = jnt_type
+ _m.light_bodyid = light_bodyid
+ _m.light_dir = light_dir
+ _m.light_dir0 = light_dir0
+ _m.light_mode = light_mode
+ _m.light_pos = light_pos
+ _m.light_pos0 = light_pos0
+ _m.light_poscom0 = light_poscom0
+ _m.light_targetbodyid = light_targetbodyid
+ _m.mapM2M = mapM2M
+ _m.mat_rgba = mat_rgba
+ _m.mesh_face = mesh_face
+ _m.mesh_faceadr = mesh_faceadr
+ _m.mesh_graph = mesh_graph
+ _m.mesh_graphadr = mesh_graphadr
+ _m.mesh_polyadr = mesh_polyadr
+ _m.mesh_polymap = mesh_polymap
+ _m.mesh_polymapadr = mesh_polymapadr
+ _m.mesh_polymapnum = mesh_polymapnum
+ _m.mesh_polynormal = mesh_polynormal
+ _m.mesh_polynum = mesh_polynum
+ _m.mesh_polyvert = mesh_polyvert
+ _m.mesh_polyvertadr = mesh_polyvertadr
+ _m.mesh_polyvertnum = mesh_polyvertnum
+ _m.mesh_vert = mesh_vert
+ _m.mesh_vertadr = mesh_vertadr
+ _m.mesh_vertnum = mesh_vertnum
+ _m.mocap_bodyid = mocap_bodyid
+ _m.nC = nC
+ _m.na = na
+ _m.nbody = nbody
+ _m.ncam = ncam
+ _m.neq = neq
+ _m.nflexedge = nflexedge
+ _m.nflexelem = nflexelem
+ _m.nflexvert = nflexvert
+ _m.ngeom = ngeom
+ _m.ngravcomp = ngravcomp
+ _m.nhfield = nhfield
+ _m.njnt = njnt
+ _m.nlight = nlight
+ _m.nlsp = nlsp
+ _m.nmeshface = nmeshface
+ _m.nmocap = nmocap
+ _m.nsite = nsite
+ _m.ntendon = ntendon
+ _m.nu = nu
+ _m.nv = nv
+ _m.nxn_geom_pair_filtered = nxn_geom_pair_filtered
+ _m.nxn_pairid = nxn_pairid
+ _m.nxn_pairid_filtered = nxn_pairid_filtered
+ _m.opt.broadphase = opt__broadphase
+ _m.opt.broadphase_filter = opt__broadphase_filter
+ _m.opt.cone = opt__cone
+ _m.opt.density = opt__density
+ _m.opt.disableflags = opt__disableflags
+ _m.opt.enableflags = opt__enableflags
+ _m.opt.epa_iterations = opt__epa_iterations
+ _m.opt.gjk_iterations = opt__gjk_iterations
+ _m.opt.graph_conditional = opt__graph_conditional
+ _m.opt.gravity = opt__gravity
+ _m.opt.has_fluid = opt__has_fluid
+ _m.opt.impratio = opt__impratio
+ _m.opt.integrator = opt__integrator
+ _m.opt.is_sparse = opt__is_sparse
+ _m.opt.iterations = opt__iterations
+ _m.opt.ls_iterations = opt__ls_iterations
+ _m.opt.ls_parallel = opt__ls_parallel
+ _m.opt.ls_tolerance = opt__ls_tolerance
+ _m.opt.magnetic = opt__magnetic
+ _m.opt.run_collision_detection = opt__run_collision_detection
+ _m.opt.sdf_initpoints = opt__sdf_initpoints
+ _m.opt.sdf_iterations = opt__sdf_iterations
+ _m.opt.solver = opt__solver
+ _m.opt.timestep = opt__timestep
+ _m.opt.tolerance = opt__tolerance
+ _m.opt.viscosity = opt__viscosity
+ _m.opt.wind = opt__wind
+ _m.pair_dim = pair_dim
+ _m.pair_friction = pair_friction
+ _m.pair_gap = pair_gap
+ _m.pair_margin = pair_margin
+ _m.pair_solimp = pair_solimp
+ _m.pair_solref = pair_solref
+ _m.pair_solreffriction = pair_solreffriction
+ _m.plugin = plugin
+ _m.plugin_attr = plugin_attr
+ _m.qLD_updates = qLD_updates
+ _m.qM_fullm_i = qM_fullm_i
+ _m.qM_fullm_j = qM_fullm_j
+ _m.qM_madr_ij = qM_madr_ij
+ _m.qM_mulm_i = qM_mulm_i
+ _m.qM_mulm_j = qM_mulm_j
+ _m.qM_tiles = qM_tiles
+ _m.qpos0 = qpos0
+ _m.qpos_spring = qpos_spring
+ _m.rangefinder_sensor_adr = rangefinder_sensor_adr
+ _m.sensor_acc_adr = sensor_acc_adr
+ _m.sensor_adr = sensor_adr
+ _m.sensor_cutoff = sensor_cutoff
+ _m.sensor_datatype = sensor_datatype
+ _m.sensor_e_kinetic = sensor_e_kinetic
+ _m.sensor_e_potential = sensor_e_potential
+ _m.sensor_limitfrc_adr = sensor_limitfrc_adr
+ _m.sensor_limitpos_adr = sensor_limitpos_adr
+ _m.sensor_limitvel_adr = sensor_limitvel_adr
+ _m.sensor_objid = sensor_objid
+ _m.sensor_objtype = sensor_objtype
+ _m.sensor_pos_adr = sensor_pos_adr
+ _m.sensor_rangefinder_adr = sensor_rangefinder_adr
+ _m.sensor_rangefinder_bodyid = sensor_rangefinder_bodyid
+ _m.sensor_refid = sensor_refid
+ _m.sensor_reftype = sensor_reftype
+ _m.sensor_rne_postconstraint = sensor_rne_postconstraint
+ _m.sensor_subtree_vel = sensor_subtree_vel
+ _m.sensor_tendonactfrc_adr = sensor_tendonactfrc_adr
+ _m.sensor_touch_adr = sensor_touch_adr
+ _m.sensor_type = sensor_type
+ _m.sensor_vel_adr = sensor_vel_adr
+ _m.site_bodyid = site_bodyid
+ _m.site_pos = site_pos
+ _m.site_quat = site_quat
+ _m.site_size = site_size
+ _m.site_type = site_type
+ _m.stat.meaninertia = stat__meaninertia
+ _m.subtree_mass = subtree_mass
+ _m.tendon_actfrclimited = tendon_actfrclimited
+ _m.tendon_actfrcrange = tendon_actfrcrange
+ _m.tendon_adr = tendon_adr
+ _m.tendon_armature = tendon_armature
+ _m.tendon_damping = tendon_damping
+ _m.tendon_frictionloss = tendon_frictionloss
+ _m.tendon_geom_adr = tendon_geom_adr
+ _m.tendon_invweight0 = tendon_invweight0
+ _m.tendon_jnt_adr = tendon_jnt_adr
+ _m.tendon_length0 = tendon_length0
+ _m.tendon_lengthspring = tendon_lengthspring
+ _m.tendon_limited_adr = tendon_limited_adr
+ _m.tendon_margin = tendon_margin
+ _m.tendon_num = tendon_num
+ _m.tendon_range = tendon_range
+ _m.tendon_site_pair_adr = tendon_site_pair_adr
+ _m.tendon_solimp_fri = tendon_solimp_fri
+ _m.tendon_solimp_lim = tendon_solimp_lim
+ _m.tendon_solref_fri = tendon_solref_fri
+ _m.tendon_solref_lim = tendon_solref_lim
+ _m.tendon_stiffness = tendon_stiffness
+ _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.act = act
+ _d.act_dot = act_dot
+ _d.act_dot_rk = act_dot_rk
+ _d.act_t0 = act_t0
+ _d.actuator_force = actuator_force
+ _d.actuator_length = actuator_length
+ _d.actuator_moment = actuator_moment
+ _d.actuator_trntype_body_ncon = actuator_trntype_body_ncon
+ _d.actuator_velocity = actuator_velocity
+ _d.cacc = cacc
+ _d.cam_xmat = cam_xmat
+ _d.cam_xpos = cam_xpos
+ _d.cdof = cdof
+ _d.cdof_dot = cdof_dot
+ _d.cfrc_ext = cfrc_ext
+ _d.cfrc_int = cfrc_int
+ _d.cinert = cinert
+ _d.collision_hftri_index = collision_hftri_index
+ _d.collision_pair = collision_pair
+ _d.collision_pairid = collision_pairid
+ _d.collision_worldid = collision_worldid
+ _d.contact.dim = contact__dim
+ _d.contact.dist = contact__dist
+ _d.contact.efc_address = contact__efc_address
+ _d.contact.frame = contact__frame
+ _d.contact.friction = contact__friction
+ _d.contact.geom = contact__geom
+ _d.contact.includemargin = contact__includemargin
+ _d.contact.pos = contact__pos
+ _d.contact.solimp = contact__solimp
+ _d.contact.solref = contact__solref
+ _d.contact.solreffriction = contact__solreffriction
+ _d.contact.worldid = contact__worldid
+ _d.crb = crb
+ _d.ctrl = ctrl
+ _d.cvel = cvel
+ _d.efc.D = efc__D
+ _d.efc.J = efc__J
+ _d.efc.Jaref = efc__Jaref
+ _d.efc.Ma = efc__Ma
+ _d.efc.Mgrad = efc__Mgrad
+ _d.efc.active = efc__active
+ _d.efc.alpha = efc__alpha
+ _d.efc.aref = efc__aref
+ _d.efc.beta = efc__beta
+ _d.efc.beta_den = efc__beta_den
+ _d.efc.beta_num = efc__beta_num
+ _d.efc.cholesky_L_tmp = efc__cholesky_L_tmp
+ _d.efc.cholesky_y_tmp = efc__cholesky_y_tmp
+ _d.efc.condim = efc__condim
+ _d.efc.cost = efc__cost
+ _d.efc.cost_candidate = efc__cost_candidate
+ _d.efc.done = efc__done
+ _d.efc.force = efc__force
+ _d.efc.frictionloss = efc__frictionloss
+ _d.efc.gauss = efc__gauss
+ _d.efc.grad = efc__grad
+ _d.efc.grad_dot = efc__grad_dot
+ _d.efc.gtol = efc__gtol
+ _d.efc.h = efc__h
+ _d.efc.hi = efc__hi
+ _d.efc.hi_alpha = efc__hi_alpha
+ _d.efc.hi_next = efc__hi_next
+ _d.efc.hi_next_alpha = efc__hi_next_alpha
+ _d.efc.id = efc__id
+ _d.efc.jv = efc__jv
+ _d.efc.lo = efc__lo
+ _d.efc.lo_alpha = efc__lo_alpha
+ _d.efc.lo_next = efc__lo_next
+ _d.efc.lo_next_alpha = efc__lo_next_alpha
+ _d.efc.ls_done = efc__ls_done
+ _d.efc.margin = efc__margin
+ _d.efc.mid = efc__mid
+ _d.efc.mid_alpha = efc__mid_alpha
+ _d.efc.mv = efc__mv
+ _d.efc.p0 = efc__p0
+ _d.efc.pos = efc__pos
+ _d.efc.prev_Mgrad = efc__prev_Mgrad
+ _d.efc.prev_cost = efc__prev_cost
+ _d.efc.prev_grad = efc__prev_grad
+ _d.efc.quad = efc__quad
+ _d.efc.quad_gauss = efc__quad_gauss
+ _d.efc.search = efc__search
+ _d.efc.search_dot = efc__search_dot
+ _d.efc.type = efc__type
+ _d.efc.u = efc__u
+ _d.efc.uu = efc__uu
+ _d.efc.uv = efc__uv
+ _d.efc.vel = efc__vel
+ _d.efc.vv = efc__vv
+ _d.energy = energy
+ _d.epa_face = epa_face
+ _d.epa_horizon = epa_horizon
+ _d.epa_index = epa_index
+ _d.epa_map = epa_map
+ _d.epa_norm2 = epa_norm2
+ _d.epa_pr = epa_pr
+ _d.epa_vert = epa_vert
+ _d.epa_vert1 = epa_vert1
+ _d.epa_vert2 = epa_vert2
+ _d.epa_vert_index1 = epa_vert_index1
+ _d.epa_vert_index2 = epa_vert_index2
+ _d.eq_active = eq_active
+ _d.flexedge_length = flexedge_length
+ _d.flexedge_velocity = flexedge_velocity
+ _d.flexvert_xpos = flexvert_xpos
+ _d.fluid_applied = fluid_applied
+ _d.geom_skip = geom_skip
+ _d.geom_xmat = geom_xmat
+ _d.geom_xpos = geom_xpos
+ _d.inverse_mul_m_skip = inverse_mul_m_skip
+ _d.light_xdir = light_xdir
+ _d.light_xpos = light_xpos
+ _d.mocap_pos = mocap_pos
+ _d.mocap_quat = mocap_quat
+ _d.ncollision = ncollision
+ _d.ncon = ncon
+ _d.ncon_hfield = ncon_hfield
+ _d.nconmax = nconmax
+ _d.ne = ne
+ _d.ne_connect = ne_connect
+ _d.ne_jnt = ne_jnt
+ _d.ne_ten = ne_ten
+ _d.ne_weld = ne_weld
+ _d.nefc = nefc
+ _d.nf = nf
+ _d.njmax = njmax
+ _d.nl = nl
+ _d.nsolving = nsolving
+ _d.qLD = qLD
+ _d.qLD_integration = qLD_integration
+ _d.qLDiagInv = qLDiagInv
+ _d.qLDiagInv_integration = qLDiagInv_integration
+ _d.qM = qM
+ _d.qM_integration = qM_integration
+ _d.qacc = qacc
+ _d.qacc_integration = qacc_integration
+ _d.qacc_rk = qacc_rk
+ _d.qacc_smooth = qacc_smooth
+ _d.qacc_warmstart = qacc_warmstart
+ _d.qfrc_actuator = qfrc_actuator
+ _d.qfrc_applied = qfrc_applied
+ _d.qfrc_bias = qfrc_bias
+ _d.qfrc_constraint = qfrc_constraint
+ _d.qfrc_damper = qfrc_damper
+ _d.qfrc_fluid = qfrc_fluid
+ _d.qfrc_gravcomp = qfrc_gravcomp
+ _d.qfrc_integration = qfrc_integration
+ _d.qfrc_passive = qfrc_passive
+ _d.qfrc_smooth = qfrc_smooth
+ _d.qfrc_spring = qfrc_spring
+ _d.qpos = qpos
+ _d.qpos_t0 = qpos_t0
+ _d.qvel = qvel
+ _d.qvel_rk = qvel_rk
+ _d.qvel_t0 = qvel_t0
+ _d.sap_cumulative_sum = sap_cumulative_sum
+ _d.sap_projection_lower = sap_projection_lower
+ _d.sap_projection_upper = sap_projection_upper
+ _d.sap_range = sap_range
+ _d.sap_segment_index = sap_segment_index
+ _d.sap_sort_index = sap_sort_index
+ _d.sensor_rangefinder_dist = sensor_rangefinder_dist
+ _d.sensor_rangefinder_geomid = sensor_rangefinder_geomid
+ _d.sensor_rangefinder_pnt = sensor_rangefinder_pnt
+ _d.sensor_rangefinder_vec = sensor_rangefinder_vec
+ _d.sensordata = sensordata
+ _d.site_xmat = site_xmat
+ _d.site_xpos = site_xpos
+ _d.solver_niter = solver_niter
+ _d.subtree_angmom = subtree_angmom
+ _d.subtree_bodyvel = subtree_bodyvel
+ _d.subtree_com = subtree_com
+ _d.subtree_linvel = subtree_linvel
+ _d.ten_J = ten_J
+ _d.ten_Jdot = ten_Jdot
+ _d.ten_actfrc = ten_actfrc
+ _d.ten_bias_coef = ten_bias_coef
+ _d.ten_length = ten_length
+ _d.ten_velocity = ten_velocity
+ _d.ten_wrapadr = ten_wrapadr
+ _d.ten_wrapnum = ten_wrapnum
+ _d.time = time
+ _d.wrap_geom_xpos = wrap_geom_xpos
+ _d.wrap_obj = wrap_obj
+ _d.wrap_xpos = wrap_xpos
+ _d.xanchor = xanchor
+ _d.xaxis = xaxis
+ _d.xfrc_applied = xfrc_applied
+ _d.ximat = ximat
+ _d.xipos = xipos
+ _d.xmat = xmat
+ _d.xpos = xpos
+ _d.xquat = xquat
+ _d.nworld = nworld
+ mjwarp.step(_m, _d)
+
+
+def _step_jax_impl(m: types.Model, d: types.Data):
+ output_dims = {
+ 'act': d.act.shape,
+ 'act_dot': d.act_dot.shape,
+ 'act_dot_rk': d._impl.act_dot_rk.shape,
+ 'act_t0': d._impl.act_t0.shape,
+ 'actuator_force': d.actuator_force.shape,
+ 'actuator_length': d._impl.actuator_length.shape,
+ 'actuator_moment': d._impl.actuator_moment.shape,
+ 'actuator_trntype_body_ncon': d._impl.actuator_trntype_body_ncon.shape,
+ 'actuator_velocity': d._impl.actuator_velocity.shape,
+ 'cacc': d._impl.cacc.shape,
+ 'cam_xmat': d.cam_xmat.shape,
+ 'cam_xpos': d.cam_xpos.shape,
+ 'cdof': d._impl.cdof.shape,
+ 'cdof_dot': d._impl.cdof_dot.shape,
+ 'cfrc_ext': d._impl.cfrc_ext.shape,
+ 'cfrc_int': d._impl.cfrc_int.shape,
+ 'cinert': d._impl.cinert.shape,
+ 'collision_hftri_index': d._impl.collision_hftri_index.shape,
+ 'collision_pair': d._impl.collision_pair.shape,
+ 'collision_pairid': d._impl.collision_pairid.shape,
+ 'collision_worldid': d._impl.collision_worldid.shape,
+ 'crb': d._impl.crb.shape,
+ 'ctrl': d.ctrl.shape,
+ 'cvel': d.cvel.shape,
+ 'energy': d._impl.energy.shape,
+ 'epa_face': d._impl.epa_face.shape,
+ 'epa_horizon': d._impl.epa_horizon.shape,
+ 'epa_index': d._impl.epa_index.shape,
+ 'epa_map': d._impl.epa_map.shape,
+ 'epa_norm2': d._impl.epa_norm2.shape,
+ 'epa_pr': d._impl.epa_pr.shape,
+ 'epa_vert': d._impl.epa_vert.shape,
+ 'epa_vert1': d._impl.epa_vert1.shape,
+ 'epa_vert2': d._impl.epa_vert2.shape,
+ 'epa_vert_index1': d._impl.epa_vert_index1.shape,
+ 'epa_vert_index2': d._impl.epa_vert_index2.shape,
+ 'eq_active': d.eq_active.shape,
+ 'flexedge_length': d._impl.flexedge_length.shape,
+ 'flexedge_velocity': d._impl.flexedge_velocity.shape,
+ 'flexvert_xpos': d._impl.flexvert_xpos.shape,
+ 'fluid_applied': d._impl.fluid_applied.shape,
+ 'geom_skip': d._impl.geom_skip.shape,
+ 'geom_xmat': d.geom_xmat.shape,
+ 'geom_xpos': d.geom_xpos.shape,
+ 'inverse_mul_m_skip': d._impl.inverse_mul_m_skip.shape,
+ 'light_xdir': d._impl.light_xdir.shape,
+ 'light_xpos': d._impl.light_xpos.shape,
+ 'mocap_pos': d.mocap_pos.shape,
+ 'mocap_quat': d.mocap_quat.shape,
+ 'ncollision': d._impl.ncollision.shape,
+ 'ncon': d._impl.ncon.shape,
+ 'ncon_hfield': d._impl.ncon_hfield.shape,
+ 'ne': d._impl.ne.shape,
+ 'ne_connect': d._impl.ne_connect.shape,
+ 'ne_jnt': d._impl.ne_jnt.shape,
+ 'ne_ten': d._impl.ne_ten.shape,
+ 'ne_weld': d._impl.ne_weld.shape,
+ 'nefc': d._impl.nefc.shape,
+ 'nf': d._impl.nf.shape,
+ 'nl': d._impl.nl.shape,
+ 'nsolving': d._impl.nsolving.shape,
+ 'qLD': d._impl.qLD.shape,
+ 'qLD_integration': d._impl.qLD_integration.shape,
+ 'qLDiagInv': d._impl.qLDiagInv.shape,
+ 'qLDiagInv_integration': d._impl.qLDiagInv_integration.shape,
+ 'qM': d._impl.qM.shape,
+ 'qM_integration': d._impl.qM_integration.shape,
+ 'qacc': d.qacc.shape,
+ 'qacc_integration': d._impl.qacc_integration.shape,
+ 'qacc_rk': d._impl.qacc_rk.shape,
+ 'qacc_smooth': d.qacc_smooth.shape,
+ 'qacc_warmstart': d.qacc_warmstart.shape,
+ 'qfrc_actuator': d.qfrc_actuator.shape,
+ 'qfrc_applied': d.qfrc_applied.shape,
+ 'qfrc_bias': d.qfrc_bias.shape,
+ 'qfrc_constraint': d.qfrc_constraint.shape,
+ 'qfrc_damper': d._impl.qfrc_damper.shape,
+ 'qfrc_fluid': d.qfrc_fluid.shape,
+ 'qfrc_gravcomp': d.qfrc_gravcomp.shape,
+ 'qfrc_integration': d._impl.qfrc_integration.shape,
+ 'qfrc_passive': d.qfrc_passive.shape,
+ 'qfrc_smooth': d.qfrc_smooth.shape,
+ 'qfrc_spring': d._impl.qfrc_spring.shape,
+ 'qpos': d.qpos.shape,
+ 'qpos_t0': d._impl.qpos_t0.shape,
+ 'qvel': d.qvel.shape,
+ 'qvel_rk': d._impl.qvel_rk.shape,
+ 'qvel_t0': d._impl.qvel_t0.shape,
+ 'sap_cumulative_sum': d._impl.sap_cumulative_sum.shape,
+ 'sap_projection_lower': d._impl.sap_projection_lower.shape,
+ 'sap_projection_upper': d._impl.sap_projection_upper.shape,
+ 'sap_range': d._impl.sap_range.shape,
+ 'sap_segment_index': d._impl.sap_segment_index.shape,
+ 'sap_sort_index': d._impl.sap_sort_index.shape,
+ 'sensor_rangefinder_dist': d._impl.sensor_rangefinder_dist.shape,
+ 'sensor_rangefinder_geomid': d._impl.sensor_rangefinder_geomid.shape,
+ 'sensor_rangefinder_pnt': d._impl.sensor_rangefinder_pnt.shape,
+ 'sensor_rangefinder_vec': d._impl.sensor_rangefinder_vec.shape,
+ 'sensordata': d.sensordata.shape,
+ 'site_xmat': d.site_xmat.shape,
+ 'site_xpos': d.site_xpos.shape,
+ 'solver_niter': d._impl.solver_niter.shape,
+ 'subtree_angmom': d._impl.subtree_angmom.shape,
+ 'subtree_bodyvel': d._impl.subtree_bodyvel.shape,
+ 'subtree_com': d.subtree_com.shape,
+ 'subtree_linvel': d._impl.subtree_linvel.shape,
+ 'ten_J': d._impl.ten_J.shape,
+ '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._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,
+ 'time': d.time.shape,
+ 'wrap_geom_xpos': d._impl.wrap_geom_xpos.shape,
+ 'wrap_obj': d._impl.wrap_obj.shape,
+ 'wrap_xpos': d._impl.wrap_xpos.shape,
+ 'xanchor': d.xanchor.shape,
+ 'xaxis': d.xaxis.shape,
+ 'xfrc_applied': d.xfrc_applied.shape,
+ 'ximat': d.ximat.shape,
+ 'xipos': d.xipos.shape,
+ 'xmat': d.xmat.shape,
+ 'xpos': d.xpos.shape,
+ 'xquat': d.xquat.shape,
+ 'contact__dim': d._impl.contact__dim.shape,
+ 'contact__dist': d._impl.contact__dist.shape,
+ 'contact__efc_address': d._impl.contact__efc_address.shape,
+ 'contact__frame': d._impl.contact__frame.shape,
+ 'contact__friction': d._impl.contact__friction.shape,
+ 'contact__geom': d._impl.contact__geom.shape,
+ 'contact__includemargin': d._impl.contact__includemargin.shape,
+ 'contact__pos': d._impl.contact__pos.shape,
+ 'contact__solimp': d._impl.contact__solimp.shape,
+ 'contact__solref': d._impl.contact__solref.shape,
+ 'contact__solreffriction': d._impl.contact__solreffriction.shape,
+ 'contact__worldid': d._impl.contact__worldid.shape,
+ 'efc__D': d._impl.efc__D.shape,
+ 'efc__J': d._impl.efc__J.shape,
+ 'efc__Jaref': d._impl.efc__Jaref.shape,
+ 'efc__Ma': d._impl.efc__Ma.shape,
+ 'efc__Mgrad': d._impl.efc__Mgrad.shape,
+ 'efc__active': d._impl.efc__active.shape,
+ 'efc__alpha': d._impl.efc__alpha.shape,
+ 'efc__aref': d._impl.efc__aref.shape,
+ 'efc__beta': d._impl.efc__beta.shape,
+ 'efc__beta_den': d._impl.efc__beta_den.shape,
+ 'efc__beta_num': d._impl.efc__beta_num.shape,
+ 'efc__cholesky_L_tmp': d._impl.efc__cholesky_L_tmp.shape,
+ 'efc__cholesky_y_tmp': d._impl.efc__cholesky_y_tmp.shape,
+ 'efc__condim': d._impl.efc__condim.shape,
+ 'efc__cost': d._impl.efc__cost.shape,
+ 'efc__cost_candidate': d._impl.efc__cost_candidate.shape,
+ 'efc__done': d._impl.efc__done.shape,
+ 'efc__force': d._impl.efc__force.shape,
+ 'efc__frictionloss': d._impl.efc__frictionloss.shape,
+ 'efc__gauss': d._impl.efc__gauss.shape,
+ 'efc__grad': d._impl.efc__grad.shape,
+ 'efc__grad_dot': d._impl.efc__grad_dot.shape,
+ 'efc__gtol': d._impl.efc__gtol.shape,
+ 'efc__h': d._impl.efc__h.shape,
+ 'efc__hi': d._impl.efc__hi.shape,
+ 'efc__hi_alpha': d._impl.efc__hi_alpha.shape,
+ 'efc__hi_next': d._impl.efc__hi_next.shape,
+ 'efc__hi_next_alpha': d._impl.efc__hi_next_alpha.shape,
+ 'efc__id': d._impl.efc__id.shape,
+ 'efc__jv': d._impl.efc__jv.shape,
+ 'efc__lo': d._impl.efc__lo.shape,
+ 'efc__lo_alpha': d._impl.efc__lo_alpha.shape,
+ 'efc__lo_next': d._impl.efc__lo_next.shape,
+ 'efc__lo_next_alpha': d._impl.efc__lo_next_alpha.shape,
+ 'efc__ls_done': d._impl.efc__ls_done.shape,
+ 'efc__margin': d._impl.efc__margin.shape,
+ 'efc__mid': d._impl.efc__mid.shape,
+ 'efc__mid_alpha': d._impl.efc__mid_alpha.shape,
+ 'efc__mv': d._impl.efc__mv.shape,
+ 'efc__p0': d._impl.efc__p0.shape,
+ 'efc__pos': d._impl.efc__pos.shape,
+ 'efc__prev_Mgrad': d._impl.efc__prev_Mgrad.shape,
+ 'efc__prev_cost': d._impl.efc__prev_cost.shape,
+ 'efc__prev_grad': d._impl.efc__prev_grad.shape,
+ 'efc__quad': d._impl.efc__quad.shape,
+ 'efc__quad_gauss': d._impl.efc__quad_gauss.shape,
+ 'efc__search': d._impl.efc__search.shape,
+ 'efc__search_dot': d._impl.efc__search_dot.shape,
+ 'efc__type': d._impl.efc__type.shape,
+ 'efc__u': d._impl.efc__u.shape,
+ 'efc__uu': d._impl.efc__uu.shape,
+ 'efc__uv': d._impl.efc__uv.shape,
+ 'efc__vel': d._impl.efc__vel.shape,
+ 'efc__vv': d._impl.efc__vv.shape,
+ }
+ jf = ffi.jax_callable_variadic_tuple(
+ _step_shim,
+ num_outputs=192,
+ output_dims=output_dims,
+ vmap_method=None,
+ graph_compatible=True,
+ in_out_argnames={
+ 'act',
+ 'act_dot',
+ 'act_dot_rk',
+ 'act_t0',
+ 'actuator_force',
+ 'actuator_length',
+ 'actuator_moment',
+ 'actuator_trntype_body_ncon',
+ 'actuator_velocity',
+ 'cacc',
+ 'cam_xmat',
+ 'cam_xpos',
+ 'cdof',
+ 'cdof_dot',
+ 'cfrc_ext',
+ 'cfrc_int',
+ 'cinert',
+ 'collision_hftri_index',
+ 'collision_pair',
+ 'collision_pairid',
+ 'collision_worldid',
+ 'crb',
+ 'ctrl',
+ 'cvel',
+ 'energy',
+ 'epa_face',
+ 'epa_horizon',
+ 'epa_index',
+ 'epa_map',
+ 'epa_norm2',
+ 'epa_pr',
+ 'epa_vert',
+ 'epa_vert1',
+ 'epa_vert2',
+ 'epa_vert_index1',
+ 'epa_vert_index2',
+ 'eq_active',
+ 'flexedge_length',
+ 'flexedge_velocity',
+ 'flexvert_xpos',
+ 'fluid_applied',
+ 'geom_skip',
+ 'geom_xmat',
+ 'geom_xpos',
+ 'inverse_mul_m_skip',
+ 'light_xdir',
+ 'light_xpos',
+ 'mocap_pos',
+ 'mocap_quat',
+ 'ncollision',
+ 'ncon',
+ 'ncon_hfield',
+ 'ne',
+ 'ne_connect',
+ 'ne_jnt',
+ 'ne_ten',
+ 'ne_weld',
+ 'nefc',
+ 'nf',
+ 'nl',
+ 'nsolving',
+ 'qLD',
+ 'qLD_integration',
+ 'qLDiagInv',
+ 'qLDiagInv_integration',
+ 'qM',
+ 'qM_integration',
+ 'qacc',
+ 'qacc_integration',
+ 'qacc_rk',
+ 'qacc_smooth',
+ 'qacc_warmstart',
+ 'qfrc_actuator',
+ 'qfrc_applied',
+ 'qfrc_bias',
+ 'qfrc_constraint',
+ 'qfrc_damper',
+ 'qfrc_fluid',
+ 'qfrc_gravcomp',
+ 'qfrc_integration',
+ 'qfrc_passive',
+ 'qfrc_smooth',
+ 'qfrc_spring',
+ 'qpos',
+ 'qpos_t0',
+ 'qvel',
+ 'qvel_rk',
+ 'qvel_t0',
+ 'sap_cumulative_sum',
+ 'sap_projection_lower',
+ 'sap_projection_upper',
+ 'sap_range',
+ 'sap_segment_index',
+ 'sap_sort_index',
+ 'sensor_rangefinder_dist',
+ 'sensor_rangefinder_geomid',
+ 'sensor_rangefinder_pnt',
+ 'sensor_rangefinder_vec',
+ 'sensordata',
+ 'site_xmat',
+ 'site_xpos',
+ 'solver_niter',
+ 'subtree_angmom',
+ 'subtree_bodyvel',
+ 'subtree_com',
+ 'subtree_linvel',
+ 'ten_J',
+ 'ten_Jdot',
+ 'ten_actfrc',
+ 'ten_bias_coef',
+ 'ten_length',
+ 'ten_velocity',
+ 'ten_wrapadr',
+ 'ten_wrapnum',
+ 'time',
+ 'wrap_geom_xpos',
+ 'wrap_obj',
+ 'wrap_xpos',
+ 'xanchor',
+ 'xaxis',
+ 'xfrc_applied',
+ 'ximat',
+ 'xipos',
+ 'xmat',
+ 'xpos',
+ 'xquat',
+ 'contact__dim',
+ 'contact__dist',
+ 'contact__efc_address',
+ 'contact__frame',
+ 'contact__friction',
+ 'contact__geom',
+ 'contact__includemargin',
+ 'contact__pos',
+ 'contact__solimp',
+ 'contact__solref',
+ 'contact__solreffriction',
+ 'contact__worldid',
+ 'efc__D',
+ 'efc__J',
+ 'efc__Jaref',
+ 'efc__Ma',
+ 'efc__Mgrad',
+ 'efc__active',
+ 'efc__alpha',
+ 'efc__aref',
+ 'efc__beta',
+ 'efc__beta_den',
+ 'efc__beta_num',
+ 'efc__cholesky_L_tmp',
+ 'efc__cholesky_y_tmp',
+ 'efc__condim',
+ 'efc__cost',
+ 'efc__cost_candidate',
+ 'efc__done',
+ 'efc__force',
+ 'efc__frictionloss',
+ 'efc__gauss',
+ 'efc__grad',
+ 'efc__grad_dot',
+ 'efc__gtol',
+ 'efc__h',
+ 'efc__hi',
+ 'efc__hi_alpha',
+ 'efc__hi_next',
+ 'efc__hi_next_alpha',
+ 'efc__id',
+ 'efc__jv',
+ 'efc__lo',
+ 'efc__lo_alpha',
+ 'efc__lo_next',
+ 'efc__lo_next_alpha',
+ 'efc__ls_done',
+ 'efc__margin',
+ 'efc__mid',
+ 'efc__mid_alpha',
+ 'efc__mv',
+ 'efc__p0',
+ 'efc__pos',
+ 'efc__prev_Mgrad',
+ 'efc__prev_cost',
+ 'efc__prev_grad',
+ 'efc__quad',
+ 'efc__quad_gauss',
+ 'efc__search',
+ 'efc__search_dot',
+ 'efc__type',
+ 'efc__u',
+ 'efc__uu',
+ 'efc__uv',
+ 'efc__vel',
+ 'efc__vv',
+ },
+ )
+ out = jf(
+ d.qpos.shape[0],
+ m._impl.M_rowadr,
+ m._impl.M_rownnz,
+ m.actuator_acc0,
+ m.actuator_actadr,
+ m.actuator_actearly,
+ m.actuator_actlimited,
+ m.actuator_actnum,
+ m.actuator_actrange,
+ m._impl.actuator_affine_bias_gain,
+ m.actuator_biasprm,
+ m.actuator_biastype,
+ m.actuator_cranklength,
+ m.actuator_ctrllimited,
+ m.actuator_ctrlrange,
+ m.actuator_dynprm,
+ m.actuator_dyntype,
+ m.actuator_forcelimited,
+ m.actuator_forcerange,
+ m.actuator_gainprm,
+ m.actuator_gaintype,
+ m.actuator_gear,
+ m.actuator_lengthrange,
+ m.actuator_trnid,
+ m.actuator_trntype,
+ m._impl.actuator_trntype_body_adr,
+ m._impl.block_dim,
+ m.body_dofadr,
+ m.body_dofnum,
+ m.body_gravcomp,
+ m.body_inertia,
+ m.body_invweight0,
+ m.body_ipos,
+ m.body_iquat,
+ m.body_jntadr,
+ m.body_jntnum,
+ m.body_mass,
+ m.body_parentid,
+ m.body_pos,
+ m.body_quat,
+ m.body_rootid,
+ m.body_subtreemass,
+ m._impl.body_tree,
+ m.body_weldid,
+ m.cam_bodyid,
+ m.cam_fovy,
+ m.cam_intrinsic,
+ m.cam_mat0,
+ m.cam_mode,
+ m.cam_pos,
+ m.cam_pos0,
+ m.cam_poscom0,
+ m.cam_quat,
+ m.cam_resolution,
+ m.cam_sensorsize,
+ m.cam_targetbodyid,
+ m._impl.condim_max,
+ m.dof_Madr,
+ m.dof_armature,
+ m.dof_bodyid,
+ m.dof_damping,
+ m.dof_frictionloss,
+ m.dof_invweight0,
+ m.dof_jntid,
+ m.dof_parentid,
+ m.dof_solimp,
+ m.dof_solref,
+ m._impl.dof_tri_col,
+ m._impl.dof_tri_row,
+ m._impl.eq_connect_adr,
+ m.eq_data,
+ m._impl.eq_jnt_adr,
+ m.eq_obj1id,
+ m.eq_obj2id,
+ m.eq_objtype,
+ m.eq_solimp,
+ m.eq_solref,
+ m._impl.eq_ten_adr,
+ m._impl.eq_wld_adr,
+ m._impl.flex_bending,
+ m._impl.flex_damping,
+ m._impl.flex_dim,
+ m._impl.flex_edge,
+ m._impl.flex_edgeadr,
+ m._impl.flex_edgeflap,
+ m._impl.flex_elem,
+ m._impl.flex_elemedge,
+ m._impl.flex_elemedgeadr,
+ m._impl.flex_stiffness,
+ m._impl.flex_vertadr,
+ m._impl.flex_vertbodyid,
+ m._impl.flexedge_length0,
+ m.geom_aabb,
+ m.geom_bodyid,
+ m.geom_condim,
+ m.geom_dataid,
+ m.geom_friction,
+ m.geom_gap,
+ m.geom_group,
+ m.geom_margin,
+ m.geom_matid,
+ m._impl.geom_pair_type_count,
+ m._impl.geom_plugin_index,
+ m.geom_pos,
+ m.geom_priority,
+ m.geom_quat,
+ m.geom_rbound,
+ m.geom_rgba,
+ m.geom_size,
+ m.geom_solimp,
+ m.geom_solmix,
+ m.geom_solref,
+ m.geom_type,
+ m._impl.geompair2hfgeompair,
+ m._impl.has_sdf_geom,
+ m.hfield_adr,
+ m.hfield_data,
+ m.hfield_ncol,
+ m.hfield_nrow,
+ m.hfield_size,
+ m.jnt_actfrclimited,
+ m.jnt_actfrcrange,
+ m.jnt_actgravcomp,
+ m.jnt_axis,
+ m.jnt_bodyid,
+ m.jnt_dofadr,
+ m._impl.jnt_limited_ball_adr,
+ m._impl.jnt_limited_slide_hinge_adr,
+ m.jnt_margin,
+ m.jnt_pos,
+ m.jnt_qposadr,
+ m.jnt_range,
+ m.jnt_solimp,
+ m.jnt_solref,
+ m.jnt_stiffness,
+ m.jnt_type,
+ m._impl.light_bodyid,
+ m.light_dir,
+ m.light_dir0,
+ m.light_mode,
+ m.light_pos,
+ m.light_pos0,
+ m.light_poscom0,
+ m._impl.light_targetbodyid,
+ m._impl.mapM2M,
+ m.mat_rgba,
+ m.mesh_face,
+ m.mesh_faceadr,
+ m.mesh_graph,
+ m.mesh_graphadr,
+ m._impl.mesh_polyadr,
+ m._impl.mesh_polymap,
+ m._impl.mesh_polymapadr,
+ m._impl.mesh_polymapnum,
+ m._impl.mesh_polynormal,
+ m._impl.mesh_polynum,
+ m._impl.mesh_polyvert,
+ m._impl.mesh_polyvertadr,
+ m._impl.mesh_polyvertnum,
+ m.mesh_vert,
+ m.mesh_vertadr,
+ m.mesh_vertnum,
+ m._impl.mocap_bodyid,
+ m.nC,
+ m.na,
+ m.nbody,
+ m.ncam,
+ m.neq,
+ m._impl.nflexedge,
+ m._impl.nflexelem,
+ m._impl.nflexvert,
+ m.ngeom,
+ m.ngravcomp,
+ m.nhfield,
+ m.njnt,
+ m.nlight,
+ m._impl.nlsp,
+ m.nmeshface,
+ m.nmocap,
+ m.nsite,
+ m.ntendon,
+ m.nu,
+ m.nv,
+ m._impl.nxn_geom_pair_filtered,
+ m._impl.nxn_pairid,
+ m._impl.nxn_pairid_filtered,
+ m.pair_dim,
+ m.pair_friction,
+ m.pair_gap,
+ m.pair_margin,
+ m.pair_solimp,
+ m.pair_solref,
+ m.pair_solreffriction,
+ m._impl.plugin,
+ m._impl.plugin_attr,
+ m._impl.qLD_updates,
+ m._impl.qM_fullm_i,
+ m._impl.qM_fullm_j,
+ m._impl.qM_madr_ij,
+ m._impl.qM_mulm_i,
+ m._impl.qM_mulm_j,
+ m._impl.qM_tiles,
+ m.qpos0,
+ m.qpos_spring,
+ m._impl.rangefinder_sensor_adr,
+ m._impl.sensor_acc_adr,
+ m.sensor_adr,
+ m.sensor_cutoff,
+ m.sensor_datatype,
+ m._impl.sensor_e_kinetic,
+ m._impl.sensor_e_potential,
+ m._impl.sensor_limitfrc_adr,
+ m._impl.sensor_limitpos_adr,
+ m._impl.sensor_limitvel_adr,
+ m.sensor_objid,
+ m.sensor_objtype,
+ m._impl.sensor_pos_adr,
+ m._impl.sensor_rangefinder_adr,
+ m._impl.sensor_rangefinder_bodyid,
+ m.sensor_refid,
+ m.sensor_reftype,
+ m._impl.sensor_rne_postconstraint,
+ m._impl.sensor_subtree_vel,
+ m._impl.sensor_tendonactfrc_adr,
+ m._impl.sensor_touch_adr,
+ m.sensor_type,
+ m._impl.sensor_vel_adr,
+ m.site_bodyid,
+ m.site_pos,
+ m.site_quat,
+ m.site_size,
+ m.site_type,
+ m._impl.subtree_mass,
+ m.tendon_actfrclimited,
+ m.tendon_actfrcrange,
+ m.tendon_adr,
+ m.tendon_armature,
+ m.tendon_damping,
+ m.tendon_frictionloss,
+ m._impl.tendon_geom_adr,
+ m.tendon_invweight0,
+ m._impl.tendon_jnt_adr,
+ m.tendon_length0,
+ m.tendon_lengthspring,
+ m._impl.tendon_limited_adr,
+ m.tendon_margin,
+ m.tendon_num,
+ m.tendon_range,
+ m._impl.tendon_site_pair_adr,
+ m.tendon_solimp_fri,
+ m.tendon_solimp_lim,
+ m.tendon_solref_fri,
+ m.tendon_solref_lim,
+ m.tendon_stiffness,
+ 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,
+ m.opt._impl.broadphase,
+ m.opt._impl.broadphase_filter,
+ m.opt.cone,
+ m.opt.density,
+ m.opt.disableflags,
+ m.opt.enableflags,
+ m.opt._impl.epa_iterations,
+ m.opt._impl.gjk_iterations,
+ m.opt._impl.graph_conditional,
+ m.opt.gravity,
+ m.opt._impl.has_fluid,
+ m.opt.impratio,
+ m.opt.integrator,
+ m.opt._impl.is_sparse,
+ m.opt.iterations,
+ m.opt.ls_iterations,
+ m.opt._impl.ls_parallel,
+ m.opt.ls_tolerance,
+ m.opt.magnetic,
+ m.opt._impl.run_collision_detection,
+ m.opt._impl.sdf_initpoints,
+ m.opt._impl.sdf_iterations,
+ m.opt.solver,
+ m.opt.timestep,
+ m.opt.tolerance,
+ m.opt.viscosity,
+ m.opt.wind,
+ m.stat.meaninertia,
+ d._impl.nconmax,
+ d._impl.njmax,
+ d.act,
+ d.act_dot,
+ d._impl.act_dot_rk,
+ d._impl.act_t0,
+ d.actuator_force,
+ d._impl.actuator_length,
+ d._impl.actuator_moment,
+ d._impl.actuator_trntype_body_ncon,
+ d._impl.actuator_velocity,
+ d._impl.cacc,
+ d.cam_xmat,
+ d.cam_xpos,
+ d._impl.cdof,
+ d._impl.cdof_dot,
+ d._impl.cfrc_ext,
+ d._impl.cfrc_int,
+ d._impl.cinert,
+ d._impl.collision_hftri_index,
+ d._impl.collision_pair,
+ d._impl.collision_pairid,
+ d._impl.collision_worldid,
+ d._impl.crb,
+ d.ctrl,
+ d.cvel,
+ d._impl.energy,
+ d._impl.epa_face,
+ d._impl.epa_horizon,
+ d._impl.epa_index,
+ d._impl.epa_map,
+ d._impl.epa_norm2,
+ d._impl.epa_pr,
+ d._impl.epa_vert,
+ d._impl.epa_vert1,
+ d._impl.epa_vert2,
+ d._impl.epa_vert_index1,
+ d._impl.epa_vert_index2,
+ d.eq_active,
+ d._impl.flexedge_length,
+ d._impl.flexedge_velocity,
+ d._impl.flexvert_xpos,
+ d._impl.fluid_applied,
+ d._impl.geom_skip,
+ d.geom_xmat,
+ d.geom_xpos,
+ d._impl.inverse_mul_m_skip,
+ d._impl.light_xdir,
+ d._impl.light_xpos,
+ d.mocap_pos,
+ d.mocap_quat,
+ d._impl.ncollision,
+ d._impl.ncon,
+ d._impl.ncon_hfield,
+ d._impl.ne,
+ d._impl.ne_connect,
+ d._impl.ne_jnt,
+ d._impl.ne_ten,
+ d._impl.ne_weld,
+ d._impl.nefc,
+ d._impl.nf,
+ d._impl.nl,
+ d._impl.nsolving,
+ d._impl.qLD,
+ d._impl.qLD_integration,
+ d._impl.qLDiagInv,
+ d._impl.qLDiagInv_integration,
+ d._impl.qM,
+ d._impl.qM_integration,
+ d.qacc,
+ d._impl.qacc_integration,
+ d._impl.qacc_rk,
+ d.qacc_smooth,
+ d.qacc_warmstart,
+ d.qfrc_actuator,
+ d.qfrc_applied,
+ d.qfrc_bias,
+ d.qfrc_constraint,
+ d._impl.qfrc_damper,
+ d.qfrc_fluid,
+ d.qfrc_gravcomp,
+ d._impl.qfrc_integration,
+ d.qfrc_passive,
+ d.qfrc_smooth,
+ d._impl.qfrc_spring,
+ d.qpos,
+ d._impl.qpos_t0,
+ d.qvel,
+ d._impl.qvel_rk,
+ d._impl.qvel_t0,
+ d._impl.sap_cumulative_sum,
+ d._impl.sap_projection_lower,
+ d._impl.sap_projection_upper,
+ d._impl.sap_range,
+ d._impl.sap_segment_index,
+ d._impl.sap_sort_index,
+ d._impl.sensor_rangefinder_dist,
+ d._impl.sensor_rangefinder_geomid,
+ d._impl.sensor_rangefinder_pnt,
+ d._impl.sensor_rangefinder_vec,
+ d.sensordata,
+ d.site_xmat,
+ d.site_xpos,
+ d._impl.solver_niter,
+ d._impl.subtree_angmom,
+ d._impl.subtree_bodyvel,
+ d.subtree_com,
+ d._impl.subtree_linvel,
+ d._impl.ten_J,
+ d._impl.ten_Jdot,
+ d._impl.ten_actfrc,
+ d._impl.ten_bias_coef,
+ d._impl.ten_length,
+ d._impl.ten_velocity,
+ d._impl.ten_wrapadr,
+ d._impl.ten_wrapnum,
+ d.time,
+ d._impl.wrap_geom_xpos,
+ d._impl.wrap_obj,
+ d._impl.wrap_xpos,
+ d.xanchor,
+ d.xaxis,
+ d.xfrc_applied,
+ d.ximat,
+ d.xipos,
+ d.xmat,
+ d.xpos,
+ d.xquat,
+ d._impl.contact__dim,
+ d._impl.contact__dist,
+ d._impl.contact__efc_address,
+ d._impl.contact__frame,
+ d._impl.contact__friction,
+ d._impl.contact__geom,
+ d._impl.contact__includemargin,
+ d._impl.contact__pos,
+ d._impl.contact__solimp,
+ d._impl.contact__solref,
+ d._impl.contact__solreffriction,
+ d._impl.contact__worldid,
+ d._impl.efc__D,
+ d._impl.efc__J,
+ d._impl.efc__Jaref,
+ d._impl.efc__Ma,
+ d._impl.efc__Mgrad,
+ d._impl.efc__active,
+ d._impl.efc__alpha,
+ d._impl.efc__aref,
+ d._impl.efc__beta,
+ d._impl.efc__beta_den,
+ d._impl.efc__beta_num,
+ d._impl.efc__cholesky_L_tmp,
+ d._impl.efc__cholesky_y_tmp,
+ d._impl.efc__condim,
+ d._impl.efc__cost,
+ d._impl.efc__cost_candidate,
+ d._impl.efc__done,
+ d._impl.efc__force,
+ d._impl.efc__frictionloss,
+ d._impl.efc__gauss,
+ d._impl.efc__grad,
+ d._impl.efc__grad_dot,
+ d._impl.efc__gtol,
+ d._impl.efc__h,
+ d._impl.efc__hi,
+ d._impl.efc__hi_alpha,
+ d._impl.efc__hi_next,
+ d._impl.efc__hi_next_alpha,
+ d._impl.efc__id,
+ d._impl.efc__jv,
+ d._impl.efc__lo,
+ d._impl.efc__lo_alpha,
+ d._impl.efc__lo_next,
+ d._impl.efc__lo_next_alpha,
+ d._impl.efc__ls_done,
+ d._impl.efc__margin,
+ d._impl.efc__mid,
+ d._impl.efc__mid_alpha,
+ d._impl.efc__mv,
+ d._impl.efc__p0,
+ d._impl.efc__pos,
+ d._impl.efc__prev_Mgrad,
+ d._impl.efc__prev_cost,
+ d._impl.efc__prev_grad,
+ d._impl.efc__quad,
+ d._impl.efc__quad_gauss,
+ d._impl.efc__search,
+ d._impl.efc__search_dot,
+ d._impl.efc__type,
+ d._impl.efc__u,
+ d._impl.efc__uu,
+ d._impl.efc__uv,
+ d._impl.efc__vel,
+ d._impl.efc__vv,
+ )
+ d = d.tree_replace({
+ 'act': out[0],
+ 'act_dot': out[1],
+ '_impl.act_dot_rk': out[2],
+ '_impl.act_t0': out[3],
+ 'actuator_force': out[4],
+ '_impl.actuator_length': out[5],
+ '_impl.actuator_moment': out[6],
+ '_impl.actuator_trntype_body_ncon': out[7],
+ '_impl.actuator_velocity': out[8],
+ '_impl.cacc': out[9],
+ 'cam_xmat': out[10],
+ 'cam_xpos': out[11],
+ '_impl.cdof': out[12],
+ '_impl.cdof_dot': out[13],
+ '_impl.cfrc_ext': out[14],
+ '_impl.cfrc_int': out[15],
+ '_impl.cinert': out[16],
+ '_impl.collision_hftri_index': out[17],
+ '_impl.collision_pair': out[18],
+ '_impl.collision_pairid': out[19],
+ '_impl.collision_worldid': out[20],
+ '_impl.crb': out[21],
+ 'ctrl': out[22],
+ 'cvel': out[23],
+ '_impl.energy': out[24],
+ '_impl.epa_face': out[25],
+ '_impl.epa_horizon': out[26],
+ '_impl.epa_index': out[27],
+ '_impl.epa_map': out[28],
+ '_impl.epa_norm2': out[29],
+ '_impl.epa_pr': out[30],
+ '_impl.epa_vert': out[31],
+ '_impl.epa_vert1': out[32],
+ '_impl.epa_vert2': out[33],
+ '_impl.epa_vert_index1': out[34],
+ '_impl.epa_vert_index2': out[35],
+ 'eq_active': out[36],
+ '_impl.flexedge_length': out[37],
+ '_impl.flexedge_velocity': out[38],
+ '_impl.flexvert_xpos': out[39],
+ '_impl.fluid_applied': out[40],
+ '_impl.geom_skip': out[41],
+ 'geom_xmat': out[42],
+ 'geom_xpos': out[43],
+ '_impl.inverse_mul_m_skip': out[44],
+ '_impl.light_xdir': out[45],
+ '_impl.light_xpos': out[46],
+ 'mocap_pos': out[47],
+ 'mocap_quat': out[48],
+ '_impl.ncollision': out[49],
+ '_impl.ncon': out[50],
+ '_impl.ncon_hfield': out[51],
+ '_impl.ne': out[52],
+ '_impl.ne_connect': out[53],
+ '_impl.ne_jnt': out[54],
+ '_impl.ne_ten': out[55],
+ '_impl.ne_weld': out[56],
+ '_impl.nefc': out[57],
+ '_impl.nf': out[58],
+ '_impl.nl': out[59],
+ '_impl.nsolving': out[60],
+ '_impl.qLD': out[61],
+ '_impl.qLD_integration': out[62],
+ '_impl.qLDiagInv': out[63],
+ '_impl.qLDiagInv_integration': out[64],
+ '_impl.qM': out[65],
+ '_impl.qM_integration': out[66],
+ 'qacc': out[67],
+ '_impl.qacc_integration': out[68],
+ '_impl.qacc_rk': out[69],
+ 'qacc_smooth': out[70],
+ 'qacc_warmstart': out[71],
+ 'qfrc_actuator': out[72],
+ 'qfrc_applied': out[73],
+ 'qfrc_bias': out[74],
+ 'qfrc_constraint': out[75],
+ '_impl.qfrc_damper': out[76],
+ 'qfrc_fluid': out[77],
+ 'qfrc_gravcomp': out[78],
+ '_impl.qfrc_integration': out[79],
+ 'qfrc_passive': out[80],
+ 'qfrc_smooth': out[81],
+ '_impl.qfrc_spring': out[82],
+ 'qpos': out[83],
+ '_impl.qpos_t0': out[84],
+ 'qvel': out[85],
+ '_impl.qvel_rk': out[86],
+ '_impl.qvel_t0': out[87],
+ '_impl.sap_cumulative_sum': out[88],
+ '_impl.sap_projection_lower': out[89],
+ '_impl.sap_projection_upper': out[90],
+ '_impl.sap_range': out[91],
+ '_impl.sap_segment_index': out[92],
+ '_impl.sap_sort_index': out[93],
+ '_impl.sensor_rangefinder_dist': out[94],
+ '_impl.sensor_rangefinder_geomid': out[95],
+ '_impl.sensor_rangefinder_pnt': out[96],
+ '_impl.sensor_rangefinder_vec': out[97],
+ 'sensordata': out[98],
+ 'site_xmat': out[99],
+ 'site_xpos': out[100],
+ '_impl.solver_niter': out[101],
+ '_impl.subtree_angmom': out[102],
+ '_impl.subtree_bodyvel': out[103],
+ 'subtree_com': out[104],
+ '_impl.subtree_linvel': out[105],
+ '_impl.ten_J': out[106],
+ '_impl.ten_Jdot': out[107],
+ '_impl.ten_actfrc': out[108],
+ '_impl.ten_bias_coef': out[109],
+ '_impl.ten_length': out[110],
+ '_impl.ten_velocity': out[111],
+ '_impl.ten_wrapadr': out[112],
+ '_impl.ten_wrapnum': out[113],
+ 'time': out[114],
+ '_impl.wrap_geom_xpos': out[115],
+ '_impl.wrap_obj': out[116],
+ '_impl.wrap_xpos': out[117],
+ 'xanchor': out[118],
+ 'xaxis': out[119],
+ 'xfrc_applied': out[120],
+ 'ximat': out[121],
+ 'xipos': out[122],
+ 'xmat': out[123],
+ 'xpos': out[124],
+ 'xquat': out[125],
+ '_impl.contact__dim': out[126],
+ '_impl.contact__dist': out[127],
+ '_impl.contact__efc_address': out[128],
+ '_impl.contact__frame': out[129],
+ '_impl.contact__friction': out[130],
+ '_impl.contact__geom': out[131],
+ '_impl.contact__includemargin': out[132],
+ '_impl.contact__pos': out[133],
+ '_impl.contact__solimp': out[134],
+ '_impl.contact__solref': out[135],
+ '_impl.contact__solreffriction': out[136],
+ '_impl.contact__worldid': out[137],
+ '_impl.efc__D': out[138],
+ '_impl.efc__J': out[139],
+ '_impl.efc__Jaref': out[140],
+ '_impl.efc__Ma': out[141],
+ '_impl.efc__Mgrad': out[142],
+ '_impl.efc__active': out[143],
+ '_impl.efc__alpha': out[144],
+ '_impl.efc__aref': out[145],
+ '_impl.efc__beta': out[146],
+ '_impl.efc__beta_den': out[147],
+ '_impl.efc__beta_num': out[148],
+ '_impl.efc__cholesky_L_tmp': out[149],
+ '_impl.efc__cholesky_y_tmp': out[150],
+ '_impl.efc__condim': out[151],
+ '_impl.efc__cost': out[152],
+ '_impl.efc__cost_candidate': out[153],
+ '_impl.efc__done': out[154],
+ '_impl.efc__force': out[155],
+ '_impl.efc__frictionloss': out[156],
+ '_impl.efc__gauss': out[157],
+ '_impl.efc__grad': out[158],
+ '_impl.efc__grad_dot': out[159],
+ '_impl.efc__gtol': out[160],
+ '_impl.efc__h': out[161],
+ '_impl.efc__hi': out[162],
+ '_impl.efc__hi_alpha': out[163],
+ '_impl.efc__hi_next': out[164],
+ '_impl.efc__hi_next_alpha': out[165],
+ '_impl.efc__id': out[166],
+ '_impl.efc__jv': out[167],
+ '_impl.efc__lo': out[168],
+ '_impl.efc__lo_alpha': out[169],
+ '_impl.efc__lo_next': out[170],
+ '_impl.efc__lo_next_alpha': out[171],
+ '_impl.efc__ls_done': out[172],
+ '_impl.efc__margin': out[173],
+ '_impl.efc__mid': out[174],
+ '_impl.efc__mid_alpha': out[175],
+ '_impl.efc__mv': out[176],
+ '_impl.efc__p0': out[177],
+ '_impl.efc__pos': out[178],
+ '_impl.efc__prev_Mgrad': out[179],
+ '_impl.efc__prev_cost': out[180],
+ '_impl.efc__prev_grad': out[181],
+ '_impl.efc__quad': out[182],
+ '_impl.efc__quad_gauss': out[183],
+ '_impl.efc__search': out[184],
+ '_impl.efc__search_dot': out[185],
+ '_impl.efc__type': out[186],
+ '_impl.efc__u': out[187],
+ '_impl.efc__uu': out[188],
+ '_impl.efc__uv': out[189],
+ '_impl.efc__vel': out[190],
+ '_impl.efc__vv': out[191],
+ })
+ return d
+
+
+@jax.custom_batching.custom_vmap
+@ffi.marshal_jax_warp_callable
+def step(m: types.Model, d: types.Data):
+ return _step_jax_impl(m, d)
+
+
+@step.def_vmap
+@ffi.marshal_custom_vmap
+def step_vmap(unused_axis_size, is_batched, m, d):
+ d = step(m, d)
+ return d, is_batched[1]
diff --git a/mjx/mujoco/mjx/warp/forward_test.py b/mjx/mujoco/mjx/warp/forward_test.py
new file mode 100644
index 00000000..c8e0ddef
--- /dev/null
+++ b/mjx/mujoco/mjx/warp/forward_test.py
@@ -0,0 +1,274 @@
+# Copyright 2025 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 forward functions."""
+
+import functools
+import logging
+
+from absl.testing import absltest
+from absl.testing import parameterized
+import jax
+import jax.numpy as jp
+import mujoco
+from mujoco import mjx
+from mujoco.mjx._src import io
+from mujoco.mjx._src import test_util
+import mujoco.mjx.warp as mjxw
+from mujoco.mjx.warp import test_util as tu
+from mujoco.mjx.warp import warp as wp # pylint: disable=g-importing-member
+import numpy as np
+
+try:
+ from mujoco.mjx.warp import forward # pylint: disable=g-import-not-at-top
+except ImportError:
+ forward = None
+
+
+class ForwardTest(parameterized.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ if mjxw.WARP_INSTALLED:
+ wp.clear_kernel_cache()
+ np.random.seed(0)
+
+ @parameterized.parameters(
+ 'pendula.xml',
+ 'humanoid/humanoid.xml',
+ )
+ def test_jit_caching(self, xml):
+ """Tests jit caching on the full step function."""
+ if not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp not installed.')
+ if not io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device available.')
+
+ batch_size = 7
+ m = test_util.load_test_file(xml)
+ mx = mjx.put_model(m, impl='warp')
+
+ keys = jp.arange(batch_size)
+ dx_batch = jax.vmap(functools.partial(tu.make_data, m))(keys)
+
+ step_fn = jax.jit(jax.vmap(forward.step, in_axes=(None, 0)))
+
+ was_logging_compiles = jax.config.jax_log_compiles
+ jax_logger = logging.getLogger('jax')
+ was_propagating = jax_logger.propagate
+ jax.config.update('jax_log_compiles', True)
+ jax_logger.propagate = False # do not print to stdout for this test
+ with self.assertLogs('jax', level='INFO') as log:
+ dx_batch1 = step_fn(mx, dx_batch)
+ jax.tree_util.tree_map(lambda x: x.block_until_ready(), dx_batch1)
+
+ # Re-generate data and run step_fn again to test jit caching.
+ keys = jp.arange(batch_size, batch_size * 2)
+ dx_batch = jax.vmap(functools.partial(tu.make_data, m))(keys)
+ dx_batch2 = step_fn(mx, dx_batch)
+ jax.tree_util.tree_map(lambda x: x.block_until_ready(), dx_batch2)
+ jax.config.update('jax_log_compiles', was_logging_compiles)
+ jax_logger.propagate = was_propagating
+
+ compilation_logs = [
+ r for r in log.records if 'Compiling jit(step)' in r.getMessage()
+ ]
+ self.assertLen(
+ compilation_logs,
+ 1,
+ msg=(
+ f'Expected 1 compilation, got {len(compilation_logs)} compilations.'
+ ),
+ )
+
+ @parameterized.product(
+ xml=(
+ 'humanoid/humanoid.xml',
+ 'pendula.xml',
+ ),
+ batch_size=(1, 7),
+ )
+ def test_forward(self, xml: str, batch_size: int):
+ if not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp not installed.')
+ if not io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device available.')
+
+ m = test_util.load_test_file(xml)
+ m.opt.iterations = 10
+ m.opt.ls_iterations = 10
+ mx = mjx.put_model(m, impl='warp')
+
+ d = mujoco.MjData(m)
+ worldids = jp.arange(batch_size)
+ dx_batch = jax.vmap(functools.partial(tu.make_data, m))(worldids)
+
+ dx_batch = jax.jit(jax.vmap(forward.forward, in_axes=(None, 0)))(
+ mx, dx_batch
+ )
+
+ for i in range(batch_size):
+ dx = dx_batch[i]
+
+ d.qpos[:] = dx.qpos
+ d.qvel[:] = dx.qvel
+ d.ctrl[:] = dx.ctrl
+ d.mocap_pos[:] = dx.mocap_pos
+ d.mocap_quat[:] = dx.mocap_quat
+ mujoco.mj_forward(m, d)
+
+ # fwd_position
+ tu.assert_attr_eq(dx, d, 'xpos')
+ tu.assert_attr_eq(dx, d, 'xquat')
+ tu.assert_attr_eq(dx, d, 'xipos')
+ tu.assert_eq(d.ximat.reshape((-1, 3, 3)), dx.ximat, 'ximat')
+ tu.assert_attr_eq(dx, d, 'xanchor')
+ tu.assert_attr_eq(dx, d, 'xaxis')
+ tu.assert_attr_eq(dx, d, 'geom_xpos')
+ tu.assert_eq(dx.geom_xmat, d.geom_xmat.reshape((-1, 3, 3)), 'geom_xmat')
+ if m.nsite:
+ tu.assert_attr_eq(dx, d, 'site_xpos')
+ tu.assert_eq(dx.site_xmat, d.site_xmat.reshape((-1, 3, 3)), 'site_xmat')
+ tu.assert_attr_eq(dx._impl, d, 'cdof')
+ tu.assert_attr_eq(dx._impl, d, 'cinert')
+ tu.assert_attr_eq(dx, d, 'subtree_com')
+ if m.nlight:
+ tu.assert_attr_eq(dx._impl, d, 'light_xpos')
+ tu.assert_attr_eq(dx._impl, d, 'light_xdir')
+ 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._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')
+ tu.assert_attr_eq(dx._impl, d, 'wrap_xpos')
+ tu.assert_attr_eq(dx._impl, d, 'wrap_obj')
+ tu.assert_attr_eq(dx._impl, d, 'crb')
+
+ if not mx.opt._impl.is_sparse:
+ qm = np.zeros((m.nv, m.nv))
+ mujoco.mj_fullM(m, qm, d.qM)
+ else:
+ qm = d.qM
+ tu.assert_eq(qm, dx._impl.qM, 'qM')
+ # qLD is fused in a cholesky factorize and solve, and not written to.
+
+ tu.assert_contact_eq(d, dx, worldid=i)
+
+ tu.assert_attr_eq(dx._impl, d, 'actuator_length')
+ actuator_moment = np.zeros((m.nu, m.nv))
+ mujoco.mju_sparse2dense(
+ actuator_moment,
+ d.actuator_moment,
+ d.moment_rownnz,
+ d.moment_rowadr,
+ d.moment_colind,
+ )
+ tu.assert_eq(dx._impl.actuator_moment, actuator_moment, 'actuator_moment')
+
+ # fwd_velocity
+ tu.assert_attr_eq(dx._impl, d, 'actuator_velocity')
+ tu.assert_attr_eq(dx, d, 'cvel')
+ tu.assert_attr_eq(dx._impl, d, 'cdof_dot')
+ tu.assert_attr_eq(dx._impl, d, 'qfrc_spring')
+ tu.assert_attr_eq(dx._impl, d, 'qfrc_damper')
+ tu.assert_attr_eq(dx, d, 'qfrc_gravcomp')
+ tu.assert_attr_eq(dx, d, 'qfrc_fluid')
+ tu.assert_attr_eq(dx, d, 'qfrc_passive')
+ tu.assert_attr_eq(dx, d, 'qfrc_bias')
+ tu.assert_efc_eq(d, dx, worldid=i)
+
+ # fwd_actuation
+ tu.assert_attr_eq(dx, d, 'act_dot')
+ tu.assert_attr_eq(dx, d, 'actuator_force')
+ tu.assert_attr_eq(dx, d, 'qfrc_actuator')
+
+ # fwd_acceleration
+ tu.assert_attr_eq(dx, d, 'qfrc_smooth')
+ tu.assert_attr_eq(dx, d, 'qacc_smooth')
+
+ # solve
+ np.testing.assert_allclose(
+ dx.qacc_warmstart,
+ d.qacc_warmstart,
+ err_msg='qacc_warmstart',
+ rtol=1e-5,
+ atol=1.0,
+ )
+ np.testing.assert_allclose(
+ dx.qacc, d.qacc, err_msg='qacc', rtol=1e-5, atol=1.0
+ )
+
+
+class StepTest(parameterized.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ if mjxw.WARP_INSTALLED:
+ wp.clear_kernel_cache()
+ np.random.seed(0)
+
+ @parameterized.product(
+ xml=(
+ 'humanoid/humanoid.xml',
+ 'pendula.xml',
+ ),
+ batch_size=(1, 7),
+ )
+ def test_step(self, xml: str, batch_size: int):
+ if not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp not installed.')
+ if not io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device available.')
+
+ m = test_util.load_test_file(xml)
+ m.opt.iterations = 10
+ m.opt.ls_iterations = 10
+ mx = mjx.put_model(m, impl='warp')
+
+ d = mujoco.MjData(m)
+ worldids = jp.arange(batch_size)
+ dx_batch = jax.vmap(functools.partial(tu.make_data, m))(worldids)
+ dx_batch_orig = dx_batch
+
+ for _ in range(10):
+ dx_batch = jax.jit(jax.vmap(forward.step, in_axes=(None, 0)))(
+ mx, dx_batch
+ )
+
+ for i in range(batch_size):
+ dx = dx_batch[i]
+ dx_orig = dx_batch_orig[i]
+
+ d.qpos[:] = dx_orig.qpos
+ d.qvel[:] = dx_orig.qvel
+ d.ctrl[:] = dx_orig.ctrl
+ d.mocap_pos[:] = dx_orig.mocap_pos
+ d.mocap_quat[:] = dx_orig.mocap_quat
+ d.time = dx_orig.time
+ mujoco.mj_step(m, d, 10)
+
+ tu.assert_attr_eq(dx, d, 'qpos')
+ tu.assert_attr_eq(dx, d, 'qvel')
+ tu.assert_attr_eq(dx, d, 'time')
+ tu.assert_attr_eq(dx, d, 'ctrl')
+ tu.assert_attr_eq(dx, d, 'act')
+ tu.assert_attr_eq(dx, d, 'mocap_pos')
+ tu.assert_attr_eq(dx, d, 'mocap_quat')
+ tu.assert_attr_eq(dx, d, 'sensordata')
+
+
+if __name__ == '__main__':
+ absltest.main()
diff --git a/mjx/mujoco/mjx/warp/smooth.py b/mjx/mujoco/mjx/warp/smooth.py
new file mode 100644
index 00000000..087d5aab
--- /dev/null
+++ b/mjx/mujoco/mjx/warp/smooth.py
@@ -0,0 +1,291 @@
+# Copyright 2025 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.
+# ==============================================================================
+
+"""DO NOT EDIT. This file is auto-generated."""
+import dataclasses
+import jax
+from mujoco.mjx._src import types
+from mujoco.mjx.warp import ffi
+import mujoco.mjx.third_party.mujoco_warp as mjwarp
+from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
+import warp as wp
+
+
+_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 _kinematics_shim(
+ # Model
+ nworld: int,
+ body_dofadr: wp.array(dtype=int),
+ body_ipos: wp.array2d(dtype=wp.vec3),
+ body_iquat: wp.array2d(dtype=wp.quat),
+ body_jntadr: wp.array(dtype=int),
+ body_jntnum: wp.array(dtype=int),
+ body_parentid: wp.array(dtype=int),
+ body_pos: wp.array2d(dtype=wp.vec3),
+ body_quat: wp.array2d(dtype=wp.quat),
+ body_tree: tuple[wp.array(dtype=int), ...],
+ flex_edge: wp.array(dtype=wp.vec2i),
+ flex_vertadr: wp.array(dtype=int),
+ flex_vertbodyid: wp.array(dtype=int),
+ geom_bodyid: wp.array(dtype=int),
+ geom_pos: wp.array2d(dtype=wp.vec3),
+ geom_quat: wp.array2d(dtype=wp.quat),
+ jnt_axis: wp.array2d(dtype=wp.vec3),
+ jnt_pos: wp.array2d(dtype=wp.vec3),
+ jnt_qposadr: wp.array(dtype=int),
+ jnt_type: wp.array(dtype=int),
+ mocap_bodyid: wp.array(dtype=int),
+ nflexedge: int,
+ nflexvert: int,
+ ngeom: int,
+ nmocap: int,
+ nsite: int,
+ qpos0: wp.array2d(dtype=float),
+ site_bodyid: wp.array(dtype=int),
+ site_pos: wp.array2d(dtype=wp.vec3),
+ site_quat: wp.array2d(dtype=wp.quat),
+ # Data
+ flexedge_length: wp.array2d(dtype=float),
+ flexedge_velocity: wp.array2d(dtype=float),
+ flexvert_xpos: wp.array2d(dtype=wp.vec3),
+ geom_skip: wp.array(dtype=bool),
+ geom_xmat: wp.array2d(dtype=wp.mat33),
+ geom_xpos: wp.array2d(dtype=wp.vec3),
+ mocap_pos: wp.array2d(dtype=wp.vec3),
+ mocap_quat: wp.array2d(dtype=wp.quat),
+ qpos: wp.array2d(dtype=float),
+ qvel: wp.array2d(dtype=float),
+ site_xmat: wp.array2d(dtype=wp.mat33),
+ site_xpos: wp.array2d(dtype=wp.vec3),
+ xanchor: wp.array2d(dtype=wp.vec3),
+ xaxis: wp.array2d(dtype=wp.vec3),
+ ximat: wp.array2d(dtype=wp.mat33),
+ xipos: wp.array2d(dtype=wp.vec3),
+ xmat: wp.array2d(dtype=wp.mat33),
+ xpos: wp.array2d(dtype=wp.vec3),
+ xquat: wp.array2d(dtype=wp.quat),
+):
+ _m.stat = _s
+ _m.opt = _o
+ _d.efc = _e
+ _d.contact = _c
+ _m.body_dofadr = body_dofadr
+ _m.body_ipos = body_ipos
+ _m.body_iquat = body_iquat
+ _m.body_jntadr = body_jntadr
+ _m.body_jntnum = body_jntnum
+ _m.body_parentid = body_parentid
+ _m.body_pos = body_pos
+ _m.body_quat = body_quat
+ _m.body_tree = body_tree
+ _m.flex_edge = flex_edge
+ _m.flex_vertadr = flex_vertadr
+ _m.flex_vertbodyid = flex_vertbodyid
+ _m.geom_bodyid = geom_bodyid
+ _m.geom_pos = geom_pos
+ _m.geom_quat = geom_quat
+ _m.jnt_axis = jnt_axis
+ _m.jnt_pos = jnt_pos
+ _m.jnt_qposadr = jnt_qposadr
+ _m.jnt_type = jnt_type
+ _m.mocap_bodyid = mocap_bodyid
+ _m.nflexedge = nflexedge
+ _m.nflexvert = nflexvert
+ _m.ngeom = ngeom
+ _m.nmocap = nmocap
+ _m.nsite = nsite
+ _m.qpos0 = qpos0
+ _m.site_bodyid = site_bodyid
+ _m.site_pos = site_pos
+ _m.site_quat = site_quat
+ _d.flexedge_length = flexedge_length
+ _d.flexedge_velocity = flexedge_velocity
+ _d.flexvert_xpos = flexvert_xpos
+ _d.geom_skip = geom_skip
+ _d.geom_xmat = geom_xmat
+ _d.geom_xpos = geom_xpos
+ _d.mocap_pos = mocap_pos
+ _d.mocap_quat = mocap_quat
+ _d.qpos = qpos
+ _d.qvel = qvel
+ _d.site_xmat = site_xmat
+ _d.site_xpos = site_xpos
+ _d.xanchor = xanchor
+ _d.xaxis = xaxis
+ _d.ximat = ximat
+ _d.xipos = xipos
+ _d.xmat = xmat
+ _d.xpos = xpos
+ _d.xquat = xquat
+ _d.nworld = nworld
+ mjwarp.kinematics(_m, _d)
+
+
+def _kinematics_jax_impl(m: types.Model, d: types.Data):
+ output_dims = {
+ 'flexedge_length': d._impl.flexedge_length.shape,
+ 'flexedge_velocity': d._impl.flexedge_velocity.shape,
+ 'flexvert_xpos': d._impl.flexvert_xpos.shape,
+ 'geom_skip': d._impl.geom_skip.shape,
+ 'geom_xmat': d.geom_xmat.shape,
+ 'geom_xpos': d.geom_xpos.shape,
+ 'mocap_pos': d.mocap_pos.shape,
+ 'mocap_quat': d.mocap_quat.shape,
+ 'qpos': d.qpos.shape,
+ 'qvel': d.qvel.shape,
+ 'site_xmat': d.site_xmat.shape,
+ 'site_xpos': d.site_xpos.shape,
+ 'xanchor': d.xanchor.shape,
+ 'xaxis': d.xaxis.shape,
+ 'ximat': d.ximat.shape,
+ 'xipos': d.xipos.shape,
+ 'xmat': d.xmat.shape,
+ 'xpos': d.xpos.shape,
+ 'xquat': d.xquat.shape,
+ }
+ jf = ffi.jax_callable_variadic_tuple(
+ _kinematics_shim,
+ num_outputs=19,
+ output_dims=output_dims,
+ vmap_method=None,
+ graph_compatible=True,
+ in_out_argnames={
+ 'flexedge_length',
+ 'flexedge_velocity',
+ 'flexvert_xpos',
+ 'geom_skip',
+ 'geom_xmat',
+ 'geom_xpos',
+ 'mocap_pos',
+ 'mocap_quat',
+ 'qpos',
+ 'qvel',
+ 'site_xmat',
+ 'site_xpos',
+ 'xanchor',
+ 'xaxis',
+ 'ximat',
+ 'xipos',
+ 'xmat',
+ 'xpos',
+ 'xquat',
+ },
+ )
+ out = jf(
+ d.qpos.shape[0],
+ m.body_dofadr,
+ m.body_ipos,
+ m.body_iquat,
+ m.body_jntadr,
+ m.body_jntnum,
+ m.body_parentid,
+ m.body_pos,
+ m.body_quat,
+ m._impl.body_tree,
+ m._impl.flex_edge,
+ m._impl.flex_vertadr,
+ m._impl.flex_vertbodyid,
+ m.geom_bodyid,
+ m.geom_pos,
+ m.geom_quat,
+ m.jnt_axis,
+ m.jnt_pos,
+ m.jnt_qposadr,
+ m.jnt_type,
+ m._impl.mocap_bodyid,
+ m._impl.nflexedge,
+ m._impl.nflexvert,
+ m.ngeom,
+ m.nmocap,
+ m.nsite,
+ m.qpos0,
+ m.site_bodyid,
+ m.site_pos,
+ m.site_quat,
+ d._impl.flexedge_length,
+ d._impl.flexedge_velocity,
+ d._impl.flexvert_xpos,
+ d._impl.geom_skip,
+ d.geom_xmat,
+ d.geom_xpos,
+ d.mocap_pos,
+ d.mocap_quat,
+ d.qpos,
+ d.qvel,
+ d.site_xmat,
+ d.site_xpos,
+ d.xanchor,
+ d.xaxis,
+ d.ximat,
+ d.xipos,
+ d.xmat,
+ d.xpos,
+ d.xquat,
+ )
+ d = d.tree_replace({
+ '_impl.flexedge_length': out[0],
+ '_impl.flexedge_velocity': out[1],
+ '_impl.flexvert_xpos': out[2],
+ '_impl.geom_skip': out[3],
+ 'geom_xmat': out[4],
+ 'geom_xpos': out[5],
+ 'mocap_pos': out[6],
+ 'mocap_quat': out[7],
+ 'qpos': out[8],
+ 'qvel': out[9],
+ 'site_xmat': out[10],
+ 'site_xpos': out[11],
+ 'xanchor': out[12],
+ 'xaxis': out[13],
+ 'ximat': out[14],
+ 'xipos': out[15],
+ 'xmat': out[16],
+ 'xpos': out[17],
+ 'xquat': out[18],
+ })
+ return d
+
+
+@jax.custom_batching.custom_vmap
+@ffi.marshal_jax_warp_callable
+def kinematics(m: types.Model, d: types.Data):
+ return _kinematics_jax_impl(m, d)
+
+
+@kinematics.def_vmap
+@ffi.marshal_custom_vmap
+def kinematics_vmap(unused_axis_size, is_batched, m, d):
+ d = kinematics(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
new file mode 100644
index 00000000..7a87c24f
--- /dev/null
+++ b/mjx/mujoco/mjx/warp/smooth_test.py
@@ -0,0 +1,229 @@
+# Copyright 2025 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 codegen'd smooth functions."""
+
+import functools
+
+from absl.testing import absltest
+import jax
+from jax import numpy as jp
+import mujoco
+from mujoco import mjx
+from mujoco.mjx._src import io
+from mujoco.mjx._src import math
+import mujoco.mjx.warp as mjxw
+from mujoco.mjx.warp import test_util as tu
+from mujoco.mjx.warp import warp as wp # pylint: disable=g-importing-member
+import numpy as np
+
+try:
+ from mujoco.mjx.warp import smooth # pylint: disable=g-import-not-at-top
+except ImportError:
+ smooth = None
+
+
+class SmoothTest(absltest.TestCase):
+
+ def setUp(self):
+ super().setUp()
+ if mjxw.WARP_INSTALLED:
+ wp.clear_kernel_cache()
+ np.random.seed(0)
+
+ def test_kinematics(self):
+ """Tests kinematics with unbatched data."""
+ if not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp not installed.')
+ if not io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device available.')
+
+ m = tu.load_test_file('pendula.xml')
+
+ d = mujoco.MjData(m)
+ mx = mjx.put_model(m, impl='warp')
+
+ rng = jax.random.PRNGKey(0)
+ dx = mjx.make_data(m, impl='warp')
+ rng, key = jax.random.split(rng)
+ qpos = jax.random.uniform(key, (m.nq,))
+ _, key1, key2 = jax.random.split(rng, 3)
+ mocap_pos = jax.random.normal(key1, (m.nmocap, 3))
+ mocap_quat = jax.random.normal(key2, (m.nmocap, 4))
+ mocap_quat = math.normalize(mocap_quat)
+ dx = dx.replace(qpos=qpos, mocap_pos=mocap_pos, mocap_quat=mocap_quat)
+
+ dx = jax.jit(smooth.kinematics)(mx, dx)
+
+ d.qpos[:] = qpos
+ d.mocap_pos[:] = mocap_pos
+ d.mocap_quat[:] = mocap_quat
+ mujoco.mj_forward(m, d)
+
+ tu.assert_attr_eq(d, dx, 'xanchor')
+ tu.assert_attr_eq(d, dx, 'xaxis')
+ tu.assert_attr_eq(d, dx, 'xpos')
+ tu.assert_attr_eq(d, dx, 'xquat')
+ tu.assert_eq(d.xmat.reshape((-1, 3, 3)), dx.xmat, 'xmat')
+ tu.assert_attr_eq(d, dx, 'xipos')
+ tu.assert_eq(d.ximat.reshape((-1, 3, 3)), dx.ximat, 'ximat')
+ tu.assert_attr_eq(d, dx, 'geom_xpos')
+ tu.assert_eq(d.geom_xmat.reshape((-1, 3, 3)), dx.geom_xmat, 'geom_xmat')
+ tu.assert_attr_eq(d, dx, 'site_xpos')
+ tu.assert_eq(d.site_xmat.reshape((-1, 3, 3)), dx.site_xmat, 'site_xmat')
+
+ def test_kinematics_vmap(self):
+ """Tests kinematics with batched data."""
+ if not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp not installed.')
+ if not io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device available.')
+
+ m = tu.load_test_file('pendula.xml')
+
+ batch_size = 7
+ d = mujoco.MjData(m)
+ mx = mjx.put_model(m, impl='warp')
+
+ worldids = jp.arange(batch_size)
+ dx_batch = jax.vmap(functools.partial(tu.make_data, m))(worldids)
+ fields = ('xanchor', 'xaxis', 'xpos', 'xquat', 'xmat', 'xipos', 'ximat',
+ 'geom_xpos', 'geom_xmat', 'site_xpos', 'site_xmat') # fmt: skip
+ for f in fields:
+ dx_batch = dx_batch.replace(**{f: jp.zeros_like(getattr(dx_batch, f))})
+
+ dx_batch = jax.jit(jax.vmap(smooth.kinematics, in_axes=(None, 0)))(
+ mx, dx_batch
+ )
+
+ for i in range(batch_size):
+ dx = dx_batch[i]
+
+ d.qpos[:] = dx.qpos
+ d.mocap_pos[:] = dx.mocap_pos
+ d.mocap_quat[:] = dx.mocap_quat
+ mujoco.mj_forward(m, d)
+
+ tu.assert_attr_eq(d, dx, 'xanchor')
+ tu.assert_attr_eq(d, dx, 'xaxis')
+ tu.assert_attr_eq(d, dx, 'xpos')
+ tu.assert_attr_eq(d, dx, 'xquat')
+ tu.assert_eq(d.xmat.reshape((-1, 3, 3)), dx.xmat, 'xmat')
+ tu.assert_attr_eq(d, dx, 'xipos')
+ tu.assert_eq(d.ximat.reshape((-1, 3, 3)), dx.ximat, 'ximat')
+ tu.assert_attr_eq(d, dx, 'geom_xpos')
+ tu.assert_eq(d.geom_xmat.reshape((-1, 3, 3)), dx.geom_xmat, 'geom_xmat')
+ tu.assert_attr_eq(d, dx, 'site_xpos')
+ tu.assert_eq(d.site_xmat.reshape((-1, 3, 3)), dx.site_xmat, 'site_xmat')
+
+ def test_kinematics_nested_vmap(self):
+ """Tests kinematics with nested batch data."""
+ if not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp not installed.')
+ if not io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device available.')
+
+ m = tu.load_test_file('pendula.xml')
+
+ d = mujoco.MjData(m)
+ mx = mjx.put_model(m, impl='warp')
+
+ worldids = jp.arange(16).reshape((4, 4))
+ dx_batch = jax.vmap(jax.vmap(functools.partial(tu.make_data, m)))(worldids)
+ fields = ('xanchor', 'xaxis', 'xpos', 'xquat', 'xmat', 'xipos', 'ximat',
+ 'geom_xpos', 'geom_xmat', 'site_xpos', 'site_xmat') # fmt: skip
+ for f in fields:
+ dx_batch = dx_batch.replace(**{f: jp.zeros_like(getattr(dx_batch, f))})
+
+ dx_batch = jax.jit(
+ jax.vmap(
+ jax.vmap(smooth.kinematics, in_axes=(None, 0)), in_axes=(None, 0)
+ )
+ )(mx, dx_batch)
+
+ for i in range(4):
+ for j in range(4):
+ dx = dx_batch[i, j]
+
+ d.qpos[:] = dx.qpos
+ d.mocap_pos[:] = dx.mocap_pos
+ d.mocap_quat[:] = dx.mocap_quat
+ mujoco.mj_forward(m, d)
+
+ tu.assert_attr_eq(d, dx, 'xanchor')
+ tu.assert_attr_eq(d, dx, 'xaxis')
+ tu.assert_attr_eq(d, dx, 'xpos')
+ tu.assert_attr_eq(d, dx, 'xquat')
+ tu.assert_eq(d.xmat.reshape((-1, 3, 3)), dx.xmat, 'xmat')
+ tu.assert_attr_eq(d, dx, 'xipos')
+ tu.assert_eq(d.ximat.reshape((-1, 3, 3)), dx.ximat, 'ximat')
+ tu.assert_attr_eq(d, dx, 'geom_xpos')
+ tu.assert_eq(d.geom_xmat.reshape((-1, 3, 3)), dx.geom_xmat, 'geom_xmat')
+ tu.assert_attr_eq(d, dx, 'site_xpos')
+ tu.assert_eq(d.site_xmat.reshape((-1, 3, 3)), dx.site_xmat, 'site_xmat')
+
+ def test_kinematics_model_vmap(self):
+ """Tests kinematics with vmap on model and data fields."""
+ if not mjxw.WARP_INSTALLED:
+ self.skipTest('Warp not installed.')
+ if not io.has_cuda_gpu_device():
+ self.skipTest('No CUDA GPU device available.')
+
+ m = tu.load_test_file('pendula.xml')
+
+ batch_size = 7
+ d = mujoco.MjData(m)
+ mx = mjx.put_model(m, impl='warp')
+ # Add batch dimension to one model field.
+ mx = mx.replace(
+ geom_pos=jax.random.normal(
+ jax.random.PRNGKey(0), (batch_size, m.ngeom, 3)
+ )
+ )
+
+ worldids = jp.arange(batch_size)
+ dx_batch = jax.vmap(functools.partial(tu.make_data, m))(worldids)
+ fields = ('xanchor', 'xaxis', 'xpos', 'xquat', 'xmat', 'xipos', 'ximat',
+ 'geom_xpos', 'geom_xmat', 'site_xpos', 'site_xmat') # fmt: skip
+ for f in fields:
+ dx_batch = dx_batch.replace(**{f: jp.zeros_like(getattr(dx_batch, f))})
+
+ dx_batch = jax.jit(jax.vmap(smooth.kinematics, in_axes=(None, 0)))(
+ mx, dx_batch
+ )
+
+ for i in range(batch_size):
+ dx = dx_batch[i]
+
+ d.qpos[:] = dx.qpos
+ d.mocap_pos[:] = dx.mocap_pos
+ d.mocap_quat[:] = dx.mocap_quat
+ m.geom_pos[:] = mx.geom_pos[i]
+ mujoco.mj_forward(m, d)
+
+ tu.assert_attr_eq(d, dx, 'xanchor')
+ tu.assert_attr_eq(d, dx, 'xaxis')
+ tu.assert_attr_eq(d, dx, 'xpos')
+ tu.assert_attr_eq(d, dx, 'xquat')
+ tu.assert_eq(d.xmat.reshape((-1, 3, 3)), dx.xmat, 'xmat')
+ tu.assert_attr_eq(d, dx, 'xipos')
+ tu.assert_eq(d.ximat.reshape((-1, 3, 3)), dx.ximat, 'ximat')
+ tu.assert_attr_eq(d, dx, 'geom_xpos')
+ tu.assert_eq(d.geom_xmat.reshape((-1, 3, 3)), dx.geom_xmat, 'geom_xmat')
+ 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
new file mode 100644
index 00000000..b5c25f3c
--- /dev/null
+++ b/mjx/mujoco/mjx/warp/test_util.py
@@ -0,0 +1,204 @@
+# Copyright 2025 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.
+# ==============================================================================
+"""Utilities for testing MJX MjWarp integration."""
+import jax
+import jax.numpy as jp
+import mujoco
+from mujoco import mjx
+from mujoco.mjx._src import test_util as mjx_test_util
+import numpy as np
+
+try:
+ from mujoco.mjx.warp import forward as mjxw_forward # pylint: disable=g-import-not-at-top
+except ImportError:
+ mjxw_forward = None
+
+# tolerance for difference between MuJoCo and MJX smooth calculations - mostly
+# due to float precision
+_TOLERANCE = 5e-5
+
+
+def assert_eq(a, b, name):
+ tol = _TOLERANCE * 10 # avoid test noise
+ err_msg = f'mismatch: {name}'
+ np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+
+
+def assert_attr_eq(a, b, attr):
+ assert_eq(getattr(a, attr), getattr(b, attr), attr)
+
+
+def make_data(
+ m: mujoco.MjModel, worldid: int, nconmax: int = 1_000, njmax: int = 100
+):
+ """Make data for a given worldid using keyframes when available."""
+ dx = mjx.make_data(m, impl='warp', nconmax=nconmax, njmax=njmax)
+
+ rng = jax.random.PRNGKey(worldid)
+ rng, key = jax.random.split(rng)
+ qpos = m.qpos0 + jax.random.uniform(key, (m.nq,), minval=-0.05, maxval=0.05)
+ key_qpos = (
+ jp.array(m.key_qpos)[worldid] if m.nkey > 0 else jp.zeros_like(qpos)
+ )
+ qpos = jp.where(worldid < m.nkey, key_qpos, qpos)
+
+ rng, key = jax.random.split(rng)
+ qvel = jax.random.uniform(key, (m.nv,), minval=-0.05, maxval=0.05)
+ key_qvel = (
+ jp.array(m.key_qvel)[worldid] if m.nkey > 0 else jp.zeros_like(qvel)
+ )
+ qvel = jp.where(worldid < m.nkey, key_qvel, qvel)
+
+ rng, key1, key2 = jax.random.split(rng, 3)
+ mocap_pos = jax.random.normal(key1, (m.nmocap, 3))
+ key_mpos = (
+ jp.array(m.key_mpos)[worldid]
+ if m.nkey > 0 and m.nmocap > 0
+ else jp.zeros_like(mocap_pos)
+ )
+ mocap_pos = jp.where(worldid < m.nkey, key_mpos, mocap_pos)
+ mocap_quat = jax.random.normal(key2, (m.nmocap, 4))
+ key_mquat = (
+ jp.array(m.key_mquat)[worldid]
+ if m.nkey > 0 and m.nmocap > 0
+ else jp.zeros_like(mocap_quat)
+ )
+ mocap_quat = jp.where(worldid < m.nkey, key_mquat, mocap_quat)
+
+ _, key = jax.random.split(rng)
+ ctrl = jax.random.uniform(key, (m.nu,), minval=-0.1, maxval=0.1)
+ key_ctrl = (
+ jp.array(m.key_ctrl)[worldid] if m.nkey > 0 else jp.zeros_like(ctrl)
+ )
+ ctrl = jp.where(worldid < m.nkey, key_ctrl, ctrl)
+
+ dx = dx.replace(
+ qpos=qpos,
+ qvel=qvel,
+ mocap_pos=mocap_pos,
+ mocap_quat=mocap_quat,
+ ctrl=ctrl,
+ )
+
+ # Mimic forward call within a vmap trace (since make_data gets called in a
+ # vmap in tests). Some fields get vmapped (qpos, qvel) while others need to be
+ # broadcasted by the custom vmap rule (e.g. geom_xpos).
+ mx = mjx.put_model(m, impl='warp')
+ dx = mjxw_forward.forward(mx, dx)
+
+ return dx
+
+
+def _mjx_contact(dx, worldid: int):
+ keys = []
+ for i in range(dx._impl.ncon[0]):
+ if dx._impl.contact__worldid[i] != worldid:
+ continue
+
+ g1, g2 = tuple(map(int, dx._impl.contact__geom[i]))
+ dist = float(dx._impl.contact__dist[i])
+ keys.append((g1, g2, -dist, i))
+
+ keys = sorted(keys)
+ geom1 = np.array([k[0] for k in keys])
+ geom2 = np.array([k[1] for k in keys])
+ dist = np.array([dx._impl.contact__dist[k[-1]] for k in keys])
+ normal = np.array([dx._impl.contact__frame[k[-1]][0] for k in keys])
+ return geom1, geom2, dist, normal
+
+
+def _mj_contact(d):
+ keys = []
+ for i in range(d.ncon):
+ g1, g2 = tuple(map(int, d.contact.geom[i]))
+ dist = float(d.contact.dist[i])
+ keys.append((g1, g2, -dist, i))
+
+ keys = sorted(keys)
+ geom1 = np.array([k[0] for k in keys])
+ geom2 = np.array([k[1] for k in keys])
+ dist = np.array([d.contact.dist[k[-1]] for k in keys])
+ normal = np.array([d.contact.frame[k[-1]][:3] for k in keys])
+ return geom1, geom2, dist, normal
+
+
+def assert_contact_eq(d, dx, worldid: int):
+ *geom, dist, normal = _mj_contact(d)
+ *geomp, distp, normalp = _mjx_contact(dx, worldid)
+ assert_eq(geomp, geom, 'geom')
+ assert_eq(distp, dist, 'dist')
+ assert_eq(normalp, normal, 'normal')
+
+
+def _mjx_efc(dx, worldid: int):
+ """Gets unpacked efc data for a given worldid."""
+ select = lambda x: x if dx._impl.nefc.ndim == 0 else x[worldid]
+ nefc = select(dx._impl.nefc)
+ keys = np.arange(nefc)
+ if not keys.size:
+ empty = np.array([])
+ 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])
+ keys_sorted = np.lexsort((-efc_pos, efc_type))
+ keys = keys[keys_sorted]
+
+ nefc = len(keys)
+ type_ = efc_type[keys]
+ pos = efc_pos[keys]
+ j = select(dx._impl.efc__J[:nefc])[keys]
+ aref = select(dx._impl.efc__aref[:nefc])[keys]
+ d_ = select(dx._impl.efc__D[:nefc])[keys]
+ return nefc, type_, pos, j, aref, d_
+
+
+def _mj_efc(d):
+ """Gets unpacked efc data."""
+ efc_j = np.zeros((d.efc_J_rownnz.shape[0], d.qvel.shape[0]))
+ if d.efc_J.shape[0] < efc_j.shape[0] * efc_j.shape[1]:
+ mujoco.mju_sparse2dense(
+ efc_j,
+ d.efc_J,
+ d.efc_J_rownnz,
+ d.efc_J_rowadr,
+ d.efc_J_colind,
+ )
+ else:
+ efc_j = d.efc_J.reshape((-1, d.qvel.shape[0]))
+
+ keys = np.lexsort((-d.efc_pos, d.efc_type))
+ type_ = d.efc_type[keys]
+ pos = d.efc_pos[keys]
+ efc_j = efc_j[keys]
+ aref = d.efc_aref[keys]
+ d_ = d.efc_D[keys]
+ return d.nefc, type_, pos, efc_j, aref, d_
+
+
+def assert_efc_eq(d, dx, worldid: int):
+ nefc, type_, pos, j, aref, d_ = _mj_efc(d)
+ nefcp, typep, posp, jp_, arefp, dp = _mjx_efc(dx, worldid)
+
+ assert_eq(nefcp, nefc, 'nefc')
+ assert_eq(typep, type_, 'type')
+ assert_eq(posp, pos, 'pos')
+ assert_eq(jp_, j, 'J')
+ assert_eq(arefp, aref, 'aref')
+ assert_eq(dp, d_, 'D')
+
+
+def load_test_file(name: str) -> mujoco.MjModel:
+ """Loads a mujoco.MjModel based on the file name."""
+ return mjx_test_util.load_test_file(name)
diff --git a/mjx/mujoco/mjx/warp/testspeed.py b/mjx/mujoco/mjx/warp/testspeed.py
new file mode 100644
index 00000000..24802a07
--- /dev/null
+++ b/mjx/mujoco/mjx/warp/testspeed.py
@@ -0,0 +1,427 @@
+# Copyright 2025 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.
+# ==============================================================================
+"""Run benchmarks."""
+
+import functools
+import os
+import time
+from typing import Any, Callable, Sequence, Tuple
+
+from absl import app
+from absl import flags
+import jax
+import mujoco
+from mujoco import mjx
+from mujoco.mjx._src import test_util
+from mujoco.mjx.warp import collision_driver as wp_collision
+from mujoco.mjx.warp import forward as wp_forward
+from mujoco.mjx.warp import smooth as wp_smooth
+import mujoco.mjx.third_party.mujoco_warp as mjwarp
+import warp as wp
+from mujoco.mjx.third_party.warp.jax_experimental import ffi as warp_ffi
+
+_MODELFILE = flags.DEFINE_string(
+ 'modelfile',
+ 'humanoid/humanoid.xml',
+ 'path to model',
+)
+_FUNCTION = flags.DEFINE_string(
+ 'function', 'kinematics', 'function to benchmark'
+)
+_NSTEP = flags.DEFINE_integer('nstep', 1000, 'number of steps per rollout')
+_NENV = flags.DEFINE_integer('nenv', 8192, 'number of environments to simulate')
+_UNROLL = flags.DEFINE_integer('unroll', 4, 'number of steps to unroll')
+_NCONMAX = flags.DEFINE_integer('nconmax', 30_000, 'max contacts')
+_NJMAX = flags.DEFINE_integer('njmax', 80_000, 'max constraints')
+_WP_KERNEL_CACHE_DIR = flags.DEFINE_string(
+ 'wp_kernel_cache_dir',
+ None,
+ 'Path to the Warp kernel cache directory.',
+)
+_COMPILER_OPTIONS = {'xla_gpu_graph_min_graph_size': 1}
+jax_jit = functools.partial(jax.jit, compiler_options=_COMPILER_OPTIONS)
+
+
+def _measure(fn, *args) -> Tuple[float, float]:
+ """Reports jit time and op time for a function."""
+
+ beg = time.perf_counter()
+ compiled_fn = fn.lower(*args).compile()
+ end = time.perf_counter()
+ jit_time = end - beg
+
+ # warmup
+ result = compiled_fn(*args)
+ jax.block_until_ready(result)
+
+ times = []
+
+ for i in range(1):
+ beg = time.perf_counter()
+ result = compiled_fn(*args)
+ jax.block_until_ready(result)
+ end = time.perf_counter()
+ run_time = end - beg
+ times.append(run_time)
+ print('Measure run: ', i, f', run time: {run_time:.3f}')
+
+ return jit_time, sum(times) / len(times)
+
+
+def benchmark(
+ m: mujoco.MjModel,
+ mx: mjx.Model,
+ step_fn: Callable[..., Any],
+ nstep: int = 1000,
+ nenv: int = 8192,
+ unroll_steps: int = 4,
+) -> Tuple[float, float, int]:
+ """Benchmark a model."""
+
+ @jax.vmap
+ def init(key):
+ d = mjx.make_data(
+ m, impl=mx.impl, nconmax=_NCONMAX.value, njmax=_NJMAX.value
+ )
+ return d
+
+ key = jax.random.split(jax.random.key(0), nenv)
+ d = jax_jit(init)(key)
+ jax.block_until_ready(d)
+
+ @jax_jit
+ def unroll(d):
+ def fn(d, _):
+ d = d.replace(qpos=d.qpos + 0 * d.qpos)
+ return step_fn(mx, d), None
+
+ return jax.lax.scan(fn, d, None, length=nstep, unroll=unroll_steps)
+
+ jit_time, run_time = _measure(unroll, d)
+ steps = nstep * nenv
+
+ return jit_time, run_time, steps
+
+
+def benchmark_raw_jax_warp(
+ m: mujoco.MjModel,
+ nstep: int = 1000,
+ nenv: int = 8192,
+ unroll_steps: int = 4,
+ function: str = 'kinematics',
+):
+ if function not in ('kinematics', 'forward', 'step', 'collision'):
+ raise NotImplementedError(
+ f'{function} is not implemented for raw warp speed test.'
+ )
+
+ def warp_fn(
+ qpos_in: wp.array2d(dtype=wp.float32),
+ qpos_out: wp.array2d(dtype=wp.float32),
+ time_: wp.array(dtype=wp.float32),
+ xpos: wp.array2d(dtype=wp.vec3),
+ xquat: wp.array2d(dtype=wp.quat),
+ xmat: wp.array2d(dtype=wp.mat33),
+ xipos: wp.array2d(dtype=wp.vec3),
+ ximat: wp.array2d(dtype=wp.mat33),
+ xanchor: wp.array2d(dtype=wp.vec3),
+ xaxis: wp.array2d(dtype=wp.vec3),
+ geom_xpos: wp.array2d(dtype=wp.vec3),
+ geom_xmat: wp.array2d(dtype=wp.mat33),
+ ):
+ wp.copy(d.qpos, qpos_in)
+ if function == 'kinematics':
+ mjwarp.kinematics(mw, d)
+ elif function == 'forward':
+ mjwarp.forward(mw, d)
+ elif function == 'step':
+ mjwarp.step(mw, d)
+ elif function == 'collision':
+ mjwarp.collision(mw, d)
+ else:
+ raise NotImplementedError(f'{function} not implemented in speed test.')
+ wp.copy(qpos_out, d.qpos)
+ wp.copy(time_, d.time)
+ wp.copy(xpos, d.xpos)
+ wp.copy(xquat, d.xquat)
+ wp.copy(xmat, d.xmat)
+ wp.copy(xipos, d.xipos)
+ wp.copy(ximat, d.ximat)
+ wp.copy(xanchor, d.xanchor)
+ wp.copy(xaxis, d.xaxis)
+ wp.copy(geom_xpos, d.geom_xpos)
+ wp.copy(geom_xmat, d.geom_xmat)
+
+ def unroll(
+ qpos,
+ qpos_out,
+ time_,
+ xpos,
+ xquat,
+ xmat,
+ xipos,
+ ximat,
+ xanchor,
+ xaxis,
+ geom_xpos,
+ geom_xmat,
+ ):
+ def step(carry, _):
+ qpos_, *_ = carry
+ out = warp_fn_jax(qpos_ + 0.0 * qpos_)
+ out = tuple(out)
+ return (out[0],) + out, None
+
+ (
+ qpos,
+ qpos_out,
+ time_,
+ xpos,
+ xquat,
+ xmat,
+ xipos,
+ ximat,
+ xanchor,
+ xaxis,
+ geom_xpos,
+ geom_xmat,
+ ), _ = jax.lax.scan(
+ step,
+ (
+ qpos,
+ qpos_out,
+ time_,
+ xpos,
+ xquat,
+ xmat,
+ xipos,
+ ximat,
+ xanchor,
+ xaxis,
+ geom_xpos,
+ geom_xmat,
+ ),
+ length=nstep,
+ unroll=unroll_steps,
+ )
+
+ return (
+ qpos,
+ qpos_out,
+ time_,
+ xpos,
+ xquat,
+ xmat,
+ xipos,
+ ximat,
+ xanchor,
+ xaxis,
+ geom_xpos,
+ geom_xmat,
+ )
+
+ output_dims = {
+ 'time_': (nenv,),
+ 'qpos_out': (nenv, m.nq),
+ 'xpos': (nenv, m.nbody, 3),
+ 'xquat': (nenv, m.nbody, 4),
+ 'xmat': (nenv, m.nbody, 3, 3),
+ 'xipos': (nenv, m.nbody, 3),
+ 'ximat': (nenv, m.nbody, 3, 3),
+ 'xanchor': (nenv, m.njnt, 3),
+ 'xaxis': (nenv, m.njnt, 3),
+ 'geom_xpos': (nenv, m.ngeom, 3),
+ 'geom_xmat': (nenv, m.ngeom, 3, 3),
+ }
+ warp_fn_jax = warp_ffi.jax_callable(
+ warp_fn,
+ num_outputs=11,
+ output_dims=output_dims,
+ )
+
+ @jax.vmap
+ def init(key):
+ d = mjx.make_data(m, impl='jax')
+ return d
+
+ key = jax.random.split(jax.random.key(0), nenv)
+ dx = jax_jit(init)(key)
+ d_ = mujoco.MjData(m)
+ mw = mjwarp.put_model(m)
+ mw.opt.graph_conditional = False
+ d = mjwarp.put_data(
+ m, d_, nworld=nenv, nconmax=_NCONMAX.value, njmax=_NJMAX.value
+ )
+
+ jax_unroll_fn = jax_jit(unroll)
+ jit_time, run_time = _measure(
+ jax_unroll_fn,
+ dx.qpos,
+ dx.qpos,
+ dx.time,
+ dx.xpos,
+ dx.xquat,
+ dx.xmat,
+ dx.xipos,
+ dx.ximat,
+ dx.xanchor,
+ dx.xaxis,
+ dx.geom_xpos,
+ dx.geom_xmat,
+ )
+ steps = nstep * nenv
+
+ return jit_time, run_time, steps
+
+
+def _compile_fn(fn, m, d):
+ fn(m, d)
+ fn(m, d)
+ with wp.ScopedCapture() as capture:
+ fn(m, d)
+ return capture.graph
+
+
+def benchmark_raw_warp(
+ m: mujoco.MjModel,
+ nstep: int = 1000,
+ nenv: int = 8192,
+ unroll_steps: int = 4,
+ function: str = 'kinematics',
+):
+ """Benchmarks raw warp."""
+ del unroll_steps
+ if function not in ('kinematics', 'forward', 'step', 'collision'):
+ raise NotImplementedError(
+ f'{function} is not implemented for raw warp speed test.'
+ )
+
+ mw = mjwarp.put_model(m)
+ # TODO(btaba): re-enable graph conditional once JAX supports it, for fair
+ # comparison.
+ mw.opt.graph_conditional = False
+ dw = mjwarp.make_data(
+ m, nworld=nenv, nconmax=_NCONMAX.value, njmax=_NJMAX.value
+ )
+
+ if function == 'kinematics':
+ fn = mjwarp.kinematics
+ elif function == 'forward':
+ fn = mjwarp.forward
+ elif function == 'step':
+ fn = mjwarp.step
+ elif function == 'collision':
+ fn = mjwarp.collision
+ else:
+ raise NotImplementedError(f'{function} not implemented in speed test.')
+
+ start = time.time()
+ graph = _compile_fn(fn, mw, dw)
+ jit_time = time.time() - start
+
+ start = time.time()
+ for _ in range(nstep):
+ wp.capture_launch(graph)
+ wp.synchronize()
+
+ run_time = time.time() - start
+ return jit_time, run_time, nstep * nenv
+
+
+def _main(_: Sequence[str]):
+ """Runs testpeed function."""
+ os.environ['MJX_WARP_ENABLED'] = 'true'
+
+ if _WP_KERNEL_CACHE_DIR.value:
+ wp.config.kernel_cache_dir = _WP_KERNEL_CACHE_DIR.value
+
+ modelfile = _MODELFILE.value
+ function_ = _FUNCTION.value
+ nstep, nenv, unroll = _NSTEP.value, _NENV.value, _UNROLL.value
+
+ try:
+ m = test_util.load_test_file(modelfile)
+ except Exception as _:
+ m = mujoco.MjModel.from_xml_path(modelfile)
+
+ mx = mjx.put_model(m, impl='jax')
+ mw = mjx.put_model(m, impl='warp')
+
+ if function_ == 'kinematics':
+ func_warp = jax.vmap(wp_smooth.kinematics, in_axes=(None, 0))
+ func_jax = jax.vmap(mjx.kinematics, in_axes=(None, 0))
+ elif function_ == 'forward':
+ func_warp = jax.vmap(wp_forward.forward, in_axes=(None, 0))
+ func_jax = jax.vmap(mjx.forward, in_axes=(None, 0))
+ elif function_ == 'step':
+ func_warp = jax.vmap(wp_forward.step, in_axes=(None, 0))
+ func_jax = jax.vmap(mjx.step, in_axes=(None, 0))
+ elif function_ == 'collision':
+ func_warp = jax.vmap(wp_collision.collision, in_axes=(None, 0))
+ func_jax = jax.vmap(mjx.collision, in_axes=(None, 0))
+ else:
+ raise ValueError(f'Unknown function: {function_}')
+
+ print('testspeed.py:\n')
+ print(f' modelfile : {modelfile}')
+ print(f' function : {function_}')
+ print(f' nenv : {nenv}')
+ print(f' nstep : {nstep}')
+ print(f' timestep : {m.opt.timestep}')
+ print(f' unroll : {unroll}\n')
+
+ for name, mx_, op in (
+ ('JAX WARP FFI', mw, func_warp),
+ ('Pure JAX', mx, func_jax),
+ ):
+ if op is not None:
+ jit_time, run_time, steps = benchmark(m, mx_, op, nstep, nenv, unroll)
+
+ print(f' {name}:')
+ print(f' JIT time : {jit_time:.2f} s')
+ print(f' simulation time : {run_time:.2f} s')
+ print(f' steps per second : {steps / run_time:,.0f}')
+ print(
+ f' realtime factor : {steps * m.opt.timestep / run_time:.2f} x'
+ )
+ print(f' time per step : {1e6 * run_time / steps:.2f} µs\n')
+
+ jit_time, run_time, steps = benchmark_raw_jax_warp(
+ m, nstep, nenv, unroll, function=function_
+ )
+ print(' Pure JAX-WARP:')
+ print(f' JIT time : {jit_time:.2f} s')
+ print(f' simulation time : {run_time:.2f} s')
+ print(f' steps per second : {steps / run_time:,.0f}')
+ print(f' realtime factor : {steps * m.opt.timestep / run_time:.2f} x')
+ print(f' time per step : {1e6 * run_time / steps:.2f} µs\n')
+
+ jit_time, run_time, steps = benchmark_raw_warp(
+ m, nstep, nenv, unroll, function=function_
+ )
+ print(' Pure WARP:')
+ print(f' JIT time : {jit_time:.2f} s')
+ print(f' simulation time : {run_time:.2f} s')
+ print(f' steps per second : {steps / run_time:,.0f}')
+ print(f' realtime factor : {steps * m.opt.timestep / run_time:.2f} x')
+ print(f' time per step : {1e6 * run_time / steps:.2f} µs\n')
+
+
+def main():
+ app.run(_main)
+
+
+if __name__ == '__main__':
+ main()
diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py
new file mode 100644
index 00000000..17cab394
--- /dev/null
+++ b/mjx/mujoco/mjx/warp/types.py
@@ -0,0 +1,1602 @@
+# Copyright 2025 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.
+# ==============================================================================
+"""MJX Warp types.
+
+DO NOT EDIT. This file is auto-generated.
+"""
+import dataclasses
+from typing import Tuple
+import jax
+from jax import tree_util
+from jax.interpreters import batching
+from mujoco.mjx._src import dataclasses as mjx_dataclasses
+import numpy as np
+
+PyTreeNode = mjx_dataclasses.PyTreeNode
+
+
+@dataclasses.dataclass(frozen=True)
+@tree_util.register_pytree_node_class
+class TileSet:
+ """Tiling configuration for decomposable block diagonal matrix.
+
+ For non-square, non-block-diagonal tiles, use two tilesets.
+
+ Attributes:
+ adr: address of each tile in the set
+ size: size of all the tiles in this set
+ """
+
+ adr: np.ndarray
+ size: int
+
+ def tree_flatten(self):
+ children = list((getattr(self, k) for k in self.__dataclass_fields__))
+ return (children, None)
+
+ @classmethod
+ def tree_unflatten(cls, aux_data, children):
+ del aux_data
+ return cls(*children)
+
+
+@dataclasses.dataclass(frozen=True)
+@tree_util.register_pytree_node_class
+class BlockDim:
+ """Block dimension 'block_dim' settings for wp.launch_tiled.
+
+ TODO(team): experimental and may be removed
+ """
+
+ actuator_velocity: int
+ cholesky_factorize: int
+ cholesky_factorize_solve: int
+ cholesky_solve: int
+ energy_vel_kinetic: int
+ euler_dense: int
+ mul_m_dense: int
+ qderiv_actuator_passive_actuation: int
+ qderiv_actuator_passive_no_actuation: int
+ qfrc_actuator: int
+ ray: int
+ segmented_sort: int
+ tendon_velocity: int
+ update_gradient_cholesky: int
+
+ def tree_flatten(self):
+ children = list((getattr(self, k) for k in self.__dataclass_fields__))
+ return (children, None)
+
+ @classmethod
+ def tree_unflatten(cls, aux_data, children):
+ del aux_data
+ return cls(*children)
+
+
+class StatisticWarp(PyTreeNode):
+ """Derived fields from Statistic."""
+
+ meaninertia: float
+
+
+class OptionWarp(PyTreeNode):
+ """Derived fields from Option."""
+
+ broadphase: int
+ broadphase_filter: int
+ epa_iterations: int
+ gjk_iterations: int
+ graph_conditional: bool
+ has_fluid: bool
+ is_sparse: bool
+ ls_parallel: bool
+ run_collision_detection: bool
+ sdf_initpoints: int
+ sdf_iterations: int
+
+
+class ModelWarp(PyTreeNode):
+ """Derived fields from Model."""
+
+ M_colind: np.ndarray
+ M_rowadr: np.ndarray
+ M_rownnz: np.ndarray
+ actuator_affine_bias_gain: bool
+ actuator_moment_tiles_nu: Tuple[TileSet, ...]
+ actuator_moment_tiles_nv: Tuple[TileSet, ...]
+ actuator_trntype_body_adr: np.ndarray
+ block_dim: BlockDim
+ body_tree: Tuple[np.ndarray, ...]
+ condim_max: int
+ dof_tri_col: np.ndarray
+ dof_tri_row: np.ndarray
+ eq_connect_adr: np.ndarray
+ eq_jnt_adr: np.ndarray
+ eq_ten_adr: np.ndarray
+ eq_wld_adr: np.ndarray
+ flex_bending: np.ndarray
+ flex_damping: np.ndarray
+ flex_dim: np.ndarray
+ flex_edge: np.ndarray
+ flex_edgeadr: np.ndarray
+ flex_edgeflap: np.ndarray
+ flex_elem: np.ndarray
+ flex_elemedge: np.ndarray
+ flex_elemedgeadr: np.ndarray
+ flex_stiffness: np.ndarray
+ flex_vertadr: np.ndarray
+ flex_vertbodyid: np.ndarray
+ flex_vertnum: np.ndarray
+ flexedge_length0: np.ndarray
+ geom_pair_type_count: Tuple[int, ...]
+ geom_plugin_index: np.ndarray
+ geompair2hfgeompair: np.ndarray
+ has_sdf_geom: bool
+ jnt_limited_ball_adr: np.ndarray
+ jnt_limited_slide_hinge_adr: np.ndarray
+ light_bodyid: np.ndarray
+ light_targetbodyid: np.ndarray
+ mapM2M: np.ndarray
+ mesh_polyadr: np.ndarray
+ mesh_polymap: np.ndarray
+ mesh_polymapadr: np.ndarray
+ mesh_polymapnum: np.ndarray
+ mesh_polynormal: np.ndarray
+ mesh_polynum: np.ndarray
+ mesh_polyvert: np.ndarray
+ mesh_polyvertadr: np.ndarray
+ mesh_polyvertnum: np.ndarray
+ mocap_bodyid: np.ndarray
+ nflex: int
+ nflexedge: int
+ nflexelem: int
+ nflexelemdata: int
+ nflexvert: int
+ nlsp: int
+ nmeshpoly: int
+ nmeshpolymap: int
+ nmeshpolyvert: int
+ nxn_geom_pair: np.ndarray
+ nxn_geom_pair_filtered: np.ndarray
+ nxn_pairid: np.ndarray
+ nxn_pairid_filtered: np.ndarray
+ plugin: np.ndarray
+ plugin_attr: np.ndarray
+ qLD_updates: Tuple[np.ndarray, ...]
+ qM_fullm_i: np.ndarray
+ qM_fullm_j: np.ndarray
+ qM_madr_ij: np.ndarray
+ qM_mulm_i: np.ndarray
+ qM_mulm_j: np.ndarray
+ qM_tiles: Tuple[TileSet, ...]
+ rangefinder_sensor_adr: np.ndarray
+ sensor_acc_adr: np.ndarray
+ sensor_e_kinetic: bool
+ sensor_e_potential: bool
+ sensor_limitfrc_adr: np.ndarray
+ sensor_limitpos_adr: np.ndarray
+ sensor_limitvel_adr: np.ndarray
+ sensor_pos_adr: np.ndarray
+ sensor_rangefinder_adr: np.ndarray
+ sensor_rangefinder_bodyid: np.ndarray
+ sensor_rne_postconstraint: bool
+ sensor_subtree_vel: bool
+ sensor_tendonactfrc_adr: np.ndarray
+ sensor_touch_adr: np.ndarray
+ sensor_vel_adr: np.ndarray
+ subtree_mass: jax.Array
+ ten_wrapadr_site: np.ndarray
+ ten_wrapnum_site: np.ndarray
+ tendon_geom_adr: np.ndarray
+ tendon_jnt_adr: np.ndarray
+ tendon_limited_adr: np.ndarray
+ tendon_site_pair_adr: np.ndarray
+ wrap_geom_adr: np.ndarray
+ wrap_jnt_adr: np.ndarray
+ wrap_pulley_scale: np.ndarray
+ wrap_site_adr: np.ndarray
+ wrap_site_pair_adr: np.ndarray
+
+
+class DataWarp(PyTreeNode):
+ """Derived fields from Data."""
+
+ act_dot_rk: jax.Array
+ act_t0: jax.Array
+ act_vel_integration: jax.Array
+ actuator_length: jax.Array
+ actuator_moment: jax.Array
+ actuator_trntype_body_ncon: jax.Array
+ actuator_velocity: jax.Array
+ cacc: jax.Array
+ cdof: jax.Array
+ cdof_dot: jax.Array
+ cfrc_ext: jax.Array
+ cfrc_int: jax.Array
+ cinert: jax.Array
+ collision_hftri_index: jax.Array
+ collision_pair: jax.Array
+ collision_pairid: jax.Array
+ collision_worldid: jax.Array
+ contact__dim: jax.Array
+ contact__dist: jax.Array
+ contact__efc_address: jax.Array
+ contact__frame: jax.Array
+ contact__friction: jax.Array
+ contact__geom: jax.Array
+ contact__includemargin: jax.Array
+ contact__pos: jax.Array
+ contact__solimp: jax.Array
+ contact__solref: jax.Array
+ contact__solreffriction: jax.Array
+ contact__worldid: jax.Array
+ crb: jax.Array
+ efc__D: jax.Array
+ efc__J: jax.Array
+ efc__Jaref: jax.Array
+ efc__Ma: jax.Array
+ efc__Mgrad: jax.Array
+ efc__active: jax.Array
+ efc__alpha: jax.Array
+ efc__aref: jax.Array
+ efc__beta: jax.Array
+ efc__beta_den: jax.Array
+ efc__beta_num: jax.Array
+ efc__cholesky_L_tmp: jax.Array
+ efc__cholesky_y_tmp: jax.Array
+ efc__condim: jax.Array
+ efc__cost: jax.Array
+ efc__cost_candidate: jax.Array
+ efc__done: jax.Array
+ efc__force: jax.Array
+ efc__frictionloss: jax.Array
+ efc__gauss: jax.Array
+ efc__grad: jax.Array
+ efc__grad_dot: jax.Array
+ efc__gtol: jax.Array
+ efc__h: jax.Array
+ efc__hi: jax.Array
+ efc__hi_alpha: jax.Array
+ efc__hi_next: jax.Array
+ efc__hi_next_alpha: jax.Array
+ efc__id: jax.Array
+ efc__jv: jax.Array
+ efc__lo: jax.Array
+ efc__lo_alpha: jax.Array
+ efc__lo_next: jax.Array
+ efc__lo_next_alpha: jax.Array
+ efc__ls_done: jax.Array
+ efc__margin: jax.Array
+ efc__mid: jax.Array
+ efc__mid_alpha: jax.Array
+ efc__mv: jax.Array
+ efc__p0: jax.Array
+ efc__pos: jax.Array
+ efc__prev_Mgrad: jax.Array
+ efc__prev_cost: jax.Array
+ efc__prev_grad: jax.Array
+ efc__quad: jax.Array
+ efc__quad_gauss: jax.Array
+ efc__search: jax.Array
+ efc__search_dot: jax.Array
+ efc__type: jax.Array
+ efc__u: jax.Array
+ efc__uu: jax.Array
+ efc__uv: jax.Array
+ efc__vel: jax.Array
+ efc__vv: jax.Array
+ energy: jax.Array
+ energy_vel_mul_m_skip: jax.Array
+ epa_face: jax.Array
+ epa_horizon: jax.Array
+ epa_index: jax.Array
+ epa_map: jax.Array
+ epa_norm2: jax.Array
+ epa_pr: jax.Array
+ epa_vert: jax.Array
+ epa_vert1: jax.Array
+ epa_vert2: jax.Array
+ epa_vert_index1: jax.Array
+ epa_vert_index2: jax.Array
+ flexedge_length: jax.Array
+ flexedge_velocity: jax.Array
+ flexvert_xpos: jax.Array
+ fluid_applied: jax.Array
+ geom_skip: jax.Array
+ inverse_mul_m_skip: jax.Array
+ light_xdir: jax.Array
+ light_xpos: jax.Array
+ ncollision: jax.Array
+ ncon: jax.Array
+ ncon_hfield: jax.Array
+ nconmax: int
+ ne: jax.Array
+ ne_connect: jax.Array
+ ne_jnt: jax.Array
+ ne_ten: jax.Array
+ ne_weld: jax.Array
+ nefc: jax.Array
+ nf: jax.Array
+ njmax: int
+ nl: jax.Array
+ nsolving: jax.Array
+ nworld: int
+ qLD: jax.Array
+ qLD_integration: jax.Array
+ qLDiagInv: jax.Array
+ qLDiagInv_integration: jax.Array
+ qM: jax.Array
+ qM_integration: jax.Array
+ qacc_discrete: jax.Array
+ qacc_integration: jax.Array
+ qacc_rk: jax.Array
+ qfrc_damper: jax.Array
+ qfrc_integration: jax.Array
+ qfrc_spring: jax.Array
+ qpos_t0: jax.Array
+ qvel_rk: jax.Array
+ qvel_t0: jax.Array
+ ray_bodyexclude: jax.Array
+ ray_dist: jax.Array
+ ray_geomid: jax.Array
+ sap_cumulative_sum: jax.Array
+ sap_projection_lower: jax.Array
+ sap_projection_upper: jax.Array
+ sap_range: jax.Array
+ sap_segment_index: jax.Array
+ sap_sort_index: jax.Array
+ sensor_rangefinder_dist: jax.Array
+ sensor_rangefinder_geomid: jax.Array
+ sensor_rangefinder_pnt: jax.Array
+ sensor_rangefinder_vec: jax.Array
+ solver_niter: jax.Array
+ subtree_angmom: jax.Array
+ subtree_bodyvel: jax.Array
+ subtree_linvel: jax.Array
+ ten_J: jax.Array
+ 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
+ wrap_geom_xpos: jax.Array
+ wrap_obj: jax.Array
+ wrap_xpos: jax.Array
+ shape = property(lambda self: self.cacc.shape)
+
+
+DATA_NON_VMAP = {
+ 'collision_hftri_index',
+ 'collision_pair',
+ 'collision_pairid',
+ 'collision_worldid',
+ 'contact__dim',
+ 'contact__dist',
+ 'contact__efc_address',
+ 'contact__frame',
+ 'contact__friction',
+ 'contact__geom',
+ 'contact__includemargin',
+ 'contact__pos',
+ 'contact__solimp',
+ 'contact__solref',
+ 'contact__solreffriction',
+ 'contact__worldid',
+ 'efc__u',
+ 'efc__uu',
+ 'efc__uv',
+ 'efc__vv',
+ 'epa_face',
+ 'epa_horizon',
+ 'epa_index',
+ 'epa_map',
+ 'epa_norm2',
+ 'epa_pr',
+ 'epa_vert',
+ 'epa_vert1',
+ 'epa_vert2',
+ 'epa_vert_index1',
+ 'epa_vert_index2',
+ 'geom_skip',
+ 'ncollision',
+ 'ncon',
+ 'nconmax',
+ 'njmax',
+ 'nsolving',
+ 'nworld',
+ 'ray_bodyexclude',
+}
+
+
+def _to_elt(cont, _, d, axis):
+ return DataWarp(**{
+ f.name: (
+ cont(getattr(d, f.name), axis)
+ if f.name not in DATA_NON_VMAP
+ else getattr(d, f.name)
+ )
+ for f in DataWarp.fields()
+ })
+
+
+def _from_elt(cont, axis_size, d, axis_dest):
+ return DataWarp(**{
+ f.name: (
+ cont(axis_size, getattr(d, f.name), axis_dest)
+ if f.name not in DATA_NON_VMAP
+ else getattr(d, f.name)
+ )
+ for f in DataWarp.fields()
+ })
+
+
+batching.register_vmappable(DataWarp, int, int, _to_elt, _from_elt, None)
+
+NDIM = {
+ 'Data': {
+ 'act': 2,
+ 'act_dot': 2,
+ 'act_dot_rk': 2,
+ 'act_t0': 2,
+ 'act_vel_integration': 2,
+ 'actuator_force': 2,
+ 'actuator_length': 2,
+ 'actuator_moment': 3,
+ 'actuator_trntype_body_ncon': 2,
+ 'actuator_velocity': 2,
+ 'cacc': 3,
+ 'cam_xmat': 4,
+ 'cam_xpos': 3,
+ 'cdof': 3,
+ 'cdof_dot': 3,
+ 'cfrc_ext': 3,
+ 'cfrc_int': 3,
+ 'cinert': 3,
+ 'collision_hftri_index': 1,
+ 'collision_pair': 2,
+ 'collision_pairid': 1,
+ 'collision_worldid': 1,
+ 'contact__dim': 1,
+ 'contact__dist': 1,
+ 'contact__efc_address': 2,
+ 'contact__frame': 3,
+ 'contact__friction': 2,
+ 'contact__geom': 2,
+ 'contact__includemargin': 1,
+ 'contact__pos': 2,
+ 'contact__solimp': 2,
+ 'contact__solref': 2,
+ 'contact__solreffriction': 2,
+ 'contact__worldid': 1,
+ 'crb': 3,
+ 'ctrl': 2,
+ 'cvel': 3,
+ 'efc__D': 2,
+ 'efc__J': 3,
+ 'efc__Jaref': 2,
+ 'efc__Ma': 2,
+ 'efc__Mgrad': 2,
+ 'efc__active': 2,
+ 'efc__alpha': 1,
+ 'efc__aref': 2,
+ 'efc__beta': 1,
+ 'efc__beta_den': 1,
+ 'efc__beta_num': 1,
+ 'efc__cholesky_L_tmp': 3,
+ 'efc__cholesky_y_tmp': 2,
+ 'efc__condim': 2,
+ 'efc__cost': 1,
+ 'efc__cost_candidate': 2,
+ 'efc__done': 1,
+ 'efc__force': 2,
+ 'efc__frictionloss': 2,
+ 'efc__gauss': 1,
+ 'efc__grad': 2,
+ 'efc__grad_dot': 1,
+ 'efc__gtol': 1,
+ 'efc__h': 3,
+ 'efc__hi': 2,
+ 'efc__hi_alpha': 1,
+ 'efc__hi_next': 2,
+ 'efc__hi_next_alpha': 1,
+ 'efc__id': 2,
+ 'efc__jv': 2,
+ 'efc__lo': 2,
+ 'efc__lo_alpha': 1,
+ 'efc__lo_next': 2,
+ 'efc__lo_next_alpha': 1,
+ 'efc__ls_done': 1,
+ 'efc__margin': 2,
+ 'efc__mid': 2,
+ 'efc__mid_alpha': 1,
+ 'efc__mv': 2,
+ 'efc__p0': 2,
+ 'efc__pos': 2,
+ 'efc__prev_Mgrad': 2,
+ 'efc__prev_cost': 1,
+ 'efc__prev_grad': 2,
+ 'efc__quad': 3,
+ 'efc__quad_gauss': 2,
+ 'efc__search': 2,
+ 'efc__search_dot': 1,
+ 'efc__type': 2,
+ 'efc__u': 2,
+ 'efc__uu': 1,
+ 'efc__uv': 1,
+ 'efc__vel': 2,
+ 'efc__vv': 1,
+ 'energy': 2,
+ 'energy_vel_mul_m_skip': 1,
+ 'epa_face': 3,
+ 'epa_horizon': 2,
+ 'epa_index': 2,
+ 'epa_map': 2,
+ 'epa_norm2': 2,
+ 'epa_pr': 3,
+ 'epa_vert': 3,
+ 'epa_vert1': 3,
+ 'epa_vert2': 3,
+ 'epa_vert_index1': 2,
+ 'epa_vert_index2': 2,
+ 'eq_active': 2,
+ 'flexedge_length': 2,
+ 'flexedge_velocity': 2,
+ 'flexvert_xpos': 3,
+ 'fluid_applied': 3,
+ 'geom_skip': 1,
+ 'geom_xmat': 4,
+ 'geom_xpos': 3,
+ 'inverse_mul_m_skip': 1,
+ 'light_xdir': 3,
+ 'light_xpos': 3,
+ 'mocap_pos': 3,
+ 'mocap_quat': 3,
+ 'ncollision': 1,
+ 'ncon': 1,
+ 'ncon_hfield': 2,
+ 'nconmax': 0,
+ 'ne': 1,
+ 'ne_connect': 1,
+ 'ne_jnt': 1,
+ 'ne_ten': 1,
+ 'ne_weld': 1,
+ 'nefc': 1,
+ 'nf': 1,
+ 'njmax': 0,
+ 'nl': 1,
+ 'nsolving': 1,
+ 'nworld': 0,
+ 'qLD': 3,
+ 'qLD_integration': 3,
+ 'qLDiagInv': 2,
+ 'qLDiagInv_integration': 2,
+ 'qM': 3,
+ 'qM_integration': 3,
+ 'qacc': 2,
+ 'qacc_discrete': 2,
+ 'qacc_integration': 2,
+ 'qacc_rk': 2,
+ 'qacc_smooth': 2,
+ 'qacc_warmstart': 2,
+ 'qfrc_actuator': 2,
+ 'qfrc_applied': 2,
+ 'qfrc_bias': 2,
+ 'qfrc_constraint': 2,
+ 'qfrc_damper': 2,
+ 'qfrc_fluid': 2,
+ 'qfrc_gravcomp': 2,
+ 'qfrc_integration': 2,
+ 'qfrc_inverse': 2,
+ 'qfrc_passive': 2,
+ 'qfrc_smooth': 2,
+ 'qfrc_spring': 2,
+ 'qpos': 2,
+ 'qpos_t0': 2,
+ 'qvel': 2,
+ 'qvel_rk': 2,
+ 'qvel_t0': 2,
+ 'ray_bodyexclude': 1,
+ 'ray_dist': 2,
+ 'ray_geomid': 2,
+ 'sap_cumulative_sum': 2,
+ 'sap_projection_lower': 3,
+ 'sap_projection_upper': 2,
+ 'sap_range': 2,
+ 'sap_segment_index': 2,
+ 'sap_sort_index': 3,
+ 'sensor_rangefinder_dist': 2,
+ 'sensor_rangefinder_geomid': 2,
+ 'sensor_rangefinder_pnt': 3,
+ 'sensor_rangefinder_vec': 3,
+ 'sensordata': 2,
+ 'site_xmat': 4,
+ 'site_xpos': 3,
+ 'solver_niter': 1,
+ 'subtree_angmom': 3,
+ 'subtree_bodyvel': 3,
+ 'subtree_com': 3,
+ 'subtree_linvel': 3,
+ 'ten_J': 3,
+ 'ten_Jdot': 3,
+ 'ten_actfrc': 2,
+ 'ten_bias_coef': 2,
+ 'ten_length': 2,
+ 'ten_velocity': 2,
+ 'ten_wrapadr': 2,
+ 'ten_wrapnum': 2,
+ 'time': 1,
+ 'wrap_geom_xpos': 3,
+ 'wrap_obj': 3,
+ 'wrap_xpos': 3,
+ 'xanchor': 3,
+ 'xaxis': 3,
+ 'xfrc_applied': 3,
+ 'ximat': 4,
+ 'xipos': 3,
+ 'xmat': 4,
+ 'xpos': 3,
+ 'xquat': 3,
+ },
+ 'Model': {
+ 'M_colind': 1,
+ 'M_rowadr': 1,
+ 'M_rownnz': 1,
+ 'actuator_acc0': 1,
+ 'actuator_actadr': 1,
+ 'actuator_actearly': 1,
+ 'actuator_actlimited': 1,
+ 'actuator_actnum': 1,
+ 'actuator_actrange': 3,
+ 'actuator_affine_bias_gain': 0,
+ 'actuator_biasprm': 3,
+ 'actuator_biastype': 1,
+ 'actuator_cranklength': 1,
+ 'actuator_ctrllimited': 1,
+ 'actuator_ctrlrange': 3,
+ 'actuator_dynprm': 3,
+ 'actuator_dyntype': 1,
+ 'actuator_forcelimited': 1,
+ 'actuator_forcerange': 3,
+ 'actuator_gainprm': 3,
+ 'actuator_gaintype': 1,
+ 'actuator_gear': 3,
+ 'actuator_lengthrange': 2,
+ 'actuator_moment_tiles_nu': -1,
+ 'actuator_moment_tiles_nv': -1,
+ 'actuator_trnid': 2,
+ 'actuator_trntype': 1,
+ 'actuator_trntype_body_adr': 1,
+ 'block_dim__actuator_velocity': 0,
+ 'block_dim__cholesky_factorize': 0,
+ 'block_dim__cholesky_factorize_solve': 0,
+ 'block_dim__cholesky_solve': 0,
+ 'block_dim__energy_vel_kinetic': 0,
+ 'block_dim__euler_dense': 0,
+ 'block_dim__mul_m_dense': 0,
+ 'block_dim__qderiv_actuator_passive_actuation': 0,
+ 'block_dim__qderiv_actuator_passive_no_actuation': 0,
+ 'block_dim__qfrc_actuator': 0,
+ 'block_dim__ray': 0,
+ 'block_dim__segmented_sort': 0,
+ 'block_dim__tendon_velocity': 0,
+ 'block_dim__update_gradient_cholesky': 0,
+ 'body_conaffinity': 1,
+ 'body_contype': 1,
+ 'body_dofadr': 1,
+ 'body_dofnum': 1,
+ 'body_geomadr': 1,
+ 'body_geomnum': 1,
+ 'body_gravcomp': 2,
+ 'body_inertia': 3,
+ 'body_invweight0': 3,
+ 'body_ipos': 3,
+ 'body_iquat': 3,
+ 'body_jntadr': 1,
+ 'body_jntnum': 1,
+ 'body_mass': 2,
+ 'body_mocapid': 1,
+ 'body_parentid': 1,
+ 'body_pos': 3,
+ 'body_quat': 3,
+ 'body_rootid': 1,
+ 'body_subtreemass': 2,
+ 'body_tree': -1,
+ 'body_weldid': 1,
+ 'cam_bodyid': 1,
+ 'cam_fovy': 1,
+ 'cam_intrinsic': 2,
+ 'cam_mat0': 4,
+ 'cam_mode': 1,
+ 'cam_pos': 3,
+ 'cam_pos0': 3,
+ 'cam_poscom0': 3,
+ 'cam_quat': 3,
+ 'cam_resolution': 2,
+ 'cam_sensorsize': 2,
+ 'cam_targetbodyid': 1,
+ 'condim_max': 0,
+ 'dof_Madr': 1,
+ 'dof_armature': 2,
+ 'dof_bodyid': 1,
+ 'dof_damping': 2,
+ 'dof_frictionloss': 2,
+ 'dof_invweight0': 2,
+ 'dof_jntid': 1,
+ 'dof_parentid': 1,
+ 'dof_solimp': 3,
+ 'dof_solref': 3,
+ 'dof_tri_col': 1,
+ 'dof_tri_row': 1,
+ 'eq_active0': 1,
+ 'eq_connect_adr': 1,
+ 'eq_data': 3,
+ 'eq_jnt_adr': 1,
+ 'eq_obj1id': 1,
+ 'eq_obj2id': 1,
+ 'eq_objtype': 1,
+ 'eq_solimp': 3,
+ 'eq_solref': 3,
+ 'eq_ten_adr': 1,
+ 'eq_type': 1,
+ 'eq_wld_adr': 1,
+ 'exclude_signature': 1,
+ 'flex_bending': 3,
+ 'flex_damping': 1,
+ 'flex_dim': 1,
+ 'flex_edge': 2,
+ 'flex_edgeadr': 1,
+ 'flex_edgeflap': 2,
+ 'flex_elem': 1,
+ 'flex_elemedge': 1,
+ 'flex_elemedgeadr': 1,
+ 'flex_stiffness': 1,
+ 'flex_vertadr': 1,
+ 'flex_vertbodyid': 1,
+ 'flex_vertnum': 1,
+ 'flexedge_length0': 1,
+ 'geom_aabb': 3,
+ 'geom_bodyid': 1,
+ 'geom_conaffinity': 1,
+ 'geom_condim': 1,
+ 'geom_contype': 1,
+ 'geom_dataid': 1,
+ 'geom_friction': 3,
+ 'geom_gap': 2,
+ 'geom_group': 1,
+ 'geom_margin': 2,
+ 'geom_matid': 2,
+ 'geom_pair_type_count': -1,
+ 'geom_plugin_index': 1,
+ 'geom_pos': 3,
+ 'geom_priority': 1,
+ 'geom_quat': 3,
+ 'geom_rbound': 2,
+ 'geom_rgba': 3,
+ 'geom_size': 3,
+ 'geom_solimp': 3,
+ 'geom_solmix': 2,
+ 'geom_solref': 3,
+ 'geom_type': 1,
+ 'geompair2hfgeompair': 1,
+ 'has_sdf_geom': 0,
+ 'hfield_adr': 1,
+ 'hfield_data': 1,
+ 'hfield_ncol': 1,
+ 'hfield_nrow': 1,
+ 'hfield_size': 2,
+ 'jnt_actfrclimited': 1,
+ 'jnt_actfrcrange': 3,
+ 'jnt_actgravcomp': 1,
+ 'jnt_axis': 3,
+ 'jnt_bodyid': 1,
+ 'jnt_dofadr': 1,
+ 'jnt_limited': 1,
+ 'jnt_limited_ball_adr': 1,
+ 'jnt_limited_slide_hinge_adr': 1,
+ 'jnt_margin': 2,
+ 'jnt_pos': 3,
+ 'jnt_qposadr': 1,
+ 'jnt_range': 3,
+ 'jnt_solimp': 3,
+ 'jnt_solref': 3,
+ 'jnt_stiffness': 2,
+ 'jnt_type': 1,
+ 'light_bodyid': 1,
+ 'light_dir': 3,
+ 'light_dir0': 3,
+ 'light_mode': 1,
+ 'light_pos': 3,
+ 'light_pos0': 3,
+ 'light_poscom0': 3,
+ 'light_targetbodyid': 1,
+ 'mapM2M': 1,
+ 'mat_rgba': 3,
+ 'mesh_face': 2,
+ 'mesh_faceadr': 1,
+ 'mesh_graph': 1,
+ 'mesh_graphadr': 1,
+ 'mesh_polyadr': 1,
+ 'mesh_polymap': 1,
+ 'mesh_polymapadr': 1,
+ 'mesh_polymapnum': 1,
+ 'mesh_polynormal': 2,
+ 'mesh_polynum': 1,
+ 'mesh_polyvert': 1,
+ 'mesh_polyvertadr': 1,
+ 'mesh_polyvertnum': 1,
+ 'mesh_vert': 2,
+ 'mesh_vertadr': 1,
+ 'mesh_vertnum': 1,
+ 'mocap_bodyid': 1,
+ 'nC': 0,
+ 'nM': 0,
+ 'na': 0,
+ 'nbody': 0,
+ 'ncam': 0,
+ 'neq': 0,
+ 'nexclude': 0,
+ 'nflex': 0,
+ 'nflexedge': 0,
+ 'nflexelem': 0,
+ 'nflexelemdata': 0,
+ 'nflexvert': 0,
+ 'ngeom': 0,
+ 'ngravcomp': 0,
+ 'nhfield': 0,
+ 'nhfielddata': 0,
+ 'njnt': 0,
+ 'nlight': 0,
+ 'nlsp': 0,
+ 'nmeshface': 0,
+ 'nmeshgraph': 0,
+ 'nmeshpoly': 0,
+ 'nmeshpolymap': 0,
+ 'nmeshpolyvert': 0,
+ 'nmeshvert': 0,
+ 'nmocap': 0,
+ 'npair': 0,
+ 'nq': 0,
+ 'nsensor': 0,
+ 'nsensordata': 0,
+ 'nsite': 0,
+ 'ntendon': 0,
+ 'nu': 0,
+ 'nv': 0,
+ 'nwrap': 0,
+ 'nxn_geom_pair': 2,
+ 'nxn_geom_pair_filtered': 2,
+ 'nxn_pairid': 1,
+ 'nxn_pairid_filtered': 1,
+ 'opt__broadphase': 0,
+ 'opt__broadphase_filter': 0,
+ 'opt__cone': 0,
+ 'opt__density': 1,
+ 'opt__disableflags': 0,
+ 'opt__enableflags': 0,
+ 'opt__epa_iterations': 0,
+ 'opt__gjk_iterations': 0,
+ 'opt__graph_conditional': 0,
+ 'opt__gravity': 2,
+ 'opt__has_fluid': 0,
+ 'opt__impratio': 1,
+ 'opt__integrator': 0,
+ 'opt__is_sparse': 0,
+ 'opt__iterations': 0,
+ 'opt__ls_iterations': 0,
+ 'opt__ls_parallel': 0,
+ 'opt__ls_tolerance': 1,
+ 'opt__magnetic': 2,
+ 'opt__run_collision_detection': 0,
+ 'opt__sdf_initpoints': 0,
+ 'opt__sdf_iterations': 0,
+ 'opt__solver': 0,
+ 'opt__timestep': 1,
+ 'opt__tolerance': 1,
+ 'opt__viscosity': 1,
+ 'opt__wind': 2,
+ 'pair_dim': 1,
+ 'pair_friction': 3,
+ 'pair_gap': 2,
+ 'pair_geom1': 1,
+ 'pair_geom2': 1,
+ 'pair_margin': 2,
+ 'pair_solimp': 3,
+ 'pair_solref': 3,
+ 'pair_solreffriction': 3,
+ 'plugin': 1,
+ 'plugin_attr': 2,
+ 'qLD_updates': -1,
+ 'qM_fullm_i': 1,
+ 'qM_fullm_j': 1,
+ 'qM_madr_ij': 1,
+ 'qM_mulm_i': 1,
+ 'qM_mulm_j': 1,
+ 'qM_tiles': -1,
+ 'qpos0': 2,
+ 'qpos_spring': 2,
+ 'rangefinder_sensor_adr': 1,
+ 'sensor_acc_adr': 1,
+ 'sensor_adr': 1,
+ 'sensor_cutoff': 1,
+ 'sensor_datatype': 1,
+ 'sensor_dim': 1,
+ 'sensor_e_kinetic': 0,
+ 'sensor_e_potential': 0,
+ 'sensor_limitfrc_adr': 1,
+ 'sensor_limitpos_adr': 1,
+ 'sensor_limitvel_adr': 1,
+ 'sensor_objid': 1,
+ 'sensor_objtype': 1,
+ 'sensor_pos_adr': 1,
+ 'sensor_rangefinder_adr': 1,
+ 'sensor_rangefinder_bodyid': 1,
+ 'sensor_refid': 1,
+ 'sensor_reftype': 1,
+ 'sensor_rne_postconstraint': 0,
+ 'sensor_subtree_vel': 0,
+ 'sensor_tendonactfrc_adr': 1,
+ 'sensor_touch_adr': 1,
+ 'sensor_type': 1,
+ 'sensor_vel_adr': 1,
+ 'site_bodyid': 1,
+ 'site_pos': 3,
+ 'site_quat': 3,
+ 'site_size': 2,
+ 'site_type': 1,
+ 'stat__meaninertia': 0,
+ 'subtree_mass': 2,
+ 'ten_wrapadr_site': 1,
+ 'ten_wrapnum_site': 1,
+ 'tendon_actfrclimited': 1,
+ 'tendon_actfrcrange': 3,
+ 'tendon_adr': 1,
+ 'tendon_armature': 2,
+ 'tendon_damping': 2,
+ 'tendon_frictionloss': 2,
+ 'tendon_geom_adr': 1,
+ 'tendon_invweight0': 2,
+ 'tendon_jnt_adr': 1,
+ 'tendon_length0': 2,
+ 'tendon_lengthspring': 3,
+ 'tendon_limited': 1,
+ 'tendon_limited_adr': 1,
+ 'tendon_margin': 2,
+ 'tendon_num': 1,
+ 'tendon_range': 3,
+ 'tendon_site_pair_adr': 1,
+ 'tendon_solimp_fri': 3,
+ 'tendon_solimp_lim': 3,
+ 'tendon_solref_fri': 3,
+ 'tendon_solref_lim': 3,
+ 'tendon_stiffness': 2,
+ 'wrap_geom_adr': 1,
+ 'wrap_jnt_adr': 1,
+ 'wrap_objid': 1,
+ 'wrap_prm': 1,
+ 'wrap_pulley_scale': 1,
+ 'wrap_site_adr': 1,
+ 'wrap_site_pair_adr': 1,
+ 'wrap_type': 1,
+ },
+ 'Option': {
+ 'broadphase': 0,
+ 'broadphase_filter': 0,
+ 'cone': 0,
+ 'density': 1,
+ 'disableflags': 0,
+ 'enableflags': 0,
+ 'epa_iterations': 0,
+ 'gjk_iterations': 0,
+ 'graph_conditional': 0,
+ 'gravity': 2,
+ 'has_fluid': 0,
+ 'impratio': 1,
+ 'integrator': 0,
+ 'is_sparse': 0,
+ 'iterations': 0,
+ 'ls_iterations': 0,
+ 'ls_parallel': 0,
+ 'ls_tolerance': 1,
+ 'magnetic': 2,
+ 'run_collision_detection': 0,
+ 'sdf_initpoints': 0,
+ 'sdf_iterations': 0,
+ 'solver': 0,
+ 'timestep': 1,
+ 'tolerance': 1,
+ 'viscosity': 1,
+ 'wind': 2,
+ },
+ 'Statistic': {'meaninertia': 0},
+}
+BATCH_DIM = {
+ 'Data': {
+ 'act': True,
+ 'act_dot': True,
+ 'act_dot_rk': True,
+ 'act_t0': True,
+ 'act_vel_integration': True,
+ 'actuator_force': True,
+ 'actuator_length': True,
+ 'actuator_moment': True,
+ 'actuator_trntype_body_ncon': True,
+ 'actuator_velocity': True,
+ 'cacc': True,
+ 'cam_xmat': True,
+ 'cam_xpos': True,
+ 'cdof': True,
+ 'cdof_dot': True,
+ 'cfrc_ext': True,
+ 'cfrc_int': True,
+ 'cinert': True,
+ 'collision_hftri_index': False,
+ 'collision_pair': False,
+ 'collision_pairid': False,
+ 'collision_worldid': False,
+ 'contact__dim': False,
+ 'contact__dist': False,
+ 'contact__efc_address': False,
+ 'contact__frame': False,
+ 'contact__friction': False,
+ 'contact__geom': False,
+ 'contact__includemargin': False,
+ 'contact__pos': False,
+ 'contact__solimp': False,
+ 'contact__solref': False,
+ 'contact__solreffriction': False,
+ 'contact__worldid': False,
+ 'crb': True,
+ 'ctrl': True,
+ 'cvel': True,
+ 'efc__D': True,
+ 'efc__J': True,
+ 'efc__Jaref': True,
+ 'efc__Ma': True,
+ 'efc__Mgrad': True,
+ 'efc__active': True,
+ 'efc__alpha': True,
+ 'efc__aref': True,
+ 'efc__beta': True,
+ 'efc__beta_den': True,
+ 'efc__beta_num': True,
+ 'efc__cholesky_L_tmp': True,
+ 'efc__cholesky_y_tmp': True,
+ 'efc__condim': True,
+ 'efc__cost': True,
+ 'efc__cost_candidate': True,
+ 'efc__done': True,
+ 'efc__force': True,
+ 'efc__frictionloss': True,
+ 'efc__gauss': True,
+ 'efc__grad': True,
+ 'efc__grad_dot': True,
+ 'efc__gtol': True,
+ 'efc__h': True,
+ 'efc__hi': True,
+ 'efc__hi_alpha': True,
+ 'efc__hi_next': True,
+ 'efc__hi_next_alpha': True,
+ 'efc__id': True,
+ 'efc__jv': True,
+ 'efc__lo': True,
+ 'efc__lo_alpha': True,
+ 'efc__lo_next': True,
+ 'efc__lo_next_alpha': True,
+ 'efc__ls_done': True,
+ 'efc__margin': True,
+ 'efc__mid': True,
+ 'efc__mid_alpha': True,
+ 'efc__mv': True,
+ 'efc__p0': True,
+ 'efc__pos': True,
+ 'efc__prev_Mgrad': True,
+ 'efc__prev_cost': True,
+ 'efc__prev_grad': True,
+ 'efc__quad': True,
+ 'efc__quad_gauss': True,
+ 'efc__search': True,
+ 'efc__search_dot': True,
+ 'efc__type': True,
+ 'efc__u': False,
+ 'efc__uu': False,
+ 'efc__uv': False,
+ 'efc__vel': True,
+ 'efc__vv': False,
+ 'energy': True,
+ 'energy_vel_mul_m_skip': True,
+ 'epa_face': False,
+ 'epa_horizon': False,
+ 'epa_index': False,
+ 'epa_map': False,
+ 'epa_norm2': False,
+ 'epa_pr': False,
+ 'epa_vert': False,
+ 'epa_vert1': False,
+ 'epa_vert2': False,
+ 'epa_vert_index1': False,
+ 'epa_vert_index2': False,
+ 'eq_active': True,
+ 'flexedge_length': True,
+ 'flexedge_velocity': True,
+ 'flexvert_xpos': True,
+ 'fluid_applied': True,
+ 'geom_skip': False,
+ 'geom_xmat': True,
+ 'geom_xpos': True,
+ 'inverse_mul_m_skip': True,
+ 'light_xdir': True,
+ 'light_xpos': True,
+ 'mocap_pos': True,
+ 'mocap_quat': True,
+ 'ncollision': False,
+ 'ncon': False,
+ 'ncon_hfield': True,
+ 'nconmax': False,
+ 'ne': True,
+ 'ne_connect': True,
+ 'ne_jnt': True,
+ 'ne_ten': True,
+ 'ne_weld': True,
+ 'nefc': True,
+ 'nf': True,
+ 'njmax': False,
+ 'nl': True,
+ 'nsolving': False,
+ 'nworld': False,
+ 'qLD': True,
+ 'qLD_integration': True,
+ 'qLDiagInv': True,
+ 'qLDiagInv_integration': True,
+ 'qM': True,
+ 'qM_integration': True,
+ 'qacc': True,
+ 'qacc_discrete': True,
+ 'qacc_integration': True,
+ 'qacc_rk': True,
+ 'qacc_smooth': True,
+ 'qacc_warmstart': True,
+ 'qfrc_actuator': True,
+ 'qfrc_applied': True,
+ 'qfrc_bias': True,
+ 'qfrc_constraint': True,
+ 'qfrc_damper': True,
+ 'qfrc_fluid': True,
+ 'qfrc_gravcomp': True,
+ 'qfrc_integration': True,
+ 'qfrc_inverse': True,
+ 'qfrc_passive': True,
+ 'qfrc_smooth': True,
+ 'qfrc_spring': True,
+ 'qpos': True,
+ 'qpos_t0': True,
+ 'qvel': True,
+ 'qvel_rk': True,
+ 'qvel_t0': True,
+ 'ray_bodyexclude': False,
+ 'ray_dist': True,
+ 'ray_geomid': True,
+ 'sap_cumulative_sum': True,
+ 'sap_projection_lower': True,
+ 'sap_projection_upper': True,
+ 'sap_range': True,
+ 'sap_segment_index': True,
+ 'sap_sort_index': True,
+ 'sensor_rangefinder_dist': True,
+ 'sensor_rangefinder_geomid': True,
+ 'sensor_rangefinder_pnt': True,
+ 'sensor_rangefinder_vec': True,
+ 'sensordata': True,
+ 'site_xmat': True,
+ 'site_xpos': True,
+ 'solver_niter': True,
+ 'subtree_angmom': True,
+ 'subtree_bodyvel': True,
+ 'subtree_com': True,
+ 'subtree_linvel': True,
+ 'ten_J': True,
+ 'ten_Jdot': True,
+ 'ten_actfrc': True,
+ 'ten_bias_coef': True,
+ 'ten_length': True,
+ 'ten_velocity': True,
+ 'ten_wrapadr': True,
+ 'ten_wrapnum': True,
+ 'time': True,
+ 'wrap_geom_xpos': True,
+ 'wrap_obj': True,
+ 'wrap_xpos': True,
+ 'xanchor': True,
+ 'xaxis': True,
+ 'xfrc_applied': True,
+ 'ximat': True,
+ 'xipos': True,
+ 'xmat': True,
+ 'xpos': True,
+ 'xquat': True,
+ },
+ 'Model': {
+ 'M_colind': False,
+ 'M_rowadr': False,
+ 'M_rownnz': False,
+ 'actuator_acc0': False,
+ 'actuator_actadr': False,
+ 'actuator_actearly': False,
+ 'actuator_actlimited': False,
+ 'actuator_actnum': False,
+ 'actuator_actrange': True,
+ 'actuator_affine_bias_gain': False,
+ 'actuator_biasprm': True,
+ 'actuator_biastype': False,
+ 'actuator_cranklength': False,
+ 'actuator_ctrllimited': False,
+ 'actuator_ctrlrange': True,
+ 'actuator_dynprm': True,
+ 'actuator_dyntype': False,
+ 'actuator_forcelimited': False,
+ 'actuator_forcerange': True,
+ 'actuator_gainprm': True,
+ 'actuator_gaintype': False,
+ 'actuator_gear': True,
+ 'actuator_lengthrange': False,
+ 'actuator_moment_tiles_nu': False,
+ 'actuator_moment_tiles_nv': False,
+ 'actuator_trnid': False,
+ 'actuator_trntype': False,
+ 'actuator_trntype_body_adr': False,
+ 'block_dim__actuator_velocity': False,
+ 'block_dim__cholesky_factorize': False,
+ 'block_dim__cholesky_factorize_solve': False,
+ 'block_dim__cholesky_solve': False,
+ 'block_dim__energy_vel_kinetic': False,
+ 'block_dim__euler_dense': False,
+ 'block_dim__mul_m_dense': False,
+ 'block_dim__qderiv_actuator_passive_actuation': False,
+ 'block_dim__qderiv_actuator_passive_no_actuation': False,
+ 'block_dim__qfrc_actuator': False,
+ 'block_dim__ray': False,
+ 'block_dim__segmented_sort': False,
+ 'block_dim__tendon_velocity': False,
+ 'block_dim__update_gradient_cholesky': False,
+ 'body_conaffinity': False,
+ 'body_contype': False,
+ 'body_dofadr': False,
+ 'body_dofnum': False,
+ 'body_geomadr': False,
+ 'body_geomnum': False,
+ 'body_gravcomp': True,
+ 'body_inertia': True,
+ 'body_invweight0': True,
+ 'body_ipos': True,
+ 'body_iquat': True,
+ 'body_jntadr': False,
+ 'body_jntnum': False,
+ 'body_mass': True,
+ 'body_mocapid': False,
+ 'body_parentid': False,
+ 'body_pos': True,
+ 'body_quat': True,
+ 'body_rootid': False,
+ 'body_subtreemass': True,
+ 'body_tree': False,
+ 'body_weldid': False,
+ 'cam_bodyid': False,
+ 'cam_fovy': False,
+ 'cam_intrinsic': False,
+ 'cam_mat0': True,
+ 'cam_mode': False,
+ 'cam_pos': True,
+ 'cam_pos0': True,
+ 'cam_poscom0': True,
+ 'cam_quat': True,
+ 'cam_resolution': False,
+ 'cam_sensorsize': False,
+ 'cam_targetbodyid': False,
+ 'condim_max': False,
+ 'dof_Madr': False,
+ 'dof_armature': True,
+ 'dof_bodyid': False,
+ 'dof_damping': True,
+ 'dof_frictionloss': True,
+ 'dof_invweight0': True,
+ 'dof_jntid': False,
+ 'dof_parentid': False,
+ 'dof_solimp': True,
+ 'dof_solref': True,
+ 'dof_tri_col': False,
+ 'dof_tri_row': False,
+ 'eq_active0': False,
+ 'eq_connect_adr': False,
+ 'eq_data': True,
+ 'eq_jnt_adr': False,
+ 'eq_obj1id': False,
+ 'eq_obj2id': False,
+ 'eq_objtype': False,
+ 'eq_solimp': True,
+ 'eq_solref': True,
+ 'eq_ten_adr': False,
+ 'eq_type': False,
+ 'eq_wld_adr': False,
+ 'exclude_signature': False,
+ 'flex_bending': False,
+ 'flex_damping': False,
+ 'flex_dim': False,
+ 'flex_edge': False,
+ 'flex_edgeadr': False,
+ 'flex_edgeflap': False,
+ 'flex_elem': False,
+ 'flex_elemedge': False,
+ 'flex_elemedgeadr': False,
+ 'flex_stiffness': False,
+ 'flex_vertadr': False,
+ 'flex_vertbodyid': False,
+ 'flex_vertnum': False,
+ 'flexedge_length0': False,
+ 'geom_aabb': False,
+ 'geom_bodyid': False,
+ 'geom_conaffinity': False,
+ 'geom_condim': False,
+ 'geom_contype': False,
+ 'geom_dataid': False,
+ 'geom_friction': True,
+ 'geom_gap': True,
+ 'geom_group': False,
+ 'geom_margin': True,
+ 'geom_matid': True,
+ 'geom_pair_type_count': False,
+ 'geom_plugin_index': False,
+ 'geom_pos': True,
+ 'geom_priority': False,
+ 'geom_quat': True,
+ 'geom_rbound': True,
+ 'geom_rgba': True,
+ 'geom_size': True,
+ 'geom_solimp': True,
+ 'geom_solmix': True,
+ 'geom_solref': True,
+ 'geom_type': False,
+ 'geompair2hfgeompair': False,
+ 'has_sdf_geom': False,
+ 'hfield_adr': False,
+ 'hfield_data': False,
+ 'hfield_ncol': False,
+ 'hfield_nrow': False,
+ 'hfield_size': False,
+ 'jnt_actfrclimited': False,
+ 'jnt_actfrcrange': True,
+ 'jnt_actgravcomp': False,
+ 'jnt_axis': True,
+ 'jnt_bodyid': False,
+ 'jnt_dofadr': False,
+ 'jnt_limited': False,
+ 'jnt_limited_ball_adr': False,
+ 'jnt_limited_slide_hinge_adr': False,
+ 'jnt_margin': True,
+ 'jnt_pos': True,
+ 'jnt_qposadr': False,
+ 'jnt_range': True,
+ 'jnt_solimp': True,
+ 'jnt_solref': True,
+ 'jnt_stiffness': True,
+ 'jnt_type': False,
+ 'light_bodyid': False,
+ 'light_dir': True,
+ 'light_dir0': True,
+ 'light_mode': False,
+ 'light_pos': True,
+ 'light_pos0': True,
+ 'light_poscom0': True,
+ 'light_targetbodyid': False,
+ 'mapM2M': False,
+ 'mat_rgba': True,
+ 'mesh_face': False,
+ 'mesh_faceadr': False,
+ 'mesh_graph': False,
+ 'mesh_graphadr': False,
+ 'mesh_polyadr': False,
+ 'mesh_polymap': False,
+ 'mesh_polymapadr': False,
+ 'mesh_polymapnum': False,
+ 'mesh_polynormal': False,
+ 'mesh_polynum': False,
+ 'mesh_polyvert': False,
+ 'mesh_polyvertadr': False,
+ 'mesh_polyvertnum': False,
+ 'mesh_vert': False,
+ 'mesh_vertadr': False,
+ 'mesh_vertnum': False,
+ 'mocap_bodyid': False,
+ 'nC': False,
+ 'nM': False,
+ 'na': False,
+ 'nbody': False,
+ 'ncam': False,
+ 'neq': False,
+ 'nexclude': False,
+ 'nflex': False,
+ 'nflexedge': False,
+ 'nflexelem': False,
+ 'nflexelemdata': False,
+ 'nflexvert': False,
+ 'ngeom': False,
+ 'ngravcomp': False,
+ 'nhfield': False,
+ 'nhfielddata': False,
+ 'njnt': False,
+ 'nlight': False,
+ 'nlsp': False,
+ 'nmeshface': False,
+ 'nmeshgraph': False,
+ 'nmeshpoly': False,
+ 'nmeshpolymap': False,
+ 'nmeshpolyvert': False,
+ 'nmeshvert': False,
+ 'nmocap': False,
+ 'npair': False,
+ 'nq': False,
+ 'nsensor': False,
+ 'nsensordata': False,
+ 'nsite': False,
+ 'ntendon': False,
+ 'nu': False,
+ 'nv': False,
+ 'nwrap': False,
+ 'nxn_geom_pair': False,
+ 'nxn_geom_pair_filtered': False,
+ 'nxn_pairid': False,
+ 'nxn_pairid_filtered': False,
+ 'opt__broadphase': False,
+ 'opt__broadphase_filter': False,
+ 'opt__cone': False,
+ 'opt__density': True,
+ 'opt__disableflags': False,
+ 'opt__enableflags': False,
+ 'opt__epa_iterations': False,
+ 'opt__gjk_iterations': False,
+ 'opt__graph_conditional': False,
+ 'opt__gravity': True,
+ 'opt__has_fluid': False,
+ 'opt__impratio': True,
+ 'opt__integrator': False,
+ 'opt__is_sparse': False,
+ 'opt__iterations': False,
+ 'opt__ls_iterations': False,
+ 'opt__ls_parallel': False,
+ 'opt__ls_tolerance': True,
+ 'opt__magnetic': True,
+ 'opt__run_collision_detection': False,
+ 'opt__sdf_initpoints': False,
+ 'opt__sdf_iterations': False,
+ 'opt__solver': False,
+ 'opt__timestep': True,
+ 'opt__tolerance': True,
+ 'opt__viscosity': True,
+ 'opt__wind': True,
+ 'pair_dim': False,
+ 'pair_friction': True,
+ 'pair_gap': True,
+ 'pair_geom1': False,
+ 'pair_geom2': False,
+ 'pair_margin': True,
+ 'pair_solimp': True,
+ 'pair_solref': True,
+ 'pair_solreffriction': True,
+ 'plugin': False,
+ 'plugin_attr': False,
+ 'qLD_updates': False,
+ 'qM_fullm_i': False,
+ 'qM_fullm_j': False,
+ 'qM_madr_ij': False,
+ 'qM_mulm_i': False,
+ 'qM_mulm_j': False,
+ 'qM_tiles': False,
+ 'qpos0': True,
+ 'qpos_spring': True,
+ 'rangefinder_sensor_adr': False,
+ 'sensor_acc_adr': False,
+ 'sensor_adr': False,
+ 'sensor_cutoff': False,
+ 'sensor_datatype': False,
+ 'sensor_dim': False,
+ 'sensor_e_kinetic': False,
+ 'sensor_e_potential': False,
+ 'sensor_limitfrc_adr': False,
+ 'sensor_limitpos_adr': False,
+ 'sensor_limitvel_adr': False,
+ 'sensor_objid': False,
+ 'sensor_objtype': False,
+ 'sensor_pos_adr': False,
+ 'sensor_rangefinder_adr': False,
+ 'sensor_rangefinder_bodyid': False,
+ 'sensor_refid': False,
+ 'sensor_reftype': False,
+ 'sensor_rne_postconstraint': False,
+ 'sensor_subtree_vel': False,
+ 'sensor_tendonactfrc_adr': False,
+ 'sensor_touch_adr': False,
+ 'sensor_type': False,
+ 'sensor_vel_adr': False,
+ 'site_bodyid': False,
+ 'site_pos': True,
+ 'site_quat': True,
+ 'site_size': False,
+ 'site_type': False,
+ 'stat__meaninertia': False,
+ 'subtree_mass': True,
+ 'ten_wrapadr_site': False,
+ 'ten_wrapnum_site': False,
+ 'tendon_actfrclimited': False,
+ 'tendon_actfrcrange': True,
+ 'tendon_adr': False,
+ 'tendon_armature': True,
+ 'tendon_damping': True,
+ 'tendon_frictionloss': True,
+ 'tendon_geom_adr': False,
+ 'tendon_invweight0': True,
+ 'tendon_jnt_adr': False,
+ 'tendon_length0': True,
+ 'tendon_lengthspring': True,
+ 'tendon_limited': False,
+ 'tendon_limited_adr': False,
+ 'tendon_margin': True,
+ 'tendon_num': False,
+ 'tendon_range': True,
+ 'tendon_site_pair_adr': False,
+ 'tendon_solimp_fri': True,
+ 'tendon_solimp_lim': True,
+ 'tendon_solref_fri': True,
+ 'tendon_solref_lim': True,
+ 'tendon_stiffness': True,
+ 'wrap_geom_adr': False,
+ 'wrap_jnt_adr': False,
+ 'wrap_objid': False,
+ 'wrap_prm': False,
+ 'wrap_pulley_scale': False,
+ 'wrap_site_adr': False,
+ 'wrap_site_pair_adr': False,
+ 'wrap_type': False,
+ },
+ 'Option': {
+ 'broadphase': False,
+ 'broadphase_filter': False,
+ 'cone': False,
+ 'density': True,
+ 'disableflags': False,
+ 'enableflags': False,
+ 'epa_iterations': False,
+ 'gjk_iterations': False,
+ 'graph_conditional': False,
+ 'gravity': True,
+ 'has_fluid': False,
+ 'impratio': True,
+ 'integrator': False,
+ 'is_sparse': False,
+ 'iterations': False,
+ 'ls_iterations': False,
+ 'ls_parallel': False,
+ 'ls_tolerance': True,
+ 'magnetic': True,
+ 'run_collision_detection': False,
+ 'sdf_initpoints': False,
+ 'sdf_iterations': False,
+ 'solver': False,
+ 'timestep': True,
+ 'tolerance': True,
+ 'viscosity': True,
+ 'wind': True,
+ },
+ 'Statistic': {'meaninertia': False},
+}
diff --git a/mjx/pyproject.toml b/mjx/pyproject.toml
index c80bcf1d..38a21389 100644
--- a/mjx/pyproject.toml
+++ b/mjx/pyproject.toml
@@ -33,6 +33,8 @@ dependencies = [
"mujoco>=3.3.5.dev0",
"scipy",
"trimesh",
+ "warp-lang==1.8.0; python_version >= '3.12' and sys_platform == 'linux'",
+ "warp-lang==1.8.0; python_version >= '3.12' and sys_platform == 'win32'",
]
[project.scripts]
@@ -65,3 +67,8 @@ pyink-use-majority-quotes = true
extend-exclude = '''(
.ipynb$
)'''
+
+[tool.pytest.ini_options]
+norecursedirs = [
+ "**/third_party",
+]
diff --git a/mjx/requirements.txt b/mjx/requirements.txt
index 9b8cf8c1..9340b6e2 100644
--- a/mjx/requirements.txt
+++ b/mjx/requirements.txt
@@ -4,31 +4,51 @@ etils[epath]==1.10.0; python_version >= '3.10' \
--hash=sha256:0777fe60a234b4c65ca53470fc64f2dd2d0c6bca7fcc623fdaa8d7fa5a317098
etils[epath]==1.5.2; python_version == '3.9' \
--hash=sha256:6dc882d355e1e98a5d1a148d6323679dc47c9a5792939b9de72615aa4737eb0b
-jax==0.4.34; python_version >= '3.10' \
- --hash=sha256:b957ca1fc91f7343f91a186af9f19c7f342c946f95a8c11c7f1e5cdfe2e58d9e
+jax==0.5.3; python_version >= '3.10' and (sys_platform != 'darwin' or platform_machine != 'x86_64') \
+ --hash=sha256:1483dc237b4f47e41755d69429e8c3c138736716147cd43bb2b99b259d4e3c41 \
+ --hash=sha256:f17fcb0fd61dc289394af6ce4de2dada2312f2689bb0d73642c6f026a95fbb2c
+jax==0.4.38; python_version >= '3.10' and sys_platform == 'darwin' and platform_machine == 'x86_64' \
+ --hash=sha256:78987306f7041ea8500d99df1a17c33ed92620c2268c4c3677fb24e06712be64
jax==0.4.30; python_version == '3.9' \
--hash=sha256:289b30ae03b52f7f4baf6ef082a9f4e3e29c1080e22d13512c5ecf02d5f1a55b
-jaxlib==0.4.34; python_version >= '3.10' \
- --hash=sha256:6b43a974c5d91a19912d138f2658dd8dbb7d30dcdff5c961d896c673e872b611 \
- --hash=sha256:87f25a477cd279840e53718403f97092eba0e8a945fcab47bcf435b6f9119dda \
- --hash=sha256:7be673a876ebd1aef440fb7e3ebaf99a91abeb550c9728c644b7d7c7b5d7c108 \
- --hash=sha256:c303f5acaf6c56ce5ff133a923c9b6247bdebedde15bd2c893c24be4d8f71306 \
- --hash=sha256:72e22e99a5dc890a64443c3fc12f13f20091f578c405a76de077ba42b4c62cd7 \
- --hash=sha256:901cb4040ed24eae40071d8114ea8d10dff436277fa74a1a5b9e7206f641151c \
- --hash=sha256:48272e9034ff868d4328cf0055a07882fd2be93f59dfb6283af7de491f9d1290 \
- --hash=sha256:1a30771d85fa77f9ab8f18e63240f455ab3a3f87660ed7b8d5eea6ceecbe5c1e \
- --hash=sha256:096f0ca309d41fa692a9d1f2f9baab1c5c8ca0749876ebb3f748e738a27c7ff4 \
- --hash=sha256:c7b3e724a30426a856070aba0192b5d199e95b4411070e7ad96ad8b196877b10 \
- --hash=sha256:133070d4fec5525ffea4dc72956398c1cf647a04dcb37f8a935ee82af78d9965 \
- --hash=sha256:3bcfa639ca3cfaf86c8ceebd5fc0d47300fd98a078014a1d0cc03133e1523d5f \
- --hash=sha256:571ef03259835458111596a71a2f4a6fabf4ec34595df4cea555035362ac5bf0 \
- --hash=sha256:c9d3adcae43a33aad4332be9c2aedc5ef751d1e755f917a5afb30c7872eacaa8 \
- --hash=sha256:8ee3f93836e53c86556ccd9449a4ea43516ee05184d031a71dd692e81259f7d9 \
- --hash=sha256:b0001c8f0e2b1c7bc99e4f314b524a340d25653505c1a1484d4041a9d3617f6f \
- --hash=sha256:d840e64b85f8865404d6d225b9bb340e158df1457152a361b05680e24792b232 \
- --hash=sha256:3e60bc826933082e99b19b87c21818a8d26fcdb01f418d47cedff554746fd6cc \
- --hash=sha256:45d719a2ce0ebf21255a277b71d756f3609b7b5be70cddc5d88fd58c35219de0 \
- --hash=sha256:b7a212a3cb5c6acc201c32ae4f4b5f5a9ac09457fbb77ba8db5ce7e7d4adc214
+jaxlib==0.5.3; python_version >= '3.10' and (sys_platform != 'darwin' or platform_machine != 'x86_64') \
+ --hash=sha256:48ff5c89fb8a0fe04d475e9ddc074b4879a91d7ab68a51cec5cd1e87f81e6c47 \
+ --hash=sha256:972400db4af6e85270d81db5e6e620d31395f0472e510c50dfcd4cb3f72b7220 \
+ --hash=sha256:52be6c9775aff738a61170d8c047505c75bb799a45518e66a7a0908127b11785 \
+ --hash=sha256:b41a6fcaeb374fabc4ee7e74cfed60843bdab607cd54f60a68b7f7655cde2b66 \
+ --hash=sha256:b62bd8b29e5a4f9bfaa57c8daf6e04820b2c994f448f3dec602d64255545e9f2 \
+ --hash=sha256:a4666f81d72c060ed3e581ded116a9caa9b0a70a148a54cb12a1d3afca3624b5 \
+ --hash=sha256:29e1530fc81833216f1e28b578d0c59697654f72ee31c7a44ed7753baf5ac466 \
+ --hash=sha256:8eb54e38d789557579f900ea3d70f104a440f8555a9681ed45f4a122dcbfd92e \
+ --hash=sha256:d394dbde4a1c6bd67501cfb29d3819a10b900cb534cc0fc603319f7092f24cfa \
+ --hash=sha256:bddf6360377aa1c792e47fd87f307c342e331e5ff3582f940b1bca00f6b4bc73 \
+ --hash=sha256:5a5e88ab1cd6fdf78d69abe3544e8f09cce200dd339bb85fbe3c2ea67f2a5e68 \
+ --hash=sha256:520665929649f29f7d948d4070dbaf3e032a4c1f7c11f2863eac73320fcee784 \
+ --hash=sha256:31321c25282a06a6dfc940507bc14d0a0ac838d8ced6c07aa00a7fae34ce7b3f \
+ --hash=sha256:e904b92dedfbc7e545725a8d7676987030ae9c069001d94701bc109c6dab4100 \
+ --hash=sha256:bb7593cb7fffcb13963f22fa5229ed960b8fb4ae5ec3b0820048cbd67f1e8e31 \
+ --hash=sha256:8019f73a10b1290f988dd3768c684f3a8a147239091c3b790ce7e47e3bbc00bd
+jaxlib==0.4.38; python_version >= '3.10' and sys_platform == 'darwin' and platform_machine == 'x86_64' \
+ --hash=sha256:55c19b9d3f33a6fc59f644aa5a21fba02639ccdd776cb4a9b5526625f57839ff \
+ --hash=sha256:30b2f52cb50d74734af2f477c2533a7a583e3bb7b2c8acdeb361ee77d940577a \
+ --hash=sha256:ee19c163a8fdf0839d4c18b88a5fbfb4e731ba7c437416d3e5483e570bb764e4 \
+ --hash=sha256:61aeccb9a27c67fdb8450f6357240019cd4511cb9d62a44e4764756d384853ad \
+ --hash=sha256:d6ab745a89d0fb737a36fe1d8b86659e3fffe6ee8303b20651b26193d5edc0ef \
+ --hash=sha256:b67fdeabd6dfed08b7768f3bdffb521160085f8305669bd197beef61d08de08b \
+ --hash=sha256:3fb0eaae7369157afecbead50aaf29e73ffddfa77a2335d721bd9794f3c510e4 \
+ --hash=sha256:43db58c4c427627296366a56c10318e1f00f503690e17f94bb4344293e1995e0 \
+ --hash=sha256:2751ff7037d6a997d0be0e77cc4be381c5a9f9bb8b314edb755c13a6fd969f45 \
+ --hash=sha256:35226968fc9de6873d1571670eac4117f5ed80e955f7a1775204d1044abe16c6 \
+ --hash=sha256:3fefea985f0415816f3bbafd3f03a437050275ef9bac9a72c1314e1644ac57c1 \
+ --hash=sha256:f33bcafe32c97a562ecf6894d7c41674c80c0acdedfa5423d49af51147149874 \
+ --hash=sha256:496f45b0e001a2341309cd0c74af0b670537dced79c168cb230cfcc773f0aa86 \
+ --hash=sha256:dad6c0a96567c06d083c0469fec40f201210b099365bd698be31a6d2ec88fd59 \
+ --hash=sha256:966cdec36cfa978f5b4582bcb4147fe511725b94c1a752dac3a5f52ce46b6fa3 \
+ --hash=sha256:41e55ae5818a882e5789e848f6f16687ac132bcfbb5a5fa114a5d18b78d05f2d \
+ --hash=sha256:6fe326b8af366387dd47ccf312583b2b17fed12712c9b74a648b18a13cbdbabf \
+ --hash=sha256:248cca3771ebf24b070f49701364ceada33e6139445b06c782cca5ac5ad92bf4 \
+ --hash=sha256:2ce77ba8cda9259a4bca97afc1c722e4291a6c463a63f8d372c6edc85117d625 \
+ --hash=sha256:4103db0b3a38a5dc132741237453c24d8547290a22079ba1b577d6c88c95300a
jaxlib==0.4.30; python_version == '3.9' \
--hash=sha256:54987e97a22db70f3829b437b9329e4799d653634bacc8b398554d3b90c76b2a \
--hash=sha256:f74a6b0e09df4b5e2ee399ebb9f0e01190e26e84ccb0a758fadb516415c07f18 \
@@ -83,6 +103,12 @@ trimesh==4.5.2 \
--hash=sha256:2e50f3a7fd135c3045da887a1b9f91230528f3ce11d2ec1ba44750d82d6b4f73
wheel==0.45.0 \
--hash=sha256:52f0baa5e6522155090a09c6bd95718cc46956d1b51d537ea5454249edb671c7
+warp-lang==1.8.0; python_version >= '3.12' \
+ --hash=sha256:75a88d2795596f06fcf79eead94e2f194a6195dadbafbd4c9b7c8b4b05456cc4 \
+ --hash=sha256:1be62e7b3e8019ccccc00916c2798afc6c4bfce17d1d6025dc8ef5c790e42c1d \
+ --hash=sha256:373464bee59be37018d134b5924bf8fdb33d1313f484fd1b7c64296971c26ae2 \
+ --hash=sha256:0ecf3b07c1d6d16592ab0318c86f452e16fb3337cfe327bcb76d793a98226dd8 \
+ --hash=sha256:2c7627bee127b522551e02f44c35d7d6a4e5632ff4b58199208012e7f2f6c608
# Transitive dependencies of etils[epath]
fsspec==2024.10.0 \