Add mujoco_warp as an implementation in mjx.
PiperOrigin-RevId: 789366860 Change-Id: I84b49e744552092434df49d855f1052b04291b87
This commit is contained in:
committed by
Copybara-Service
parent
700da7c5bd
commit
47bc16a37c
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.'
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
+1037
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()
|
||||
@@ -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()
|
||||
+1598
File diff suppressed because it is too large
Load Diff
@@ -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()
|
||||
@@ -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
@@ -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
|
||||
@@ -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()
|
||||
@@ -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
@@ -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,
|
||||
)
|
||||
@@ -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()
|
||||
+2113
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()
|
||||
+3128
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()
|
||||
+2578
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()
|
||||
@@ -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()
|
||||
@@ -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
|
||||
+1667
File diff suppressed because it is too large
Load Diff
@@ -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()
|
||||
@@ -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>
|
||||
+26
@@ -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>
|
||||
Binary file not shown.
Binary file not shown.
@@ -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>
|
||||
+36
@@ -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>
|
||||
+32
@@ -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",
|
||||
)
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user