Add mujoco_warp as an implementation in mjx.

PiperOrigin-RevId: 789366860
Change-Id: I84b49e744552092434df49d855f1052b04291b87
This commit is contained in:
Baruch Tabanpour
2025-07-31 09:32:23 -07:00
committed by Copybara-Service
parent 700da7c5bd
commit 47bc16a37c
112 changed files with 43209 additions and 198 deletions
+2 -1
View File
@@ -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
+10
View File
@@ -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)
-----------------------------
+14 -11
View File
@@ -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
+12 -1
View File
@@ -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
+11 -2
View File
@@ -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.'
)
+25 -5
View File
@@ -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)
+69
View File
@@ -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()
+11 -1
View File
@@ -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
+10
View File
@@ -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:
+309 -32
View File
@@ -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.'
)
+159 -86
View File
@@ -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("""
<mujoco>
<worldbody>
<flexcomp name="flex" type="grid" dim="1" count="3 1 1" mass="1" spacing=".1 .1 .1">
<pin id="0"/>
</flexcomp>
</worldbody>
</mujoco>
""")
mjx.put_model(m, impl='jax')
if __name__ == '__main__':
absltest.main()
+11 -2
View File
@@ -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)
+11 -2
View File
@@ -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)
+68 -28
View File
@@ -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)
@@ -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"/>
<key name="no_efc" qpos="0 0 2
1 0 0 0
0 0 0
0 0 0 0 0 0
0 0 0 0 0 0
0 0 0 0 0 0"/>
</keyframe>
</mujoco>
+79
View File
@@ -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
+14
View File
@@ -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.
# ==============================================================================
@@ -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
@@ -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 = """
<mujoco>
<worldbody>
<body>
<freejoint/>
<geom type="sphere" size="0.1"/>
</body>
<body>
<freejoint/>
<geom type="sphere" size="0.1"/>
</body>
<body>
<freejoint/>
<geom type="capsule" size="0.1 0.1"/>
</body>
<body>
<freejoint/>
<geom type="sphere" size="0.1"/>
</body>
<body>
<freejoint/>
<!-- self collision -->
<geom type="sphere" size="0.1"/>
<geom type="sphere" size="0.1"/>
<!-- parent-child self collision -->
<body>
<geom type="sphere" size="0.1"/>
<joint type="hinge"/>
</body>
</body>
</worldbody>
<keyframe>
<key qpos='0 0 0 1 0 0 0
1 0 0 1 0 0 0
2 0 0 1 0 0 0
3 0 0 1 0 0 0
4 0 0 1 0 0 0
0'/>
<key qpos='0 0 0 1 0 0 0
.05 0 0 1 0 0 0
2 0 0 1 0 0 0
3 0 0 1 0 0 0
4 0 0 1 0 0 0
0'/>
<key qpos='0 0 0 1 0 0 0
.01 0 0 1 0 0 0
.02 0 0 1 0 0 0
3 0 0 1 0 0 0
4 0 0 1 0 0 0
0'/>
<key qpos='0 0 0 1 0 0 0
1 0 0 1 0 0 0
2 0 0 1 0 0 0
2 0 0 1 0 0 0
4 0 0 1 0 0 0
0'/>
</keyframe>
</mujoco>
"""
# 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"""
<mujoco>
<worldbody>
<body>
<geom type="sphere" size=".1" margin="{margin1}"/>
<joint type="slide" axis="1 0 0"/>
</body>
<body>
<geom type="sphere" size=".1" margin="{margin2}"/>
<joint type="slide" axis="1 0 0"/>
</body>
</worldbody>
<keyframe>
<key qpos="0 .21"/>
</keyframe>
</mujoco>
"""
_, _, 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 = """
<mujoco>
<worldbody>
<body>
<geom type="sphere" size=".1"/>
<joint type="slide"/>
<body>
<geom type="sphere" size=".1"/>
<joint type="slide"/>
</body>
</body>
</worldbody>
<keyframe>
<key qpos="0 0"/>
</keyframe>
</mujoco>
"""
_, _, 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 = """
<mujoco>
<option gravity="0 0 0"/>
<worldbody>
<light type="directional" pos="0 0 1"/>
<geom name="floor" size="10 10 .001" type="plane"/>
<body>
<geom type="capsule" size=".05 .1" rgba="0 1 0 1"/>
<joint type="slide" axis="1 0 0"/>
<joint type="slide" axis="0 0 1"/>
<joint type="hinge" axis="0 1 0"/>
</body>
<body>
<geom type="capsule" size=".05 .1" rgba="1 0 0 1"/>
<joint type="slide" axis="1 0 0"/>
<joint type="slide" axis="0 0 1"/>
<joint type="hinge" axis="0 1 0"/>
</body>
</worldbody>
<keyframe>
<key qpos="-.5 .25 0 .5 .25 0"/>
<key qpos="-.5 .075 1.57 .5 .25 0"/>
<key qpos="-.075 .25 0 .075 .25 0"/>
<key qpos="0 .25 .7853 0 .45 .7853"/>
</keyframe>
</mujoco>
"""
_, _, 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()
@@ -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,
],
)
@@ -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 `<exclude>` 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)
@@ -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": """
<mujoco>
<extension>
<plugin plugin="mujoco.sdf.nut">
<instance name="nut1">
<config key="radius" value="0.27"/>
</instance>
<instance name="nut2">
<config key="radius" value="0.27"/>
</instance>
</plugin>
</extension>
<compiler autolimits="true"/>
<option sdf_iterations="10"/> <!-- Increased for better collision accuracy -->
<asset>
<mesh name="nut_mesh1">
<plugin instance="nut1"/>
</mesh>
<mesh name="nut_mesh2">
<plugin instance="nut2"/>
</mesh>
</asset>
<worldbody>
<!-- First nut (floating) -->
<body pos="0 0 0.5">
<joint type="free"/>
<geom type="sdf" name="nut1" mesh="nut_mesh1" rgba="0.83 0.68 0.4 1">
<plugin instance="nut1"/>
</geom>
</body>
<!-- Second nut (positioned to intersect) -->
<body pos="0 0 0.4">
<geom type="sdf" name="nut2" mesh="nut_mesh2" rgba="0.9 0.4 0.2 1">
<plugin instance="nut2"/>
</geom>
</body>
<light name="left" pos="-1 0 2" cutoff="80"/>
<light name="right" pos="1 0 2" cutoff="80"/>
</worldbody>
</mujoco>
""",
"NUT_BOLT": """<mujoco>
<extension>
<plugin plugin="mujoco.sdf.nut">
<instance name="nut">
<config key="radius" value="0.26"/>
</instance>
</plugin>
<plugin plugin="mujoco.sdf.bolt">
<instance name="bolt">
<config key="radius" value="0.255"/>
</instance>
</plugin>
</extension>
<compiler autolimits="true"/>
<visual>
<map force="0.05"/>
</visual>
<asset>
<mesh name="nut">
<plugin instance="nut"/>
</mesh>
<mesh name="bolt">
<plugin instance="bolt"/>
</mesh>
</asset>
<option sdf_iterations="10" sdf_initpoints="20"/>
<default>
<geom solref="0.01 1" solimp=".95 .99 .0001" friction="0.01"/>
</default>
<statistic meansize=".1"/>
<worldbody>
<body pos="-0.0012496 0.00329058 0.830362" quat="-0.000212626 0.999996 -0.00200453 0.00185878">
<joint type="free" damping="30"/>
<geom type="sdf" name="nut" mesh="nut" rgba="0.83 0.68 0.4 1">
<plugin instance="nut"/>
</geom>
</body>
<body euler="180 0 0">
<geom type="sdf" name="bolt" mesh="bolt" rgba="0.7 0.7 0.7 1">
<plugin instance="bolt"/>
</geom>
</body>
<light name="left" pos="-1 0 2" cutoff="80"/>
<light name="right" pos="1 0 2" cutoff="80"/>
</worldbody>
</mujoco>
""",
}
_FIXTURES = {
"box_plane": """
<mujoco>
<worldbody>
<geom size="40 40 40" type="plane"/>
<body pos="0 0 0.3" euler="45 0 0">
<freejoint/>
<geom size="0.5 0.5 0.5" type="box"/>
</body>
</worldbody>
</mujoco>
""",
"box_box_vf": """
<mujoco>
<worldbody>
<body pos="0 -0.2 1.2" euler="44 46 0">
<freejoint/>
<geom size="0.6 0.4 0.7" type="box"/>
</body>
<body pos="0 0 0" euler="0 0 0">
<geom size="0.5 0.5 0.5" type="box"/>
</body>
</worldbody>
</mujoco>
""",
"box_box_vf_flat": """
<mujoco>
<worldbody>
<body pos="0 0 0" >
<geom size="0.6 0.6 0.5" type="box"/>
</body>
<body pos="-.28 0.4 1.199" >
<freejoint/>
<geom size="0.6 0.4 0.7" type="box"/>
</body>
</worldbody>
</mujoco>
""",
"box_box_ee": """
<mujoco>
<worldbody>
<body pos="0 0 0" euler="0 45 0">
<geom size="0.5 0.5 0.5" type="box"/>
</body>
<body pos="0 0 1.6" euler="44 0 90">
<freejoint/>
<geom size="0.6 0.4 0.7" type="box"/>
</body>
</worldbody>
</mujoco>
""",
"box_box_ee_deep": """
<mujoco>
<worldbody>
<body pos="0 0 0" euler="0 74 0">
<geom size="0.4 0.45 0.4" type="box"/>
</body>
<body pos="0 0 1.2" euler="24 0 90">
<freejoint/>
<geom size="0.6 0.4 0.7" type="box"/>
</body>
</worldbody>
</mujoco>
""",
"plane_sphere": """
<mujoco>
<worldbody>
<geom size="40 40 40" type="plane"/>
<body pos="0 0 0.2" euler="45 0 0">
<freejoint/>
<geom size="0.5" type="sphere"/>
</body>
</worldbody>
</mujoco>
""",
"plane_ellipsoid": """
<mujoco>
<worldbody>
<geom type="plane" size="10 10 .001"/>
<body pos="0 0 .299">
<geom type="ellipsoid" size=".1 .2 .3"/>
<freejoint/>
</body>
</worldbody>
</mujoco>
""",
"plane_capsule": """
<mujoco>
<worldbody>
<geom size="40 40 40" type="plane"/>
<body pos="0 0 0.0" euler="30 30 0">
<freejoint/>
<geom size="0.05 0.05" type="capsule"/>
</body>
</worldbody>
</mujoco>
""",
"convex_convex": """
<mujoco>
<asset>
<mesh name="poly"
vertex="0.3 0 0 0 0.5 0 -0.3 0 0 0 -0.5 0 0 -1 1 0 1 1"
face="0 1 5 0 5 4 0 4 3 3 4 2 2 4 5 1 2 5 0 2 1 0 3 2"/>
</asset>
<worldbody>
<body pos="0.0 2.0 0.35" euler="0 0 90">
<freejoint/>
<geom size="0.2 0.2 0.2" type="mesh" mesh="poly"/>
</body>
<body pos="0.0 2.0 2.281" euler="180 0 0">
<freejoint/>
<geom size="0.2 0.2 0.2" type="mesh" mesh="poly"/>
</body>
</worldbody>
</mujoco>
""",
"capsule_capsule": """
<mujoco model="two_capsules">
<worldbody>
<body>
<joint type="free"/>
<geom fromto="0.62235904 0.58846647 0.651046 1.5330081 0.33564585 0.977849"
size="0.05" type="capsule"/>
</body>
<body>
<joint type="free"/>
<geom fromto="0.5505271 0.60345304 0.476661 1.3900293 0.30709633 0.932082"
size="0.05" type="capsule"/>
</body>
</worldbody>
</mujoco>
""",
"sphere_sphere": """
<mujoco>
<worldbody>
<body>
<joint type="free"/>
<geom pos="0 0 0" size="0.2" type="sphere"/>
</body>
<body >
<joint type="free"/>
<geom pos="0 0.3 0" size="0.11" type="sphere"/>
</body>
</worldbody>
</mujoco>
""",
"sphere_capsule": """
<mujoco>
<worldbody>
<body>
<joint type="free"/>
<geom pos="0 0 0" size="0.25" type="sphere"/>
</body>
<body>
<joint type="free"/>
<geom fromto="0.3 0 0 0.7 0 0" size="0.1" type="capsule"/>
</body>
</worldbody>
</mujoco>
""",
"sphere_cylinder_corner": """
<mujoco>
<worldbody>
<body>
<joint type="slide" axis="1 0 0"/>
<joint type="slide" axis="0 1 0"/>
<joint type="slide" axis="0 0 1"/>
<geom size="0.1" type="sphere" pos=".33 0 0"/>
</body>
<body>
<geom size="0.15 0.2" type="cylinder" euler="30 45 0"/>
</body>
</worldbody>
</mujoco>
""",
"sphere_cylinder_cap": """
<mujoco>
<worldbody>
<body>
<joint type="slide" axis="1 0 0"/>
<joint type="slide" axis="0 1 0"/>
<joint type="slide" axis="0 0 1"/>
<geom size="0.1" type="sphere" pos=".26 -.14 .1"/>
</body>
<body>
<geom size="0.15 0.2" type="cylinder" euler="30 45 0"/>
</body>
</worldbody>
</mujoco>
""",
"sphere_cylinder_side": """
<mujoco>
<worldbody>
<body>
<joint type="slide" axis="1 0 0"/>
<joint type="slide" axis="0 1 0"/>
<joint type="slide" axis="0 0 1"/>
<geom size="0.1" type="sphere" pos="0 -.26 0"/>
</body>
<body>
<geom size="0.15 0.2" type="cylinder" euler="30 45 0"/>
</body>
</worldbody>
</mujoco>
""",
"plane_cylinder_1": """
<mujoco>
<worldbody>
<geom size="40 40 40" type="plane" euler="3 0 0"/>
<body pos="0 0 0.1" euler="30 30 0">
<freejoint/>
<geom size="0.05 0.1" type="cylinder"/>
</body>
</worldbody>
</mujoco>
""",
"plane_cylinder_2": """
<mujoco>
<worldbody>
<geom size="40 40 40" type="plane" euler="3 0 0"/>
<body pos="0.2 0 0.04" euler="90 0 0">
<freejoint/>
<geom size="0.05 0.1" type="cylinder"/>
</body>
</worldbody>
</mujoco>
""",
"plane_cylinder_3": """
<mujoco>
<worldbody>
<geom size="40 40 40" type="plane" euler="3 0 0"/>
<body pos="0.5 0 0.1" euler="3 0 0">
<freejoint/>
<geom size="0.05 0.1" type="cylinder"/>
</body>
</worldbody>
</mujoco>
""",
"mesh_plane_simple": """
<mujoco>
<asset>
<mesh name="cube" vertex="1 1 1 1 1 -1 1 -1 1 1 -1 -1 -1 1 1 -1 1 -1 -1 -1 1 -1 -1 -1"/>
</asset>
<worldbody>
<geom size="40 40 40" type="plane"/>
<body pos="0 0 1" euler="45 0 0">
<freejoint/>
<geom type="mesh" mesh="cube"/>
</body>
</worldbody>
</mujoco>
""",
"mesh_plane_complex": """
<mujoco>
<asset>
<mesh name="poly"
vertex="
0.5 0.0 0.0 0.4 0.3 0.0 0.2 0.45 0.0 0.0 0.5 0.0
-0.2 0.45 0.0 -0.4 0.3 0.0 -0.5 0.0 0.0 -0.4 -0.3 0.0
-0.2 -0.45 0.0 0.0 -0.5 0.0 0.2 -0.45 0.0 0.4 -0.3 0.0
0.0 0.0 1.0
"
face="
0 1 12 1 2 12 2 3 12 3 4 12 4 5 12 5 6 12
6 7 12 7 8 12 8 9 12 9 10 12 10 11 12 11 0 12
0 1 2 0 2 3 0 3 4 0 4 5 0 5 6 0 6 7 0 7 8 0 8 9 0 9 10 0 10 11 0 11 1
"/>
</asset>
<worldbody>
<geom size="40 40 40" type="plane"/>
<body pos="0.0 2.0 0.0" euler="90 90 0">
<freejoint/>
<geom size="0.2 0.2 0.2" type="mesh" mesh="poly"/>
</body>
</worldbody>
</mujoco>
""",
"sphere_box_shallow": """
<mujoco>
<worldbody>
<geom type="box" pos="0 0 0" size=".5 .5 .5" />
<body pos="-0.6 -0.6 0.7">
<geom type="sphere" size="0.5"/>
<freejoint/>
</body>
</worldbody>
</mujoco>
""",
"sphere_box_deep": """
<mujoco>
<worldbody>
<geom type="box" pos="0 0 0" size=".5 .5 .5" />
<body pos="-0.6 -0.6 0.7">
<geom type="sphere" size="0.5"/>
<freejoint/>
</body>
</worldbody>
</mujoco>
""",
"capsule_box_edge": """
<mujoco>
<worldbody>
<geom type="box" pos="0 0 0" size=".5 .4 .9" />
<body pos="0.4 0.2 0.8" euler="0 -40 0" >
<geom type="capsule" size="0.5 0.8"/>
<freejoint/>
</body>
</worldbody>
</mujoco>
""",
"capsule_box_corner": """
<mujoco>
<worldbody>
<geom type="box" pos="0 0 0" size=".5 .55 .6" />
<body pos="0.55 0.6 0.65" euler="0 0 0" >
<geom type="capsule" size="0.4 0.6"/>
<freejoint/>
</body>
</worldbody>
</mujoco>
""",
"capsule_box_face_tip": """
<mujoco>
<worldbody>
<geom type="box" pos="0 0 0" size=".5 .4 .9" />
<body pos="0 0 1.5" >
<geom type="capsule" size="0.5 0.8"/>
<freejoint/>
</body>
</worldbody>
</mujoco>
""",
"capsule_box_face_flat": """
<mujoco>
<worldbody>
<geom type="box" pos="0 0 0" size=".5 .7 .9" />
<body pos="0.5 0.2 0.0" euler="0 0 0" >
<geom type="capsule" size="0.2 0.4"/>
<freejoint/>
</body>
</worldbody>
</mujoco>
""",
}
@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": """
<mujoco>
<asset>
<hfield name="terrain" nrow="2" ncol="2" size="1 1 0.1 0.1"
elevation="0 0
0 0"/>
</asset>
<worldbody>
<geom type="hfield" hfield="terrain" pos="0 0 0"/>
<body pos=".0 .0 .1">
<freejoint/>
<geom type="box" size=".1 .1 .11"/>
</body>
</worldbody>
</mujoco>
""",
}
@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="""
<mujoco>
<worldbody>
<body name="body1">
<freejoint/>
<geom type="sphere" size=".1"/>
</body>
<body name="body2">
<freejoint/>
<geom type="sphere" size=".1"/>
</body>
<body name="body3">
<freejoint/>
<geom type="sphere" size=".1"/>
</body>
</worldbody>
<contact>
<exclude body1="body1" body2="body2"/>
</contact>
</mujoco>
"""
)
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="""
<mujoco>
<worldbody>
<body>
<freejoint/>
<geom type="sphere" size=".1"/>
</body>
</worldbody>
</mujoco>
"""
)
self.assertTrue((m.nxn_pairid.numpy() == -1).all())
# 1 pair
_, _, m, d = test_util.fixture(
xml="""
<mujoco>
<worldbody>
<body>
<freejoint/>
<geom name="geom1" type="sphere" size=".1"/>
</body>
<body>
<freejoint/>
<geom name="geom2" type="sphere" size=".1"/>
</body>
</worldbody>
<contact>
<pair geom1="geom1" geom2="geom2" margin="2" gap="3" condim="6" friction="5 4 3 2 1" solref="-.25 -.5" solreffriction="2 4" solimp=".1 .2 .3 .4 .5"/>
</contact>
</mujoco>
""",
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="""
<mujoco>
<worldbody>
<body name="body1">
<freejoint/>
<geom name="geom1" type="sphere" size=".1" contype="0" conaffinity="0"/>
</body>
<body name="body2">
<freejoint/>
<geom name="geom2" type="sphere" size=".1" contype="0" conaffinity="0"/>
</body>
</worldbody>
<contact>
<pair geom1="geom1" geom2="geom2" margin="2" gap="3" condim="6" friction="5 4 3 2 1" solref="-.25 -.5" solreffriction="2 4" solimp=".1 .2 .3 .4 .5"/>
</contact>
</mujoco>
""",
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="""
<mujoco>
<worldbody>
<body name="body1">
<freejoint/>
<geom name="geom1" type="sphere" size=".1"/>
</body>
<body name="body2">
<freejoint/>
<geom name="geom2" type="sphere" size=".1"/>
</body>
</worldbody>
<contact>
<exclude body1="body1" body2="body2"/>
<pair geom1="geom1" geom2="geom2" margin="2" gap="3" condim="6" friction="5 4 3 2 1" solref="-.25 -.5" solreffriction="2 4" solimp=".1 .2 .3 .4 .5"/>
</contact>
</mujoco>
""",
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="""
<mujoco>
<worldbody>
<body name="body1">
<freejoint/>
<geom name="geom1" type="sphere" size=".1"/>
</body>
<body name="body2">
<freejoint/>
<geom name="geom2" type="sphere" size=".1"/>
</body>
<body name="body3">
<freejoint/>
<geom name="geom3" type="sphere" size=".1"/>
</body>
</worldbody>
<contact>
<exclude body1="body1" body2="body2"/>
<pair geom1="geom2" geom2="geom3" margin="2" gap="3" condim="6" friction="5 4 3 2 1" solref="-.25 -.5" solreffriction="2 4" solimp=".1 .2 .3 .4 .5"/>
</contact>
</mujoco>
""",
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"""
<mujoco>
<asset>
<hfield name="hfield" nrow="10" ncol="10" size="1e-6 1e-6 1 1"/>
</asset>
<worldbody>
<body>
<joint type="slide" axis="0 0 1"/>
<geom type="sphere" size=".1"/>
</body>
<geom type="hfield" hfield="hfield"/>
</worldbody>
<keyframe>
<key qpos=".0999"/>
</keyframe>
</mujoco>
"""
_, _, 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="""
<mujoco>
<worldbody>
<body>
<geom type="sphere" size=".1" friction="0 0 0"/>
<joint type="slide"/>
</body>
<body>
<geom type="sphere" size=".1" friction="0 0 0"/>
<joint type="slide"/>
</body>
</worldbody>
<keyframe>
<key qpos="0 .1"/>
</keyframe>
</mujoco>
""",
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()
File diff suppressed because it is too large Load Diff
@@ -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
@@ -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"""
<mujoco>
<worldbody>
<geom name="geom1" type="sphere" pos="-1.5 0 0" size="1"/>
<geom name="geom2" type="sphere" pos="1.5 0 0" size="1"/>
</worldbody>
</mujoco>
"""
)
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"""
<mujoco>
<worldbody>
<geom type="sphere" pos="-1 0 0" size="1"/>
<geom type="sphere" pos="1 0 0" size="1"/>
</worldbody>
</mujoco>
"""
)
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"""
<mujoco model="MuJoCo Model">
<asset>
<mesh name="smallbox" scale="0.1 0.1 0.1"
vertex="-1 -1 -1
1 -1 -1
1 1 -1
1 1 1
1 -1 1
-1 1 -1
-1 1 1
-1 -1 1"/>
</asset>
<worldbody>
<geom pos="0 0 .90" type="box" size="0.5 0.5 0.1"/>
<geom pos="0 0 1.2" type="mesh" mesh="smallbox"/>
</worldbody>
</mujoco>
"""
)
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"""
<mujoco>
<worldbody>
<geom type="sphere" pos="-1 0 0" size="3"/>
<geom type="sphere" pos=" 3 0 0" size="3"/>
</worldbody>
</mujoco>
"""
)
# 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"""
<mujoco>
<worldbody>
<geom type="box" pos="-1 0 0" size="2.5 2.5 2.5"/>
<geom type="box" pos="1.5 0 0" size="1 1 1"/>
</worldbody>
</mujoco>
"""
)
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"""
<mujoco>
<asset>
<mesh name="box" scale=".5 .5 .1"
vertex="-1 -1 -1
1 -1 -1
1 1 -1
1 1 1
1 -1 1
-1 1 -1
-1 1 1
-1 -1 1"/>
<mesh name="smallbox" scale=".1 .1 .1"
vertex="-1 -1 -1
1 -1 -1
1 1 -1
1 1 1
1 -1 1
-1 1 -1
-1 1 1
-1 -1 1"/>
</asset>
<worldbody>
<geom pos="0 0 .09" type="mesh" mesh="smallbox"/>
<geom pos="0 0 -.1" type="mesh" mesh="box"/>
</worldbody>
</mujoco>
"""
)
dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
self.assertAlmostEqual(-0.01, dist)
if __name__ == "__main__":
wp.init()
absltest.main()
@@ -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,
],
)
File diff suppressed because it is too large Load Diff
@@ -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,
],
)
File diff suppressed because it is too large Load Diff
@@ -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"""
<mujoco>
<worldbody>
<body pos="0.0 0 0">
<freejoint/>
<geom type="sphere" size=".1" condim="{condim1}"/>
</body>
<body pos="0.05 0 0">
<freejoint/>
<geom type="sphere" size=".1" condim="{condim2}"/>
</body>
</worldbody>
</mujoco>
"""
_, 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="""
<mujoco>
<option>
<flag contact="disable"/>
</option>
<worldbody>
<body>
<geom type="sphere" size=".1"/>
<joint name="joint0" type="hinge"/>
</body>
<body>
<geom type="sphere" size=".1"/>
<joint name="joint1" type="hinge"/>
</body>
<body>
<geom type="sphere" size=".1"/>
<joint name="joint2" type="hinge"/>
</body>
</worldbody>
<tendon>
<fixed name="tendon0">
<joint joint="joint0" coef=".1"/>
</fixed>
<fixed name="tendon1">
<joint joint="joint1" coef=".2"/>
</fixed>
<fixed name="tendon2">
<joint joint="joint2" coef=".3"/>
</fixed>
</tendon>
<equality>
<tendon tendon1="tendon0" tendon2="tendon1" polycoef=".1 .2 .3 .4 .5"/>
<tendon tendon1="tendon2" polycoef="-.1 0 0 0 0"/>
</equality>
<keyframe>
<key qpos=".1 .2 .3"/>
</keyframe>
</mujoco>
"""
)
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()
@@ -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
File diff suppressed because it is too large Load Diff
@@ -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="""
<mujoco>
<worldbody>
<body>
<joint name="joint" type="slide"/>
<geom type="sphere" size=".1"/>
</body>
</worldbody>
<actuator>
<motor joint="joint" ctrlrange="-1 1"/>
</actuator>
<keyframe>
<key ctrl="2"/>
</keyframe>
</mujoco>
""",
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="""
<mujoco>
<option integrator="RK4" iterations="1" ls_iterations="1">
<flag constraint="disable"/>
</option>
<worldbody>
<body>
<joint type="hinge"/>
<geom type="sphere" size=".1"/>
<body pos="0.1 0 0">
<joint type="hinge"/>
<geom type="sphere" size=".1"/>
</body>
</body>
</worldbody>
<keyframe>
<key qpos=".1 .2" qvel=".025 .05"/>
</keyframe>
</mujoco>
""",
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()
+147
View File
@@ -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)
@@ -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 = """
<mujoco>
<option timestep=".01" gravity="-1 -1 -1"/>
<worldbody>
<body>
<geom type="sphere" size=".1" pos=".5 0 0"/>
<joint name="joint1" type="hinge" axis="0 1 0" damping=".1"/>
<body>
<geom type="sphere" size=".2" pos="1 0 0"/>
<joint name="joint2" type="hinge" axis="0 1 0" damping=".2"/>
</body>
</body>
</worldbody>
<actuator>
<motor joint="joint1"/>
</actuator>
<equality>
<joint joint1="joint1" joint2="joint2"/>
</equality>
</mujoco>
"""
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()
File diff suppressed because it is too large Load Diff
+354
View File
@@ -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("""
<mujoco>
<worldbody>
<body pos="0 0 0" >
<geom type="box" pos="0 0 0" size=".5 .5 .5" />
<joint type="hinge" />
</body>
<body pos="0 0 0.1">
<geom type="sphere" size="0.5"/>
<freejoint/>
</body>
</worldbody>
</mujoco>
""")
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(
"""
<mujoco>
<option density="1"/>
<worldbody>
<body>
<geom type="sphere" size=".1" fluidshape="ellipsoid"/>
<freejoint/>
</body>
</worldbody>
</mujoco>
"""
)
mjwarp.put_model(mjm)
def test_jacobian_auto(self):
mjm = mujoco.MjModel.from_xml_string("""
<mujoco>
<option jacobian="auto"/>
<worldbody>
<replicate count="11">
<body>
<geom type="sphere" size=".1"/>
<freejoint/>
</body>
</replicate>
</worldbody>
</mujoco>
""")
mjwarp.put_model(mjm)
def test_put_data_qLD(self):
mjm = mujoco.MjModel.from_xml_string("""
<mujoco>
<worldbody>
<body>
<geom type="sphere" size="1"/>
<joint type="hinge"/>
</body>
</worldbody>
</mujoco>
""")
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="""
<mujoco>
<option noslip_iterations="1"/>
</mujoco>
"""
)
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()
+101
View File
@@ -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()
+289
View File
@@ -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
+131
View File
@@ -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()
+617
View File
@@ -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,
],
)
@@ -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"""
<mujoco>
<option density="{density}" viscosity="{viscosity}" wind="{wind0} {wind1} {wind2}"/>
<worldbody>
<body>
<geom type="box" size=".1 .1 .1"/>
<freejoint/>
</body>
</worldbody>
<keyframe>
<key qvel="1 1 1 1 1 1"/>
</keyframe>
</mujoco>
""",
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="""
<mujoco>
<option gravity="1 2 3">
<flag contact="disable"/>
</option>
<worldbody>
<body gravcomp="1">
<geom type="sphere" size=".1" pos="1 0 0"/>
<joint name="joint0" type="hinge" axis="0 1 0" actuatorgravcomp="true"/>
</body>
<body gravcomp="1">
<geom type="sphere" size=".1"/>
<joint name="joint1" type="hinge" axis="1 0 0"/>
<joint type="hinge" axis="0 1 0"/>
<joint type="hinge" axis="0 0 1"/>
</body>
<body gravcomp="1">
<geom type="sphere" size=".1"/>
<joint type="hinge" axis="0 1 0"/>
</body>
<body gravcomp="0">
<geom type="sphere" size=".1"/>
<joint type="hinge" axis="0 1 0"/>
</body>
</worldbody>
<actuator>
<motor joint="joint0"/>
<motor joint="joint1"/>
</actuator>
</mujoco>
""",
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()
+903
View File
@@ -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,
)
+312
View File
@@ -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()
File diff suppressed because it is too large Load Diff
@@ -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="""
<mujoco>
<option gravity="-1 -1 -1"/>
<worldbody>
<body name="body0" pos="0.1 0.2 0.3" quat=".05 .1 .15 .2">
<joint name="slide" type="slide"/>
<geom name="geom0" type="sphere" size="0.1"/>
<site name="site0"/>
<camera name="cam0"/>
</body>
<body name="body1" pos=".5 .6 .7">
<joint name="ballquat" type="ball"/>
<geom type="sphere" size="0.1"/>
<site name="site1" pos=".1 .2 .3"/>
</body>
<body name="body2" pos="1 1 1">
<freejoint/>
<geom type="sphere" size="0.1"/>
<site name="site2"/>
</body>
<body name="body3" pos="2 2 2">
<joint name="hinge0" type="hinge" axis="1 0 0"/>
<geom type="sphere" size="0.1" pos=".1 0 0"/>
<body pos="2 2 2">
<joint name="hinge1" type="hinge" axis="1 0 0"/>
<geom type="sphere" size="0.1" pos=".1 0 0"/>
</body>
</body>
<body name="body4" pos="1 0 0">
<joint type="ball"/>
<geom type="sphere" size="0.1" pos=".1 0 0"/>
<body>
<joint type="ball"/>
<geom type="sphere" size="0.1" pos=".1 0 0"/>
</body>
</body>
<body pos="10 0 0">
<joint type="hinge" axis="1 2 3"/>
<geom type="sphere" size="0.1"/>
<site name="force_site" pos="1 2 3"/>
</body>
<body pos="20 0 0">
<joint type="slide" axis="1 2 3"/>
<geom type="sphere" size="0.1"/>
<site name="torque_site" pos="1 2 3"/>
</body>
<body name="body8">
<joint type="hinge"/>
<geom type="sphere" size="0.1" pos="1 2 3"/>
<body name="body9">
<joint type="hinge"/>
<geom name="geom9" type="sphere" size="0.1" pos="1 2 3"/>
<site name="site9" pos=".2 .4 .6"/>
</body>
</body>
<camera name="camera"/>
<site name="camera_site" pos="0 0 -1"/>
<!-- limit pos: slide -->
<body pos="1 1 1">
<joint name="limitslide" type="slide" limited="true" range="-.5 .5" margin=".1"/>
<geom type="sphere" size=".1"/>
</body>
<!-- limit pos: hinge -->
<body pos="2 2 2">
<joint name="limithinge" type="hinge" limited="true" range="-.4 .4" margin=".09"/>
<geom type="sphere" size=".1"/>
</body>
<!-- limit pos: ball -->
<body pos="3 3 3">
<joint name="limitball" type="ball" limited="true" range="0 .1" margin=".05"/>
<geom type="sphere" size=".1"/>
</body>
<!-- tendon limit pos -->
<site name="sitetendon0" pos="5 4 4"/>
<body pos="4 4 4">
<joint type="slide" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
<site name="sitetendon1"/>
</body>
<body pos="3 3 3">
<joint name="tendonjoint" type="hinge"/>
<geom type="sphere" size=".1"/>
</body>
</worldbody>
<tendon>
<spatial name="limittendon" limited="true" range="0 .5" margin=".1">
<site site="sitetendon0"/>
<site site="sitetendon1"/>
</spatial>
<fixed name="tendon">
<joint joint="tendonjoint" coef=".123"/>
</fixed>
</tendon>
<actuator>
<motor name="slide" joint="slide"/>
<motor tendon="tendon"/>
<motor tendon="tendon" gear="2"/>
</actuator>
<sensor>
<magnetometer site="site0"/>
<camprojection camera="camera" site="camera_site"/>
<camprojection camera="camera" site="camera_site" cutoff=".001"/>
<jointpos joint="slide"/>
<jointpos joint="slide" cutoff=".001"/>
<actuatorpos actuator="slide"/>
<actuatorpos actuator="slide" cutoff=".001"/>
<ballquat joint="ballquat"/>
<jointlimitpos joint="limitslide"/>
<jointlimitpos joint="limitslide" cutoff=".001"/>
<jointlimitpos joint="limithinge"/>
<jointlimitpos joint="limithinge" cutoff=".001"/>
<jointlimitpos joint="limitball"/>
<jointlimitpos joint="limitball" cutoff=".001"/>
<tendonlimitpos tendon="limittendon"/>
<tendonlimitpos tendon="limittendon" cutoff=".001"/>
<framepos objtype="body" objname="body1"/>
<framepos objtype="body" objname="body1" cutoff=".001"/>
<framepos objtype="body" objname="body1" reftype="body" refname="body0"/>
<framepos objtype="body" objname="body1" reftype="geom" refname="geom0"/>
<framepos objtype="body" objname="body1" reftype="site" refname="site0"/>
<framepos objtype="body" objname="body1" reftype="cam" refname="cam0"/>
<framepos objtype="xbody" objname="body1"/>
<framepos objtype="geom" objname="geom0"/>
<framepos objtype="site" objname="site0"/>
<framepos objtype="camera" objname="cam0"/>
<framexaxis objtype="body" objname="body1"/>
<framexaxis objtype="body" objname="body1" reftype="body" refname="body0"/>
<framexaxis objtype="body" objname="body1" reftype="geom" refname="geom0"/>
<framexaxis objtype="body" objname="body1" reftype="site" refname="site0"/>
<framexaxis objtype="body" objname="body1" reftype="cam" refname="cam0"/>
<framexaxis objtype="xbody" objname="body1"/>
<framexaxis objtype="geom" objname="geom0"/>
<framexaxis objtype="site" objname="site0"/>
<framexaxis objtype="camera" objname="cam0"/>
<frameyaxis objtype="body" objname="body1"/>
<frameyaxis objtype="body" objname="body1" reftype="body" refname="body0"/>
<frameyaxis objtype="body" objname="body1" reftype="geom" refname="geom0"/>
<frameyaxis objtype="body" objname="body1" reftype="site" refname="site0"/>
<frameyaxis objtype="body" objname="body1" reftype="cam" refname="cam0"/>
<frameyaxis objtype="xbody" objname="body1"/>
<frameyaxis objtype="geom" objname="geom0"/>
<frameyaxis objtype="site" objname="site0"/>
<frameyaxis objtype="camera" objname="cam0"/>
<framezaxis objtype="body" objname="body1"/>
<framezaxis objtype="body" objname="body1" reftype="body" refname="body0"/>
<framezaxis objtype="body" objname="body1" reftype="geom" refname="geom0"/>
<framezaxis objtype="body" objname="body1" reftype="site" refname="site0"/>
<framezaxis objtype="body" objname="body1" reftype="cam" refname="cam0"/>
<framezaxis objtype="xbody" objname="body1"/>
<framezaxis objtype="geom" objname="geom0"/>
<framezaxis objtype="site" objname="site0"/>
<framezaxis objtype="camera" objname="cam0"/>
<framequat objtype="body" objname="body1"/>
<framequat objtype="body" objname="body1" reftype="body" refname="body0"/>
<framequat objtype="body" objname="body1" reftype="geom" refname="geom0"/>
<framequat objtype="body" objname="body1" reftype="site" refname="site0"/>
<framequat objtype="body" objname="body1" reftype="cam" refname="cam0"/>
<framequat objtype="xbody" objname="body1"/>
<framequat objtype="geom" objname="geom0"/>
<framequat objtype="site" objname="site0"/>
<framequat objtype="camera" objname="cam0"/>
<subtreecom body="body3"/>
<subtreecom body="body3" cutoff=".001"/>
<e_potential/>
<e_potential cutoff=".001"/>
<e_kinetic/>
<e_kinetic cutoff=".001"/>
<clock/>
<clock cutoff=".001"/>
<velocimeter site="site2"/>
<velocimeter site="site2" cutoff=".001"/>
<gyro site="site2"/>
<gyro site="site2" cutoff=".001"/>
<jointvel joint="slide"/>
<jointvel joint="slide" cutoff=".001"/>
<actuatorvel actuator="slide"/>
<actuatorvel actuator="slide" cutoff=".001"/>
<ballangvel joint="ballquat"/>
<ballangvel joint="ballquat" cutoff=".001"/>
<jointlimitvel joint="limithinge"/>
<jointlimitvel joint="limithinge" cutoff=".001"/>
<tendonlimitvel tendon="limittendon"/>
<tendonlimitvel tendon="limittendon" cutoff=".001"/>
<framelinvel objtype="body" objname="body9"/>
<framelinvel objtype="body" objname="body9" cutoff=".001"/>
<frameangvel objtype="body" objname="body9"/>
<frameangvel objtype="body" objname="body9" cutoff=".001"/>
<framelinvel objtype="xbody" objname="body9"/>
<frameangvel objtype="xbody" objname="body9"/>
<framelinvel objtype="geom" objname="geom9"/>
<frameangvel objtype="geom" objname="geom9"/>
<framelinvel objtype="site" objname="site9"/>
<frameangvel objtype="site" objname="site9"/>
<framelinvel objtype="camera" objname="cam0"/>
<frameangvel objtype="camera" objname="cam0"/>
<framelinvel objtype="body" objname="body9" reftype="xbody" refname="body0"/>
<frameangvel objtype="body" objname="body9" reftype="xbody" refname="body0"/>
<framelinvel objtype="body" objname="body9" reftype="geom" refname="geom0"/>
<frameangvel objtype="body" objname="body9" reftype="geom" refname="geom0"/>
<framelinvel objtype="body" objname="body9" reftype="site" refname="site0"/>
<frameangvel objtype="body" objname="body9" reftype="site" refname="site0"/>
<framelinvel objtype="body" objname="body9" reftype="cam" refname="cam0"/>
<frameangvel objtype="body" objname="body9" reftype="cam" refname="cam0"/>
<subtreelinvel body="body4"/>
<subtreelinvel body="body4" cutoff=".001"/>
<subtreeangmom body="body4"/>
<subtreeangmom body="body4" cutoff=".001"/>
<accelerometer site="force_site"/>
<accelerometer site="force_site" cutoff=".001"/>
<force site="force_site"/>
<force site="force_site" cutoff=".001"/>
<torque site="torque_site"/>
<torque site="torque_site" cutoff=".001"/>
<actuatorfrc actuator="slide"/>
<actuatorfrc actuator="slide" cutoff=".001"/>
<jointactuatorfrc joint="slide"/>
<jointactuatorfrc joint="slide" cutoff=".001"/>
<jointlimitfrc joint="limitslide"/>
<jointlimitfrc joint="limitslide" cutoff=".001"/>
<jointlimitfrc joint="limithinge"/>
<jointlimitfrc joint="limithinge" cutoff=".001"/>
<jointlimitfrc joint="limitball"/>
<jointlimitfrc joint="limitball" cutoff=".001"/>
<tendonlimitfrc tendon="limittendon"/>
<tendonlimitfrc tendon="limittendon" cutoff=".001"/>
<tendonactuatorfrc tendon="tendon"/>
<tendonactuatorfrc tendon="tendon" cutoff=".001"/>
<framelinacc objtype="body" objname="body9"/>
<framelinacc objtype="body" objname="body9" cutoff=".001"/>
<frameangacc objtype="body" objname="body9"/>
<frameangacc objtype="body" objname="body9" cutoff=".001"/>
<framelinacc objtype="xbody" objname="body9"/>
<frameangacc objtype="xbody" objname="body9"/>
<framelinacc objtype="geom" objname="geom9"/>
<frameangacc objtype="geom" objname="geom9"/>
<framelinacc objtype="site" objname="site9"/>
<frameangacc objtype="site" objname="site9"/>
<framelinacc objtype="camera" objname="cam0"/>
<frameangacc objtype="camera" objname="cam0"/>
</sensor>
<keyframe>
<key qpos="1 .1 .2 .3 .4 1 1 1 1 0 0 0 .25 .35 1 0 0 0 1 0 0 0 0 0 1 1 .6 .5 1 2 3 4 .5 0" qvel="2 .2 -.1 .4 .25 .35 .45 -0.1 -0.2 -0.3 .1 -.2 -.5 -0.75 -1 .1 .2 .3 0 0 2 2 0 0 0 0 0 0 0" ctrl="3 .1 .2"/>
</keyframe>
</mujoco>
""",
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="""
<mujoco>
<compiler angle="degree"/>
<worldbody>
<geom type="sphere" size=".1" pos="0 0 1"/>
<body>
<joint name="joint0" type="hinge" axis="1 0 0"/>
<joint type="hinge" axis="0 1 0"/>
<joint type="hinge" axis="0 0 1"/>
<geom type="sphere" size="0.1"/>
<site name="site0" size=".1"/>
<site name="site1" size=".1" pos="0 0 -.2"/>
</body>
</worldbody>
<sensor>
<rangefinder site="site0"/>
<jointpos joint="joint0"/>
<rangefinder site="site1"/>
</sensor>
<keyframe>
<key qpos="0 0 0"/>
<key qpos="0 90 0"/>
</keyframe>
</mujoco>
""",
)
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="""
<mujoco>
<worldbody>
<geom type="plane" size="10 10 .001"/>
<body pos="0 0 .25">
<geom type="sphere" size="0.1" pos=".1 0 0"/>
<geom type="sphere" size="0.1" pos="-.1 0 0"/>
<geom type="sphere" size="0.1" pos="0 0 .11"/>
<geom type="sphere" size="0.1" pos="0 0 10"/>
<site name="site_sphere" type="sphere" size=".2"/>
<site name="site_capsule" type="capsule" size=".2 .2"/>
<site name="site_ellipsoid" type="ellipsoid" size=".2 .2 .2"/>
<site name="site_cylinder" type="cylinder" size=".2 .2"/>
<site name="site_box" type="box" size=".2 .2 .2"/>
<freejoint/>
</body>
</worldbody>
<sensor>
<touch site="site_sphere"/>
<touch site="site_capsule"/>
<touch site="site_ellipsoid"/>
<touch site="site_cylinder"/>
<touch site="site_box"/>
</sensor>
<keyframe>
<key qpos="0 0 10 1 0 0 0"/>
<key qpos="0 0 .05 1 0 0 0"/>
<key qpos="0 0 0 1 0 0 0"/>
</keyframe>
</mujoco>
""",
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()
File diff suppressed because it is too large Load Diff
@@ -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 = """
<mujoco>
<option gravity="1 1 -1">
<flag contact="disable"/>
</option>
<worldbody>
<site name="siteworld"/>
<body name="body0">
<geom type="sphere" size=".1"/>
<freejoint/>
</body>
<body name="body1">
<geom type="sphere" size=".1"/>
<site name="site1"/>
<freejoint/>
</body>
<body name="body2">
<geom type="sphere" size=".1"/>
<freejoint/>
</body>
<body name="body3">
<geom type="sphere" size=".1"/>
<site name="site3" quat="0 1 0 0"/>
<freejoint/>
</body>
</worldbody>
<equality>
<connect body1="body0" anchor="1 1 1"/>
<connect site1="siteworld" site2="site1"/>
<weld body1="body2" relpose="1 1 1 0 1 0 0"/>
<weld site1="siteworld" site2="site3"/>
</equality>
<keyframe>
<key qpos="0 0 0 1 0 0 0 1 1 1 1 0 0 0 0 0 0 1 0 0 0 1 1 1 1 0 0 0"/>
</keyframe>
</mujoco>
"""
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="""
<mujoco>
<worldbody>
<body>
<geom type="sphere" size=".1"/>
<freejoint/>
</body>
</worldbody>
</mujoco>
""",
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()
File diff suppressed because it is too large Load Diff
@@ -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()
+554
View File
@@ -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
@@ -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"""
<mujoco>
<worldbody>
<geom type="plane" size="10 10 .001"/>
<body pos="0 0 1">
<freejoint/>
<geom fromto="-.4 0 0 .4 0 0" size=".05 .1" type="capsule" condim="{condim}" friction="1 1 1"/>
</body>
</worldbody>
<keyframe>
<key qpos="0 0 0.04 1 0 0 0" qvel="-1 -1 -1 .1 .1 .1"/>
</keyframe>
</mujoco>
"""
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()
+267
View File
@@ -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
File diff suppressed because it is too large Load Diff
+603
View File
@@ -0,0 +1,603 @@
# 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.
# ==============================================================================
"""Miscellaneous utilities."""
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.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
@wp.func
def is_intersect(p1: wp.vec2, p2: wp.vec2, p3: wp.vec2, p4: wp.vec2) -> 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)
@@ -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()
+196
View File
@@ -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.<locals>.nested_kernel"
qualname = f.__qualname__
parts = [part for part in qualname.split(".") if part != "<locals>"]
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
@@ -0,0 +1,40 @@
<mujoco>
<option iterations="4" ls_iterations="4"/>
<worldbody>
<body pos="0 0 0">
<joint name="hinge0" type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
<body>
<joint name="hinge1" type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
<body>
<joint name="hinge2" type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
</body>
</body>
</body>
<body pos="1 0 0">
<joint name="hinge3" type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
<body>
<joint name="hinge4" type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
<body>
<joint name="hinge5" type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
</body>
</body>
</body>
</worldbody>
<actuator>
<motor joint="hinge4"/>
<motor joint="hinge1"/>
<motor joint="hinge2"/>
<motor joint="hinge5"/>
<motor joint="hinge3"/>
<motor joint="hinge0"/>
</actuator>
<keyframe>
<key ctrl="1 2 3 4 5 6"/>
</keyframe>
</mujoco>
@@ -0,0 +1,28 @@
<mujoco>
<option iterations="4" ls_iterations="4"/>
<worldbody>
<body>
<joint name="hinge" type="hinge"/>
<geom type="sphere" size=".1"/>
</body>
</worldbody>
<actuator>
<!-- gaintype -->
<general joint="hinge" gaintype="fixed" gainprm="1.2345 0 0 0 0 0 0 0 0 0"/>
<general joint="hinge" gaintype="affine" gainprm="1.2345 2.3456 3.4567 0 0 0 0 0 0 0"/>
<!-- biastype -->
<general joint="hinge" biastype="none"/>
<general joint="hinge" biastype="affine" biasprm="1.2345 2.3456 3.4567 0 0 0 0 0 0 0"/>
<!-- dyntype -->
<general joint="hinge" dyntype="none"/>
<general joint="hinge" dyntype="integrator"/>
<general joint="hinge" dyntype="integrator" actearly="true"/>
<general joint="hinge" dyntype="filter"/>
<general joint="hinge" dyntype="filter" actearly="true"/>
<general joint="hinge" dyntype="filterexact"/>
<general joint="hinge" dyntype="filterexact" actearly="true"/>
</actuator>
<keyframe>
<key ctrl="1 1 1 1 1 1 1 1 1 1 1" act=".1 .1 .2 .2 .3 .3" qpos=".1234" qvel=".2345"/>
</keyframe>
</mujoco>
@@ -0,0 +1,21 @@
<mujoco>
<worldbody>
<geom type="plane" size="10 10 .001"/>
<body name="body">
<geom type="sphere" size=".1" pos=".1 0 0" margin=".01"/>
<geom type="sphere" size=".1" pos="0 .1 0" margin=".01" gap=".01"/>
<geom type="sphere" size=".1" pos=".1 .1 0" gap=".01"/>
<joint type="slide" axis="0 0 1"/>
<joint type="hinge" axis="1 0 1"/>
</body>
</worldbody>
<actuator>
<adhesion body="body" gain=".123" ctrlrange="0 1"/>
</actuator>
<keyframe>
<key qpos=".11 0"/>
<key qpos=".12 0"/>
<key qpos=".09 0"/>
<key qpos=".08 0"/>
</keyframe>
</mujoco>
@@ -0,0 +1,22 @@
<mujoco>
<worldbody>
<body>
<joint type="slide" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
<site name="site"/>
</body>
<site name="siteworld" pos="1 0 0"/>
</worldbody>
<tendon>
<spatial name="spatial">
<site site="site"/>
<site site="siteworld"/>
</spatial>
</tendon>
<actuator>
<muscle tendon="spatial" lengthrange="0 1"/>
</actuator>
<keyframe>
<key ctrl=".123" act=".456" qpos=".1" qvel=".01"/>
</keyframe>
</mujoco>
@@ -0,0 +1,40 @@
<mujoco>
<option iterations="4" ls_iterations="4"/>
<worldbody>
<body pos="0 0 0">
<joint name="hinge0" type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
<body>
<joint name="hinge1" type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
<body>
<joint name="hinge2" type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
</body>
</body>
</body>
<body pos="1 0 0">
<joint name="hinge3" type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
<body>
<joint name="hinge4" type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
<body>
<joint name="hinge5" type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1"/>
</body>
</body>
</body>
</worldbody>
<actuator>
<position joint="hinge0" kp="1" kv="2"/>
<position joint="hinge3" kp="7" kv="8"/>
<position joint="hinge4" kp="9" kv="10"/>
<position joint="hinge5" kp="11" kv="12"/>
<position joint="hinge2" kp="5" kv="6"/>
<position joint="hinge1" kp="3" kv="4"/>
</actuator>
<keyframe>
<key ctrl="1 2 3 4 5 6"/>
</keyframe>
</mujoco>
@@ -0,0 +1,31 @@
<mujoco>
<worldbody>
<site name="siteworld"/>
<body>
<joint type="hinge" axis="1 0 0"/>
<geom type="sphere" size=".1" pos="0 .1 0"/>
<site name="site0"/>
<body>
<joint type="hinge" axis="0 1 0"/>
<geom type="sphere" size=".1" pos="0 0 .2"/>
<site name="site1"/>
<body>
<joint type="hinge" axis="0 0 1"/>
<geom type="sphere" size=".1" pos=".3 0 0"/>
<site name="site2"/>
</body>
</body>
</body>
</worldbody>
<actuator>
<motor site="site0" gear="1 2 3 4 5 6"/>
<motor site="site0" refsite="siteworld" gear="1 2 3 4 5 6"/>
<motor site="site1" gear="1 2 3 4 5 6"/>
<motor site="site1" refsite="siteworld" gear="1 2 3 4 5 6"/>
<motor site="site2" gear="1 2 3 4 5 6"/>
<motor site="site2" refsite="siteworld" gear="1 2 3 4 5 6"/>
</actuator>
<keyframe>
<key qpos=".1 .2 .3"/>
</keyframe>
</mujoco>
@@ -0,0 +1,15 @@
<mujoco>
<compiler angle="degree"/>
<worldbody>
<site name="slidersite" euler="0 90 0"/>
<body pos="1 0 0">
<joint type="hinge" axis="0 0 1"/>
<geom type="sphere" size=".05" pos="0 0.2 0"/>
<geom type="cylinder" size=".2 .01"/>
<site name="cranksite" pos="0 0.2 0"/>
</body>
</worldbody>
<actuator>
<motor cranklength="0.25" cranksite="cranksite" slidersite="slidersite"/>
</actuator>
</mujoco>
@@ -0,0 +1,26 @@
<mujoco>
<worldbody>
<body>
<joint name="joint" type="hinge"/>
<geom type="sphere" size=".1"/>
</body>
</worldbody>
<tendon>
<fixed name="fixed" actuatorfrclimited="true" actuatorfrcrange="-1 1">
<joint joint="joint" coef="1"/>
</fixed>
</tendon>
<actuator>
<motor tendon="fixed"/>
<motor tendon="fixed"/>
</actuator>
<keyframe>
<key ctrl="0 0"/>
<key ctrl="1 1"/>
<key ctrl="2 0"/>
<key ctrl="0 2"/>
<key ctrl=".5 .5"/>
<key ctrl="-1 -1"/>
<key ctrl="-2 0"/>
</keyframe>
</mujoco>
@@ -0,0 +1,43 @@
<mujoco>
<worldbody>
<geom name="floor" size="10 10 .001" type="plane"/>
<body pos="0 0 .1">
<joint type="slide" axis="1 0 0"/>
<joint type="slide" axis="0 0 1"/>
<geom type="sphere" size=".1"/>
</body>
<body pos=".2 0 .1">
<joint type="slide" axis="1 0 0"/>
<joint type="slide" axis="0 0 1"/>
<geom type="sphere" size=".1"/>
</body>
<body pos="0 0 .6">
<joint type="slide" axis="1 0 0"/>
<joint type="slide" axis="0 0 1"/>
<geom type="capsule" size=".1 .2"/>
</body>
<body pos=".2 0 .6">
<joint type="slide" axis="1 0 0"/>
<joint type="slide" axis="0 0 1"/>
<geom type="capsule" size=".1 .2"/>
</body>
<body pos="0 0 1">
<joint type="slide" axis="1 0 0"/>
<joint type="slide" axis="0 0 1"/>
<geom type="box" size=".1 .1 .1"/>
</body>
<body pos=".2 0 1">
<joint type="slide" axis="1 0 0"/>
<joint type="slide" axis="0 0 1"/>
<geom type="box" size=".1 .1 .1"/>
</body>
</worldbody>
<keyframe>
<key qpos=" 0 -.1
-.2 -.1
0 -.6
-.2 -.6
0 -1
-.2 -1"/>
</keyframe>
</mujoco>
@@ -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
@@ -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
@@ -0,0 +1,56 @@
<mujoco>
<extension>
<plugin plugin="mujoco.sdf.nut">
<instance name="nut">
<config key="radius" value="0.26"/>
</instance>
</plugin>
<plugin plugin="mujoco.sdf.bolt">
<instance name="bolt">
<config key="radius" value="0.255"/>
</instance>
</plugin>
</extension>
<compiler autolimits="true"/>
<include file="scene.xml"/>
<visual>
<map force="0.05"/>
</visual>
<asset>
<mesh name="nut">
<plugin instance="nut"/>
</mesh>
<mesh name="bolt">
<plugin instance="bolt"/>
</mesh>
</asset>
<option sdf_iterations="10" sdf_initpoints="40"/>
<default>
<geom solref="0.01 1" solimp=".95 .99 .0001" friction="0.01"/>
</default>
<statistic meansize=".1"/>
<worldbody>
<body pos="-0.0012496 0.00329058 0.830362" quat="-0.000212626 0.999996 -0.00200453 0.00185878">
<joint type="free" damping="30"/>
<geom type="sdf" name="nut" mesh="nut" rgba="0.83 0.68 0.4 1">
<plugin instance="nut"/>
</geom>
</body>
<body euler="180 0 0">
<geom type="sdf" name="bolt" mesh="bolt" rgba="0.7 0.7 0.7 1">
<plugin instance="bolt"/>
</geom>
</body>
<light name="left" pos="-1 0 2" cutoff="80"/>
<light name="right" pos="1 0 2" cutoff="80"/>
</worldbody>
</mujoco>
@@ -0,0 +1,38 @@
<!-- Copyright 2021 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.
-->
<mujoco>
<compiler texturedir="asset"/>
<statistic meansize=".05"/>
<visual>
<rgba haze="0.15 0.25 0.35 1"/>
<map stiffness="700" shadowscale="0.5" fogstart="1" fogend="15" zfar="40" haze="1" shadowclip="3"/>
</visual>
<asset>
<texture type="skybox" builtin="gradient" rgb1="0.3 0.5 0.7" rgb2="0 0 0" width="512" height="512"/>
<texture name="texplane" type="2d" builtin="checker" rgb1=".2 .3 .4" rgb2=".1 0.15 0.2"
width="512" height="512" mark="cross" markrgb=".8 .8 .8"/>
<material name="matplane" reflectance="0.3" texture="texplane" texrepeat="1 1" texuniform="true"/>
</asset>
<worldbody>
<light directional="true" diffuse=".8 .8 .8" specular="0.2 0.2 0.2" pos="0 0 4" dir="0 0 -1"/>
<geom name="ground" type="plane" pos="0 0 -.02" size="10 10 .01" material="matplane" condim="1"/>
</worldbody>
</mujoco>
@@ -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 = """<mujoco>
<extension>
<plugin plugin="mujoco.sdf.nut"><instance name="n"/></plugin>
<plugin plugin="mujoco.sdf.bolt"><instance name="b"/></plugin>
</extension>
<asset>
<mesh name="nm"><plugin instance="n"/></mesh>
<mesh name="bm"><plugin instance="b"/></mesh>
</asset>
<worldbody>
<body><geom type="sdf" name="ng" mesh="nm"><plugin instance="n"/></geom></body>
<body><geom type="sdf" name="bg" mesh="bm"><plugin instance="b"/></geom></body>
</worldbody>
</mujoco>"""
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
@@ -0,0 +1,114 @@
<!-- For validating constraint dynamics:
* connect, weld, joint constraints
* collision and joint limits for ball and 1d joints
* solref, solimp
-->
<mujoco>
<option timestep="0.015" impratio="1.5" iterations="4" ls_iterations="4"/>
<default>
<default class="box">
<geom type="box" size=".2" fromto="0 0 0 0 -2 0" rgba=".4 .7 .6 .3" contype="0"/>
</default>
</default>
<worldbody>
<geom pos="0 0 -1" type="plane" size="20 20 .01" condim="1"/>
<site name="site0" pos="0 0 1"/>
<body pos="0 0 1">
<freejoint/>
<geom class="box"/>
<site name="site1"/>
</body>
<site name="site2" pos="0 -1 1"/>
<body pos="0 -1 1">
<freejoint/>
<geom class="box"/>
<site name="site3"/>
</body>
<body name="anchor1" pos="-3 0 0"/>
<body name="beam1" pos="-3 0 0">
<joint name="joint1" type="ball" limited="true" range="0 .9" solreflimit="0.03 0.9" solimplimit="0.89 0.9 0.01 2.1"/>
<geom class="box"/>
</body>
<body name="anchor2" pos="-1 0 0"/>
<body name="beam2" pos="-1 0 0">
<freejoint align="false"/>
<geom class="box"/>
</body>
<body name="beam3" pos="1 0 0">
<joint name="joint3" axis="1 0 0" type="hinge" range="-20 20" frictionloss=".1"/>
<geom class="box"/>
</body>
<body name="beam4" pos="3 0 0">
<joint name="joint4" axis="1 0 0" type="hinge" damping="1.0"/> <!-- tests no joint range -->
<geom class="box"/>
</body>
<body name="beam5" pos="-2 0 0" euler="45 45 0">
<joint name="joint5" type="ball" solreflimit="-100 -1"/> <!-- test negative solref -->
<geom class="box" contype="0" conaffinity="0"/>
</body>
<body name="box_condim1" pos="4 0 0">
<freejoint align="false"/>
<geom class="box" condim="3"/>
</body>
<body name="box_condim3" pos="5 0 0">
<freejoint align="false"/>
<geom class="box" condim="3"/>
</body>
<body name="box_condim4" pos="6 0 0">
<freejoint align="false"/>
<geom class="box" condim="3"/>
</body>
<body name="box_condim6" pos="7 0 0">
<freejoint align="false"/>
<geom class="box" condim="3"/>
</body>
</worldbody>
<tendon>
<fixed name="tendon_1" limited="true" range="-0.3 0.1" stiffness=".1" damping=".2" frictionloss=".1">
<joint joint="joint3" coef=".1"/>
<joint joint="joint4" coef="-.2"/>
</fixed>
<fixed name="tendon_2" limited="true" range="-0.3 2" solreflimit="0.03 0.9" solimplimit="0.89 0.9 0.01 2.1">
<joint joint="joint4" coef=".3"/>
<joint joint="joint5" coef="-.4"/>
</fixed>
</tendon>
<equality>
<connect name="c_site" site1="site0" site2="site1"/>
<connect name="connect" body1="anchor1" body2="beam1" anchor="1 0 -1" />
<weld name="w_site" site1="site2" site2="site3" torquescale="0.1"/>
<weld name="weld" body1="anchor2" body2="beam2" relpose="0 0 0 1 -.3 0 0" torquescale="0.002" anchor="0 -2 0"/>
<joint name="joint" joint1="joint3" joint2="joint4" polycoef="0.5 -1 0.1 0.15 0.2" />
<tendon name="tendon" tendon1="tendon_1" tendon2="tendon_2" polycoef="0.5 -1.0 0.1 0.15 0.2"/>
</equality>
<actuator>
<position ctrlrange="-20 20" gear="500" joint="joint1" kv="0.5" name="act1"/>
<motor gear="50000" joint="joint3" name="act2"/>
<motor gear="75000" joint="joint4" name="act3"/>
</actuator>
<keyframe>
<!-- keyframe 0: default position with some motion, zero contacts -->
<key qpos='-1 -1 1 1 0 0 0 -2 -2 1 1 0 0 0 0 0 0 0 -1 0 0 1 0 0 0 0 0 1 0 0 0 4 0 0 1 0 0 0 5 0 0 1 0 0 0 6 0 0 1 0 0 0 7 0 0 1 0 0 0' qvel='1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1'/>
<!-- keyframe 1: some contacts but constraints are quadratic -->
<key qpos='-1 -1 1 1 0 0 0 -2 -2 1 1 0 0 0 1 0.0011 -8.2e-07 -0.0011 -0.82 1.9 -0.58 0.97 -0.14 0.17 -0.069 0.22 0.29 0.93 0.2 0.15 0.28 4.3 0.23 -0.26 0.97 0.12 0.11 0.18 5.4 0.25 -0.26 0.97 0.11 0.16 0.17 6.4 0.25 -0.26 0.97 0.11 0.16 0.17 7.4 0.25 -0.26 0.97 0.11 0.16 0.17' qvel='1 1 1 1 1 1 1 1 1 1 1 1 -0.14 3.3e-05 0.14 -0.14 -0.93 -3.6 -1.8 1 -0.57 -0.16 0.21 1.1 1 2.2 0.65 1.3 -4.2 -2.6 -5.6 0.54 1.6 2.8 -3.7 -2.3 0.7 -1.8 1.6 2.8 -3.7 -2.3 0.62 -1.8 1.6 2.8 -3.7 -2.3 0.61 -1.8'/>
<!-- keyframe 2: some contacts and some constraints are in cone state (for elliptic) -->
<key qpos='-1 -1 1 1 0 0 0 -2 -2 1 1 0 0 0 1 0.0087 2.4e-07 -0.0086 -0.89 1.8 -0.77 0.98 -0.2 -0.0022 -0.026 0.19 0.33 0.86 0.32 0.064 0.38 4.4 0.36 -0.81 0.97 -0.0013 -0.0011 0.25 5.6 0.52 -0.75 0.98 -0.018 0.17 0.094 6.6 0.52 -0.75 0.98 -0.017 0.16 0.094 7.6 0.52 -0.76 0.98 -0.017 0.16 0.094' qvel='1 1 1 1 1 1 1 1 1 1 1 1 0.2 -1.8e-05 -0.2 -0.72 0.072 0.025 0.015 -4.9 0.35 -0.22 0.26 0.99 -4.6 1.7 1.1 0.52 -0.73 0.16 7.9 1 0.043 -0.042 1.3 1.6 -1.6 0.98 0.027 -0.024 1.3 1.6 -1.8 0.96 0.025 -0.022 1.3 1.6 -1.8 0.96'/>
</keyframe>
</mujoco>
@@ -0,0 +1,45 @@
<!-- Copyright 2021 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.
-->
<mujoco model="Poncho">
<include file="mannequin.xml"/>
<option timestep="0.001" integrator="Euler" solver="CG" tolerance="1e-6" jacobian="sparse">
<flag energy="enable"/>
</option>
<visual>
<map force="0.1" zfar="30"/>
<rgba haze="0.15 0.25 0.35 1"/>
<quality shadowsize="4096"/>
<global offwidth="800" offheight="800"/>
</visual>
<default>
<geom solref="0.003 1"/>
</default>
<worldbody>
<light directional="false" diffuse=".2 .2 .2" specular="0 0 0" pos="0 0 5" dir="0 0 -1"/>
<flexcomp name="towel" type="grid" count="15 15 1" spacing="0.1 0.1 0.1"
radius="0.03" dim="2" rgba="1 0.5 0.5 1" pos="0 0 2" mass=".1">
<edge equality="false"/>
<elasticity young="3e2" poisson="0" thickness="1e-1" damping="1e-3" elastic2d="both"/>
<contact vertcollide="true" conaffinity="0" contype="0"/>
</flexcomp>
</worldbody>
</mujoco>
@@ -0,0 +1,45 @@
<!-- Copyright 2021 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.
-->
<mujoco model="Floppy">
<include file="scene.xml"/>
<compiler autolimits="true"/>
<option solver="CG" tolerance="1e-6" timestep=".001"/>
<size memory="100M"/>
<visual>
<map stiffness="100"/>
</visual>
<worldbody>
<flexcomp type="grid" count="24 4 4" spacing=".1 .1 .1" pos=".1 0 1.5"
radius=".03" rgba="0 .7 .7 1" name="softbody" dim="3" mass="25">
<contact condim="3" solref="0.01 1" solimp=".95 .99 .0001" selfcollide="none" vertcollide="true" conaffinity="0" contype="0"/>
<elasticity young="5e4" damping="0.002" poisson="0.2"/>
</flexcomp>
<body>
<joint name="hinge" pos="0 0 .5" axis="0 1 0" damping="50"/>
<geom type="cylinder" size=".4" fromto="0 -.5 .5 0 .5 .5" density="300"/>
</body>
</worldbody>
<actuator>
<motor name="cylinder" joint="hinge" gear="1 0 0 0 0 0" ctrlrange="-100 100"/>
</actuator>
</mujoco>
@@ -0,0 +1,174 @@
<!-- Copyright 2021 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.
-->
<mujoco model="Mannequin">
<!-- Static humanoid model with all joints removed -->
<option timestep="0.005"/>
<visual>
<map force="0.1" zfar="30"/>
<rgba haze="0.15 0.25 0.35 1"/>
<global offwidth="2560" offheight="1440" elevation="-20" azimuth="120"/>
</visual>
<statistic center="0 0 0.7"/>
<asset>
<texture type="skybox" builtin="gradient" rgb1=".3 .5 .7" rgb2="0 0 0" width="32" height="512"/>
<texture name="body" type="cube" builtin="flat" mark="cross" width="128" height="128" rgb1="0.8 0.6 0.4" rgb2="0.8 0.6 0.4" markrgb="1 1 1" random="0.01"/>
<material name="body" texture="body" texuniform="true" rgba="0.8 0.6 .4 1"/>
<texture name="grid" type="2d" builtin="checker" width="512" height="512" rgb1=".1 .2 .3" rgb2=".2 .3 .4"/>
<material name="grid" texture="grid" texrepeat="1 1" texuniform="true" reflectance=".2"/>
</asset>
<default>
<motor ctrlrange="-1 1" ctrllimited="true"/>
<default class="body">
<!-- geoms -->
<geom type="capsule" condim="1" friction=".7" solimp=".9 .99 .003" solref=".003 1" material="body"/>
<default class="thigh">
<geom size=".06"/>
</default>
<default class="shin">
<geom fromto="0 0 0 0 0 -.3" size=".049"/>
</default>
<default class="foot">
<geom size=".027"/>
<default class="foot1">
<geom fromto="-.07 -.01 0 .14 -.03 0"/>
</default>
<default class="foot2">
<geom fromto="-.07 .01 0 .14 .03 0"/>
</default>
</default>
<default class="arm_upper">
<geom size=".04"/>
</default>
<default class="arm_lower">
<geom size=".031"/>
</default>
<default class="hand">
<geom type="sphere" size=".04"/>
</default>
<!-- joints -->
<joint type="hinge" damping=".2" stiffness="1" armature=".01" limited="true" solimplimit="0 .99 .01"/>
<default class="joint_big">
<joint damping="5" stiffness="10"/>
<default class="hip_x">
<joint range="-30 10"/>
</default>
<default class="hip_z">
<joint range="-60 35"/>
</default>
<default class="hip_y">
<joint axis="0 1 0" range="-150 20"/>
</default>
<default class="joint_big_stiff">
<joint stiffness="20"/>
</default>
</default>
<default class="knee">
<joint pos="0 0 .02" axis="0 -1 0" range="-160 2"/>
</default>
<default class="ankle">
<joint range="-50 50"/>
<default class="ankle_y">
<joint pos="0 0 .08" axis="0 1 0" stiffness="6"/>
</default>
<default class="ankle_x">
<joint pos="0 0 .04" stiffness="3"/>
</default>
</default>
<default class="shoulder">
<joint range="-85 60"/>
</default>
<default class="elbow">
<joint range="-100 50" stiffness="0"/>
</default>
</default>
</default>
<worldbody>
<geom name="floor" size="0 0 .05" type="plane" material="grid" condim="3"/>
<light name="spotlight" mode="targetbodycom" target="torso" diffuse=".8 .8 .8" specular="0.3 0.3 0.3" pos="0 -6 4" cutoff="30"/>
<body name="torso" pos="0 0 1.282" childclass="body">
<light name="top" pos="0 0 2" mode="trackcom"/>
<camera name="back" pos="-3 0 1" xyaxes="0 -1 0 1 0 2" mode="trackcom"/>
<camera name="side" pos="0 -3 1" xyaxes="1 0 0 0 1 2" mode="trackcom"/>
<freejoint name="root"/>
<geom name="torso" fromto="0 -.07 0 0 .07 0" size=".07"/>
<geom name="waist_upper" fromto="-.01 -.06 -.12 -.01 .06 -.12" size=".06"/>
<body name="head" pos="0 0 .19">
<geom name="head" type="sphere" size=".09"/>
<camera name="egocentric" pos=".09 0 0" xyaxes="0 -1 0 .1 0 1" fovy="80"/>
</body>
<body name="neck">
<geom type="cylinder" size=".03" fromto="0 0 .07 0 0 .1"/>
</body>
<body name="waist_lower" pos="-.01 0 -.26">
<geom name="waist_lower" fromto="0 -.06 0 0 .06 0" size=".06"/>
<body name="pelvis" pos="0 0 -.165">
<geom name="butt" fromto="-.02 -.07 0 -.02 .07 0" size=".09"/>
<body name="thigh_right" pos="0 -.1 -.04">
<geom name="thigh_right" fromto="0 0 0 0 .01 -.34" class="thigh"/>
<body name="shin_right" pos="0 .01 -.4">
<geom name="shin_right" class="shin"/>
<body name="foot_right" pos="0 0 -.39">
<geom name="foot1_right" class="foot1"/>
<geom name="foot2_right" class="foot2"/>
</body>
</body>
</body>
<body name="thigh_left" pos="0 .1 -.04">
<geom name="thigh_left" fromto="0 0 0 0 -.01 -.34" class="thigh"/>
<body name="shin_left" pos="0 -.01 -.4">
<geom name="shin_left" fromto="0 0 0 0 0 -.3" class="shin"/>
<body name="foot_left" pos="0 0 -.39">
<geom name="foot1_left" class="foot1"/>
<geom name="foot2_left" class="foot2"/>
</body>
</body>
</body>
</body>
</body>
<body name="upper_arm_right" pos="0 -.17 .06">
<geom name="upper_arm_right" fromto="0 0 0 0 -.22 0" class="arm_upper"/>
<body name="lower_arm_right" pos="0 -.22 0">
<geom name="lower_arm_right" fromto=".01 .01 .01 0 -.21 -.05" class="arm_lower"/>
<body name="hand_right" pos="0 -.22 -.05">
<geom name="hand_right" zaxis="1 1 1" class="hand"/>
</body>
</body>
</body>
<body name="upper_arm_left" pos="0 .17 .06">
<geom name="upper_arm_left" fromto="0 0 0 0 .22 0" class="arm_upper"/>
<body name="lower_arm_left" pos="0 .22 0">
<geom name="lower_arm_left" fromto=".01 -.01 .01 0 .21 -.05" class="arm_lower"/>
<body name="hand_left" pos="0 .22 -.05">
<geom name="hand_left" zaxis="1 -1 1" class="hand"/>
</body>
</body>
</body>
</body>
</worldbody>
<contact>
<exclude body1="waist_lower" body2="thigh_right"/>
<exclude body1="waist_lower" body2="thigh_left"/>
</contact>
</mujoco>
@@ -0,0 +1,41 @@
<!-- Copyright 2021 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.
-->
<mujoco>
<compiler meshdir="asset" texturedir="asset"/>
<statistic meansize=".05"/>
<visual>
<rgba haze="0.15 0.25 0.35 1"/>
<quality shadowsize="4096"/>
<map stiffness="700" shadowscale="0.5" fogstart="1" fogend="15" zfar="40" haze="1"/>
</visual>
<asset>
<texture type="skybox" builtin="gradient" rgb1="0.3 0.5 0.7" rgb2="0 0 0" width="512" height="512"/>
<texture name="texplane" type="2d" builtin="checker" rgb1=".2 .3 .4" rgb2=".1 0.15 0.2"
width="512" height="512" mark="cross" markrgb=".8 .8 .8"/>
<material name="matplane" reflectance="0.3" texture="texplane" texrepeat="10 10" texuniform="true"/>
</asset>
<worldbody>
<light diffuse=".4 .4 .4" specular="0.1 0.1 0.1" pos="0 0 2.0" dir="0 0 -1" castshadow="false"/>
<light directional="true" diffuse=".8 .8 .8" specular="0.2 0.2 0.2" pos="0 0 4" dir="0 0 -1"/>
<geom name="ground" type="plane" size="0 0 1" pos="0 0 0" quat="1 0 0 0" material="matplane" condim="1"/>
</worldbody>
</mujoco>
@@ -0,0 +1,84 @@
<mujoco>
<asset>
<!--
<hfield name="terrain" nrow="5" ncol="5" size="1 1 0.5 0.3"
elevation="5 3 2 3 5
3 2 1 2 3
2 1 0 1 2
3 2 1 2 3
5 3 2 3 5"/>
-->
<!-- -->
<hfield name="terrain" nrow="3" ncol="3" size="1 1 0.1 0.3"
elevation="0 0 0
0 1 0
0 0 0"/>
<!-- -->
<!--
<hfield name="terrain" nrow="3" ncol="5" size="2 1 0.5 0.3"
elevation="8 4 2 4 8
4 2 1 2 4
8 4 2 4 8"/>
-->
<!--
<hfield name="terrain" nrow="2" ncol="3" size="2 1 0.5 0.3"
elevation="8 2 8
2 8 2"/>
-->
<!-- Simple box mesh, size will be scaled by the geom size
<mesh name="box_mesh" vertex="1 1 0 1 1 -1 1 -1 0.2 1 -1 -1 -1 1 0 -1 1 -1 -1 -1 0 -1 -1 -1"/>
-->
</asset>
<worldbody>
<!-- Ground plane -->
<geom type="plane" size="2 2 0.1" pos="0 0 -0.3" rgba="0.8 0.8 0.8 1"/>
<!-- Height fields -->
<geom type="hfield" hfield="terrain" pos="0 0 0" rgba="0.4 0.6 0.4 1"/>
<!--geom type="hfield" hfield="terrain" pos="2 0 0" rgba="0.4 0.6 0.4 1"/-->
<!--geom type="box" size="1.0 1.0 0.1" rgba="0.8 0.2 0.2 1"/-->
<!-- Spheres
<body name="sphere1" pos="-0.9 -0.9 0.5">
<freejoint/>
<geom type="sphere" size="0.1" rgba="0.8 0.2 0.2 1"/>
</body>
<body name="sphere2" pos="-0.85 -0.85 0.8">
<freejoint/>
<geom type="sphere" size="0.1" rgba="0.8 0.2 0.2 1"/>
</body>
<body name="sphere3" pos="-0.8 -0.8 1.1">
<freejoint/>
<geom type="sphere" size="0.1" rgba="0.8 0.2 0.2 1"/>
</body>
-->
<body pos="0 0.1 0.3">
<freejoint/>
<geom type="sphere" size="0.1 0.1 0.1" rgba="0.8 0.2 0.2 1"/>
</body>
<!-- Boxes ->
<body name="box1" pos="0.8 0.8 0.6">
<freejoint/>
<geom type="box" size="0.1 0.1 0.1" rgba="0.8 0.2 0.2 1"/>
</body>
<body name="box2" pos="0.75 0.75 0.9">
<freejoint/>
<geom type="box" size="0.1 0.1 0.1" rgba="0.8 0.2 0.2 1"/>
</body>
<body name="box3" pos="0.7 0.7 1.2">
<freejoint/>
<geom type="box" size="0.1 0.1 0.1" rgba="0.8 0.2 0.2 1"/>
</body>
<body name="box4" pos="0.65 0.65 1.5">
<freejoint/>
<geom type="box" size="0.1 0.1 0.1" rgba="0.8 0.2 0.2 1"/>
</body>
-->
</worldbody>
</mujoco>
@@ -0,0 +1,252 @@
<!-- Copyright 2021 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.
-->
<mujoco model="Humanoid">
<option timestep="0.005" iterations="10" ls_iterations="20">
<flag eulerdamp="disable"/>
</option>
<visual>
<map force="0.1" zfar="30"/>
<rgba haze="0.15 0.25 0.35 1"/>
<global offwidth="2560" offheight="1440" elevation="-20" azimuth="120"/>
</visual>
<statistic center="0 0 0.7"/>
<asset>
<texture type="skybox" builtin="gradient" rgb1=".3 .5 .7" rgb2="0 0 0" width="32" height="512"/>
<texture name="body" type="cube" builtin="flat" mark="cross" width="128" height="128" rgb1="0.8 0.6 0.4" rgb2="0.8 0.6 0.4" markrgb="1 1 1" random="0.01"/>
<material name="body" texture="body" texuniform="true" rgba="0.8 0.6 .4 1"/>
<texture name="grid" type="2d" builtin="checker" width="512" height="512" rgb1=".1 .2 .3" rgb2=".2 .3 .4"/>
<material name="grid" texture="grid" texrepeat="1 1" texuniform="true" reflectance=".2"/>
</asset>
<default>
<motor ctrlrange="-1 1" ctrllimited="true"/>
<default class="body">
<!-- geoms -->
<geom type="capsule" condim="1" friction=".7" solimp=".9 .99 .003" solref=".015 1" material="body"/>
<default class="thigh">
<geom size=".06"/>
</default>
<default class="shin">
<geom fromto="0 0 0 0 0 -.3" size=".049"/>
</default>
<default class="foot">
<geom size=".027"/>
<default class="foot1">
<geom fromto="-.07 -.01 0 .14 -.03 0"/>
</default>
<default class="foot2">
<geom fromto="-.07 .01 0 .14 .03 0"/>
</default>
</default>
<default class="arm_upper">
<geom size=".04"/>
</default>
<default class="arm_lower">
<geom size=".031"/>
</default>
<default class="hand">
<geom type="sphere" size=".04"/>
</default>
<!-- joints -->
<joint type="hinge" damping=".2" stiffness="1" armature=".01" limited="true" solimplimit="0 .99 .01"/>
<default class="joint_big">
<joint damping="5" stiffness="10"/>
<default class="hip_x">
<joint range="-30 10"/>
</default>
<default class="hip_z">
<joint range="-60 35"/>
</default>
<default class="hip_y">
<joint axis="0 1 0" range="-150 20"/>
</default>
<default class="joint_big_stiff">
<joint stiffness="20"/>
</default>
</default>
<default class="knee">
<joint pos="0 0 .02" axis="0 -1 0" range="-160 2"/>
</default>
<default class="ankle">
<joint range="-50 50"/>
<default class="ankle_y">
<joint pos="0 0 .08" axis="0 1 0" stiffness="6"/>
</default>
<default class="ankle_x">
<joint pos="0 0 .04" stiffness="3"/>
</default>
</default>
<default class="shoulder">
<joint range="-85 60"/>
</default>
<default class="elbow">
<joint range="-100 50" stiffness="0"/>
</default>
</default>
</default>
<worldbody>
<geom name="floor" size="0 0 .05" type="plane" material="grid" condim="3"/>
<light name="spotlight" mode="targetbodycom" target="torso" diffuse=".8 .8 .8" specular="0.3 0.3 0.3" pos="0 -6 4" cutoff="30"/>
<body name="torso" pos="0 0 1.282" childclass="body">
<light name="top" pos="0 0 2" mode="trackcom"/>
<camera name="back" pos="-3 0 1" xyaxes="0 -1 0 1 0 2" mode="trackcom"/>
<camera name="side" pos="0 -3 1" xyaxes="1 0 0 0 1 2" mode="trackcom"/>
<freejoint name="root"/>
<geom name="torso" fromto="0 -.07 0 0 .07 0" size=".07"/>
<geom name="waist_upper" fromto="-.01 -.06 -.12 -.01 .06 -.12" size=".06"/>
<body name="head" pos="0 0 .19">
<geom name="head" type="sphere" size=".09"/>
<camera name="egocentric" pos=".09 0 0" xyaxes="0 -1 0 .1 0 1" fovy="80"/>
</body>
<body name="waist_lower" pos="-.01 0 -.26">
<geom name="waist_lower" fromto="0 -.06 0 0 .06 0" size=".06"/>
<joint name="abdomen_z" pos="0 0 .065" axis="0 0 1" range="-45 45" class="joint_big_stiff"/>
<joint name="abdomen_y" pos="0 0 .065" axis="0 1 0" range="-75 30" class="joint_big"/>
<body name="pelvis" pos="0 0 -.165">
<joint name="abdomen_x" pos="0 0 .1" axis="1 0 0" range="-35 35" class="joint_big"/>
<geom name="butt" fromto="-.02 -.07 0 -.02 .07 0" size=".09"/>
<body name="thigh_right" pos="0 -.1 -.04">
<joint name="hip_x_right" axis="1 0 0" class="hip_x"/>
<joint name="hip_z_right" axis="0 0 1" class="hip_z"/>
<joint name="hip_y_right" class="hip_y"/>
<geom name="thigh_right" fromto="0 0 0 0 .01 -.34" class="thigh"/>
<body name="shin_right" pos="0 .01 -.4">
<joint name="knee_right" class="knee"/>
<geom name="shin_right" class="shin"/>
<body name="foot_right" pos="0 0 -.39">
<joint name="ankle_y_right" class="ankle_y"/>
<joint name="ankle_x_right" class="ankle_x" axis="1 0 .5"/>
<geom name="foot1_right" class="foot1"/>
<geom name="foot2_right" class="foot2"/>
</body>
</body>
</body>
<body name="thigh_left" pos="0 .1 -.04">
<joint name="hip_x_left" axis="-1 0 0" class="hip_x"/>
<joint name="hip_z_left" axis="0 0 -1" class="hip_z"/>
<joint name="hip_y_left" class="hip_y"/>
<geom name="thigh_left" fromto="0 0 0 0 -.01 -.34" class="thigh"/>
<body name="shin_left" pos="0 -.01 -.4">
<joint name="knee_left" class="knee"/>
<geom name="shin_left" fromto="0 0 0 0 0 -.3" class="shin"/>
<body name="foot_left" pos="0 0 -.39">
<joint name="ankle_y_left" class="ankle_y"/>
<joint name="ankle_x_left" class="ankle_x" axis="-1 0 -.5"/>
<geom name="foot1_left" class="foot1"/>
<geom name="foot2_left" class="foot2"/>
</body>
</body>
</body>
</body>
</body>
<body name="upper_arm_right" pos="0 -.17 .06">
<joint name="shoulder1_right" axis="2 1 1" class="shoulder"/>
<joint name="shoulder2_right" axis="0 -1 1" class="shoulder"/>
<geom name="upper_arm_right" fromto="0 0 0 .16 -.16 -.16" class="arm_upper"/>
<body name="lower_arm_right" pos=".18 -.18 -.18">
<joint name="elbow_right" axis="0 -1 1" class="elbow"/>
<geom name="lower_arm_right" fromto=".01 .01 .01 .17 .17 .17" class="arm_lower"/>
<body name="hand_right" pos=".18 .18 .18">
<geom name="hand_right" zaxis="1 1 1" class="hand"/>
</body>
</body>
</body>
<body name="upper_arm_left" pos="0 .17 .06">
<joint name="shoulder1_left" axis="-2 1 -1" class="shoulder"/>
<joint name="shoulder2_left" axis="0 -1 -1" class="shoulder"/>
<geom name="upper_arm_left" fromto="0 0 0 .16 .16 -.16" class="arm_upper"/>
<body name="lower_arm_left" pos=".18 .18 -.18">
<joint name="elbow_left" axis="0 -1 -1" class="elbow"/>
<geom name="lower_arm_left" fromto=".01 -.01 .01 .17 -.17 .17" class="arm_lower"/>
<body name="hand_left" pos=".18 -.18 .18">
<geom name="hand_left" zaxis="1 -1 1" class="hand"/>
</body>
</body>
</body>
</body>
</worldbody>
<!-- <tendon>
<fixed name="hamstring_right" limited="true" range="-0.3 2">
<joint joint="hip_y_right" coef=".5"/>
<joint joint="knee_right" coef="-.5"/>
</fixed>
<fixed name="hamstring_left" limited="true" range="-0.3 2">
<joint joint="hip_y_left" coef=".5"/>
<joint joint="knee_left" coef="-.5"/>
</fixed>
</tendon> -->
<actuator>
<motor name="abdomen_y" gear="40" joint="abdomen_y"/>
<motor name="abdomen_z" gear="40" joint="abdomen_z"/>
<motor name="abdomen_x" gear="40" joint="abdomen_x"/>
<motor name="hip_x_right" gear="40" joint="hip_x_right"/>
<motor name="hip_z_right" gear="40" joint="hip_z_right"/>
<motor name="hip_y_right" gear="120" joint="hip_y_right"/>
<motor name="knee_right" gear="80" joint="knee_right"/>
<motor name="ankle_x_right" gear="20" joint="ankle_x_right"/>
<motor name="ankle_y_right" gear="20" joint="ankle_y_right"/>
<motor name="hip_x_left" gear="40" joint="hip_x_left"/>
<motor name="hip_z_left" gear="40" joint="hip_z_left"/>
<motor name="hip_y_left" gear="120" joint="hip_y_left"/>
<motor name="knee_left" gear="80" joint="knee_left"/>
<motor name="ankle_x_left" gear="20" joint="ankle_x_left"/>
<motor name="ankle_y_left" gear="20" joint="ankle_y_left"/>
<motor name="shoulder1_right" gear="20" joint="shoulder1_right"/>
<motor name="shoulder2_right" gear="20" joint="shoulder2_right"/>
<motor name="elbow_right" gear="40" joint="elbow_right"/>
<motor name="shoulder1_left" gear="20" joint="shoulder1_left"/>
<motor name="shoulder2_left" gear="20" joint="shoulder2_left"/>
<motor name="elbow_left" gear="40" joint="elbow_left"/>
</actuator>
<keyframe>
<!--
The values below are split into rows for readibility:
torso position
torso orientation
spinal
right leg
left leg
arms
-->
<key name="squat" qpos="0 0 0.596
0.988015 0 0.154359 0
0 0.4 0
-0.25 -0.5 -2.5 -2.65 -0.8 0.56
-0.25 -0.5 -2.5 -2.65 -0.8 0.56
0 0 0 0 0 0"/>
<key name="stand_on_left_leg" qpos="0 0 1.21948
0.971588 -0.179973 0.135318 -0.0729076
-0.0516 -0.202 0.23
-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"/>
<key name="no_efc" qpos="0 0 2
1 0 0 0
0 0 0
0 0 0 0 0 0
0 0 0 0 0 0
0 0 0 0 0 0"/>
</keyframe>
</mujoco>
@@ -0,0 +1,167 @@
<!-- For validating dynamics of joints:
* free, ball, slide, hinge joints
* stacked joints (e.g. hinge + slide, ball + slide, etc)
* n-link kinematic chains
* limits, armature, damping
-->
<mujoco model="pendula">
<compiler autolimits="true"/>
<option timestep="0.02">
<flag contact="disable" />
</option>
<default>
<geom type="box" pos=".1 .2 .3" size=".1 .2 .3"/>
<joint damping="0.25" stiffness="0.1"/>
</default>
<worldbody>
<site name="origin"/>
<camera pos="0 1.5 0.8" xyaxes="-1 0 0 0 -0.2 .8"/>
<light pos="0 1.5 0.8" dir="0.1 -0.9 0"/>
<camera mode="targetbodycom" pos="1 -1 2"/>
<light mode="targetbodycom" pos="1 -1 2"/>
<camera mode="targetbodycom" target="body" pos="0 -1 2"/>
<light mode="targetbodycom" target="body" pos="0 -1 2"/>
<camera mode="targetbody" target="body" pos="0 -1 -2"/>
<light mode="targetbody" target="body" pos="0 -1 -2"/>
<!-- a single free body -->
<body pos="0 0 0">
<camera mode="fixed" pos="0 0.1 0.1" euler="2.7 0 0"/>
<light mode="fixed" pos="0 0.1 0.1" dir="0 0.02 -0.98"/>
<camera pos="-3 0 1" xyaxes="0 -1 0 1 0 2" mode="trackcom"/>
<light pos="-3 0 1" dir="0 -1 0" mode="trackcom"/>
<camera pos="-3 0 1" xyaxes="0 -1 0 1 0 2" mode="track"/>
<light pos="-3 0 1" dir="0 -1 0" mode="track"/>
<freejoint/>
<geom/>
</body>
<!-- a single ball joint with a limit -->
<body name="body" pos="0.5 0 0">
<joint name="joint1" type="ball" range="0 35"/>
<site name="s1" pos="1.3 0 0"/>
<geom/>
</body>
<!-- a single slide joint with a limit -->
<body pos="1.0 0 0">
<joint name="joint2" type="slide" axis="0.1 0.2 0.3" range="-1 1"/>
<site name="s2"/>
<geom/>
</body>
<!-- a single hinge joint with a limit -->
<body pos="1.5 0 0">
<joint name="joint3" type="hinge" axis="0.1 0.2 0.3" range="-35 50"/>
<site name="s3" quat="1 0 1 0"/>
<geom/>
</body>
<!-- stacked joint: hinge + slide -->
<body pos="2.0 0 0">
<joint name="joint4" type="hinge" axis="0.1 0.2 0.3"/>
<joint name="joint5" type="slide" axis="0.4 0.5 0.6" range="0 1"/>
<geom/>
</body>
<!-- stacked joint: slide + ball -->
<body pos="2.5 0 0">
<joint name="joint6" type="slide" axis="0.4 0.5 0.6" range="0 1"/>
<joint type="ball"/>
<geom/>
</body>
<!-- triple pendulum of hinges -->
<body pos="3.0 0 0">
<joint name="joint7" axis="0.1 0.2 0.3" type="hinge"/>
<geom/>
<body pos="0 0 -0.8">
<joint name="joint8" axis="0.4 0.5 0.6" type="hinge" armature="0.02" range="-20 20"/>
<geom/>
<body pos="0 -0.7 0">
<joint name="joint9" axis="0.7 0.8 0.9" type="hinge" damping="0.75" range="-30 30"/>
<site name="s4" pos="0 1.5 0" quat="1 0 1 0"/>
<geom/>
</body>
</body>
</body>
<!-- cherry pendulum: two bodies attached to same parent body -->
<body pos="3.5 0 0">
<joint name="joint10" type="ball" damping="0.5" />
<geom/>
<body pos="0 0 -0.8">
<joint name="joint11" axis="0.4 0.5 0.6" type="hinge" armature="0.02" range="-20 20"/>
<geom/>
</body>
<body pos="0 -0.7 0">
<joint name="joint12" axis="0.7 0.8 0.9" type="hinge" damping="0.75" range="-30 30"/>
<geom/>
</body>
</body>
<!-- falling pendulum -->
<body pos="4.0 0 0">
<freejoint name="freejoint"/>
<geom/>
<body pos="0 0 -0.8">
<joint name="joint13" axis="0.4 0.5 0.6" type="slide" armature="0.02" range="-0.4 0.6"/>
<geom/>
<body pos="0 -0.7 0">
<joint name="joint14" axis="0.7 0.8 0.9" type="hinge" damping="0.75" range="-30 30"/>
<geom/>
</body>
</body>
</body>
<!-- triple pendulum of hinges with gravcomp -->
<body pos="3.0 0 0" gravcomp="1">
<joint name="joint15" axis="0.1 0.2 0.3" type="hinge"/>
<geom/>
<body pos="0 0 -0.8" gravcomp="2">
<joint name="joint16" axis="0.4 0.5 0.6" type="hinge" armature="0.02" range="-20 20"/>
<geom/>
<body pos="0 -0.7 0" gravcomp="3">
<joint name="joint17" axis="0.7 0.8 0.9" type="hinge" damping="0.75" range="-30 30" actuatorgravcomp="true"/>
<geom/>
</body>
</body>
</body>
<body mocap="true">
<geom type="sphere" size="0.05"/>
</body>
<body mocap="true" pos="1.0 2.0 0.5" quat="1 0 1 0">
<geom type="sphere" size="0.05"/>
</body>
</worldbody>
<tendon>
<fixed name="tendon_1" limited="true" range="-0.3 0.1" stiffness=".1" damping=".2">
<joint joint="joint2" coef=".1"/>
<joint joint="joint3" coef="-.2"/>
</fixed>
<fixed name="tendon_2" limited="true" range="-0.3 2" springlength="0 0.05" stiffness=".3" damping=".4">
<joint joint="joint4" coef=".3"/>
<joint joint="joint5" coef="-.4"/>
</fixed>
</tendon>
<actuator>
<motor gear="250 0 0" joint="joint1" name="act1"/>
<motor gear="0 275 0" joint="joint1" name="act2"/>
<motor gear="0 0 300" joint="joint1" name="act3"/>
<motor gear="275" joint="joint2" name="act4"/>
<motor gear="275" joint="joint3" name="act5"/>
<motor gear="150" joint="joint15" name="act6"/>
<motor gear="150" joint="joint16" name="act7"/>
<motor gear="150" joint="joint17" name="act8"/>
<position tendon="tendon_2" kp="100"/>
<motor gear="2 3 4" jointinparent="joint1"/>
<motor gear="5 6 7 8 9 10" jointinparent="freejoint"/>
</actuator>
</mujoco>
@@ -0,0 +1,21 @@
<mujoco model="ray">
<asset>
<mesh name="tetrahedron" file="meshes/tetrahedron.stl" scale="0.4 0.4 0.4" />
<mesh name="dodecahedron" file="meshes/dodecahedron.stl" scale="0.04 0.04 0.04" />
<hfield name="hfield" nrow="3" ncol="2" elevation="1 2 3 3 2 1" size=".6 .4 .1 .1"/>
<texture builtin="checker" height="100" name="texplane" rgb1="0 0 0" rgb2="0.8 0.8 0.8" type="2d" width="100"/>
<material name="MatPlane" reflectance="0.5" shininess="1" specular="1" texrepeat="60 60" texture="texplane"/>
</asset>
<worldbody>
<light cutoff="100" diffuse="1 1 1" dir="-0 0 -1.3" directional="true" exponent="1" pos="0 0 1.3" specular=".1 .1 .1"/>
<geom name="plane" pos="0 0 0" quat="1 0 0 0" size="4 4 4" type="plane" rgba="0.1 0.1 0.1 1"/>
<geom name="sphere" pos="0 0 1" quat="1 0 0 0" size="0.5" type="sphere" rgba="1 0 0 1"/>
<geom name="capsule" pos="0 1 1" quat="0 0.3826834 0 0.9238795 " size="0.25 0.5" type="capsule" rgba="0 1 0 1"/>
<geom name="box" pos="1 0 1" quat="0 0.3826834 0 0.9238795" size="0.5 0.25 0.3" type="box" rgba="0 0 1 1"/>
<geom name="mesh" pos="1 1 1" quat="0 0 0.3826834 0.9238795" type="mesh" mesh="tetrahedron" rgba="1 1 0 1"/>
<geom name="mesh2" pos="2 1 1" type="mesh" mesh="dodecahedron" rgba="1 0 1 1"/>
<geom name="hfield" pos="0 2 1" type="hfield" hfield="hfield" rgba=".2 .4 .6 1"/>
<geom name="cylinder" pos="2 0 1" quat=" 0 0 .3826834 .9238796" type="cylinder" size=".25 .5" rgba="1 1 1 1"/>
</worldbody>
</mujoco>
@@ -0,0 +1,19 @@
<mujoco>
<worldbody>
<body>
<joint type="hinge" axis="0 1 0"/>
<geom type="sphere" size=".2" pos="1 0 0"/>
<site name="site0" pos="1 0 0"/>
</body>
<site name="site1" pos="1 0 0"/>
</worldbody>
<tendon>
<spatial armature="2">
<site site="site0"/>
<site site="site1"/>
</spatial>
</tendon>
<keyframe>
<key qpos="1" qvel=".5"/>
</keyframe>
</mujoco>
@@ -0,0 +1,20 @@
<mujoco>
<option timestep=".1"/>
<worldbody>
<body>
<joint type="hinge" axis="0 1 0"/>
<geom type="sphere" size=".2" pos="1 0 0"/>
<site name="site0" pos="1 0 0"/>
</body>
<site name="site1" pos="1 0 0"/>
</worldbody>
<tendon>
<spatial damping="5">
<site site="site0"/>
<site site="site1"/>
</spatial>
</tendon>
<keyframe>
<key qpos="1" qvel=".5"/>
</keyframe>
</mujoco>
@@ -0,0 +1,29 @@
<mujoco>
<worldbody>
<body>
<joint name="joint0" type="hinge"/>
<geom type="sphere" size="0.1"/>
<body>
<joint name="joint1" type="hinge"/>
<geom type="sphere" size="0.1"/>
<body>
<joint name="joint2" type="hinge"/>
<geom type="sphere" size="0.1"/>
</body>
</body>
</body>
</worldbody>
<tendon>
<fixed name="fixed">
<joint joint="joint0" coef=".25"/>
<joint joint="joint1" coef=".5"/>
<joint joint="joint2" coef=".75"/>
</fixed>
</tendon>
<actuator>
<motor tendon="fixed"/>
</actuator>
<keyframe>
<key qpos=".2 .4 .6" qvel=".1 .2 .3"/>
</keyframe>
</mujoco>
@@ -0,0 +1,35 @@
<mujoco>
<worldbody>
<site name="siteworld"/>
<body>
<joint name="jnt0" type="slide" axis="1 0 0"/>
<geom type="sphere" size="0.1"/>
<site name="site0" pos="0 0 .1"/>
<body>
<joint name="jnt1" type="slide" axis="1 0 0"/>
<joint name="jnt2" type="hinge" axis="0 1 0"/>
<geom type="sphere" size="0.1"/>
<site name="site1" pos="0 0 .2"/>
</body>
</body>
</worldbody>
<tendon>
<fixed name="fixed">
<joint joint="jnt0" coef=".1"/>
<joint joint="jnt1" coef=".2"/>
<joint joint="jnt2" coef=".3"/>
</fixed>
<spatial name="spatial">
<site site="siteworld"/>
<site site="site0"/>
<site site="site1"/>
</spatial>
</tendon>
<actuator>
<motor tendon="fixed"/>
<motor tendon="spatial"/>
</actuator>
<keyframe>
<key qpos=".25 .5 .75"/>
</keyframe>
</mujoco>
@@ -0,0 +1,36 @@
<mujoco>
<worldbody>
<site name="siteworld"/>
<body>
<joint name="jnt0" type="slide" axis="1 0 0"/>
<geom type="sphere" size="0.1"/>
<site name="site0" pos="0 0 .1"/>
<body>
<joint name="jnt1" type="slide" axis="1 0 0"/>
<joint name="jnt2" type="hinge" axis="0 1 0"/>
<geom type="sphere" size="0.1"/>
<site name="site1" pos="0 0 .2"/>
</body>
</body>
</worldbody>
<tendon>
<fixed name="fixed">
<joint joint="jnt0" coef=".1"/>
<joint joint="jnt1" coef=".2"/>
<joint joint="jnt2" coef=".3"/>
</fixed>
<spatial name="spatial">
<pulley divisor="2"/>
<site site="siteworld"/>
<site site="site0"/>
<site site="site1"/>
</spatial>
</tendon>
<actuator>
<motor tendon="fixed"/>
<motor tendon="spatial"/>
</actuator>
<keyframe>
<key qpos=".25 .5 .75"/>
</keyframe>
</mujoco>
@@ -0,0 +1,46 @@
<mujoco>
<worldbody>
<site name="siteworld"/>
<body>
<joint type="slide" axis="1 0 0"/>
<geom type="sphere" size="0.1"/>
<site name="site0" pos="0 0 .1"/>
<body>
<joint type="slide" axis="1 0 0"/>
<joint type="hinge" axis="0 1 0"/>
<geom type="sphere" size="0.1"/>
<site name="site1" pos="0 0 .2"/>
</body>
</body>
</worldbody>
<tendon>
<spatial name="spatial0">
<pulley divisor="2"/>
<site site="siteworld"/>
<site site="site0"/>
<site site="site1"/>
</spatial>
<spatial name="spatial1">
<site site="siteworld"/>
<site site="site0"/>
</spatial>
<spatial name="spatial2">
<site site="siteworld"/>
<site site="site1"/>
</spatial>
<spatial name="spatial3">
<pulley divisor="3"/>
<site site="site0"/>
<site site="site1"/>
</spatial>
</tendon>
<actuator>
<motor tendon="spatial0"/>
<motor tendon="spatial1"/>
<motor tendon="spatial2"/>
<motor tendon="spatial3"/>
</actuator>
<keyframe>
<key qpos=".25 .5 .75"/>
</keyframe>
</mujoco>
@@ -0,0 +1,32 @@
<mujoco>
<worldbody>
<site name="siteworld"/>
<body>
<joint name="jnt0" type="slide" axis="1 0 0"/>
<geom type="sphere" size="0.1"/>
<site name="site0" pos="0 0 .1"/>
<body>
<joint name="jnt1" type="slide" axis="1 0 0"/>
<joint name="jnt2" type="hinge" axis="0 1 0"/>
<geom type="sphere" size="0.1"/>
<site name="site1" pos="0 0 .2"/>
</body>
</body>
</worldbody>
<tendon>
<spatial>
<pulley divisor="2"/>
<site site="siteworld"/>
<site site="site0"/>
<site site="site1"/>
</spatial>
<fixed>
<joint joint="jnt0" coef=".1"/>
<joint joint="jnt1" coef=".2"/>
<joint joint="jnt2" coef=".3"/>
</fixed>
</tendon>
<keyframe>
<key qpos=".25 .5 .75"/>
</keyframe>
</mujoco>
@@ -0,0 +1,42 @@
<mujoco>
<worldbody>
<site name="siteworld0"/>
<site name="siteworld1" pos="3 0 0"/>
<geom name="sphereworld0" type="sphere" size=".1" pos=".5 0 .025"/>
<geom name="sphereworld1" type="cylinder" size=".2 .3" pos="1.25 0 .1"/>
<body>
<joint type="slide" axis="1 0 0"/>
<geom type="sphere" size="0.1"/>
<site name="site0"/>
</body>
<body>
<joint type="slide" axis="1 0 0"/>
<geom type="cylinder" size=".2 .3"/>
<site name="site1"/>
</body>
</worldbody>
<tendon>
<spatial name="spatial0">
<pulley divisor="2"/>
<site site="siteworld0"/>
<geom geom="sphereworld0"/>
<site site="site0"/>
</spatial>
<spatial name="spatial1">
<site site="siteworld0"/>
<site site="site0"/>
<geom geom="sphereworld1"/>
<site site="site1"/>
</spatial>
<spatial>
<pulley divisor="3"/>
<site site="siteworld0"/>
<geom geom="sphereworld1"/>
<site site="site1"/>
<site site="siteworld1"/>
</spatial>
</tendon>
<keyframe>
<key qpos="1 2"/>
</keyframe>
</mujoco>
@@ -0,0 +1,44 @@
<mujoco>
<worldbody>
<site name="siteworld"/>
<body>
<joint type="slide" axis="1 0 0"/>
<geom type="sphere" size="0.1"/>
<site name="site0" pos="0 0 .1"/>
<body>
<joint type="slide" axis="1 0 0"/>
<joint type="hinge" axis="0 1 0"/>
<geom type="sphere" size="0.1"/>
<site name="site1" pos="0 0 .2"/>
</body>
</body>
</worldbody>
<tendon>
<spatial name="spatial0">
<site site="siteworld"/>
<site site="site0"/>
<site site="site1"/>
</spatial>
<spatial name="spatial1">
<site site="siteworld"/>
<site site="site0"/>
</spatial>
<spatial name="spatial2">
<site site="siteworld"/>
<site site="site1"/>
</spatial>
<spatial name="spatial3">
<site site="site0"/>
<site site="site1"/>
</spatial>
</tendon>
<actuator>
<motor tendon="spatial0"/>
<motor tendon="spatial1"/>
<motor tendon="spatial2"/>
<motor tendon="spatial3"/>
</actuator>
<keyframe>
<key qpos=".25 .5 .75"/>
</keyframe>
</mujoco>
@@ -0,0 +1,31 @@
<mujoco>
<worldbody>
<site name="siteworld"/>
<body>
<joint name="jnt0" type="slide" axis="1 0 0"/>
<geom type="sphere" size="0.1"/>
<site name="site0" pos="0 0 .1"/>
<body>
<joint name="jnt1" type="slide" axis="1 0 0"/>
<joint name="jnt2" type="hinge" axis="0 1 0"/>
<geom type="sphere" size="0.1"/>
<site name="site1" pos="0 0 .2"/>
</body>
</body>
</worldbody>
<tendon>
<spatial>
<site site="siteworld"/>
<site site="site0"/>
<site site="site1"/>
</spatial>
<fixed>
<joint joint="jnt0" coef=".1"/>
<joint joint="jnt1" coef=".2"/>
<joint joint="jnt2" coef=".3"/>
</fixed>
</tendon>
<keyframe>
<key qpos=".25 .5 .75"/>
</keyframe>
</mujoco>
@@ -0,0 +1,29 @@
<mujoco>
<option>
<flag contact="disable"/>
</option>
<worldbody>
<body>
<joint name="joint0" type="hinge"/>
<geom type="sphere" size=".1"/>
<body>
<joint name="joint1" type="hinge"/>
<geom type="sphere" size=".1"/>
<body>
<joint name="joint2" type="hinge"/>
<geom type="sphere" size=".1"/>
</body>
</body>
</body>
</worldbody>
<tendon>
<fixed limited="true" range="0 .5">
<joint joint="joint0" coef=".25"/>
<joint joint="joint1" coef=".5"/>
<joint joint="joint2" coef=".75"/>
</fixed>
</tendon>
<keyframe>
<key qpos=".2 .4 .6"/>
</keyframe>
</mujoco>
@@ -0,0 +1,40 @@
<mujoco>
<worldbody>
<site name="siteworld0"/>
<site name="siteworld1" pos="3 0 0"/>
<geom name="sphereworld0" type="sphere" size=".1" pos=".5 0 .025"/>
<geom name="sphereworld1" type="cylinder" size=".2 .3" pos="1.25 0 .1"/>
<body>
<joint type="slide" axis="1 0 0"/>
<geom type="sphere" size="0.1"/>
<site name="site0"/>
</body>
<body>
<joint type="slide" axis="1 0 0"/>
<geom type="cylinder" size=".2 .3"/>
<site name="site1"/>
</body>
</worldbody>
<tendon>
<spatial name="spatial0">
<site site="siteworld0"/>
<geom geom="sphereworld0"/>
<site site="site0"/>
</spatial>
<spatial name="spatial1">
<site site="siteworld0"/>
<site site="site0"/>
<geom geom="sphereworld1"/>
<site site="site1"/>
</spatial>
<spatial>
<site site="siteworld0"/>
<geom geom="sphereworld1"/>
<site site="site1"/>
<site site="siteworld1"/>
</spatial>
</tendon>
<keyframe>
<key qpos="1 2"/>
</keyframe>
</mujoco>
@@ -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
@@ -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
# <kernel-id>|<launch-dims>|<arg-dims-list>
# 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",
)
+804
View File
@@ -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)
@@ -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,
}
+34 -4
View File
@@ -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
+70
View File
@@ -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()

Some files were not shown because too many files have changed in this diff Show More