From 47bc16a37c6213a849c5f2a6134b7a58f5e8c529 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Thu, 31 Jul 2025 09:32:23 -0700 Subject: [PATCH] Add mujoco_warp as an implementation in mjx. PiperOrigin-RevId: 789366860 Change-Id: I84b49e744552092434df49d855f1052b04291b87 --- .github/workflows/build.yml | 3 +- doc/changelog.rst | 10 + mjx/cuda_requirements.txt | 25 +- mjx/mujoco/mjx/_src/collision_driver.py | 13 +- mjx/mujoco/mjx/_src/constraint.py | 13 +- mjx/mujoco/mjx/_src/dataclasses.py | 30 +- mjx/mujoco/mjx/_src/dataclasses_test.py | 69 + mjx/mujoco/mjx/_src/derivative.py | 12 +- mjx/mujoco/mjx/_src/forward.py | 10 + mjx/mujoco/mjx/_src/io.py | 341 +- mjx/mujoco/mjx/_src/io_test.py | 245 +- mjx/mujoco/mjx/_src/passive.py | 13 +- mjx/mujoco/mjx/_src/solver.py | 13 +- mjx/mujoco/mjx/_src/types.py | 96 +- .../mjx/test_data/humanoid/humanoid.xml | 6 + .../mjx/third_party/mujoco_warp/__init__.py | 79 + .../third_party/mujoco_warp/_src/__init__.py | 14 + .../mujoco_warp/_src/block_cholesky.py | 218 + .../mujoco_warp/_src/broadphase_test.py | 335 ++ .../mujoco_warp/_src/collision_convex.py | 495 ++ .../mujoco_warp/_src/collision_driver.py | 767 +++ .../mujoco_warp/_src/collision_driver_test.py | 870 ++++ .../mujoco_warp/_src/collision_gjk.py | 1361 ++++++ .../mujoco_warp/_src/collision_gjk_legacy.py | 714 +++ .../mujoco_warp/_src/collision_gjk_test.py | 320 ++ .../mujoco_warp/_src/collision_hfield.py | 397 ++ .../mujoco_warp/_src/collision_primitive.py | 2949 ++++++++++++ .../mujoco_warp/_src/collision_sdf.py | 599 +++ .../mujoco_warp/_src/constraint.py | 1908 ++++++++ .../mujoco_warp/_src/constraint_test.py | 217 + .../mujoco_warp/_src/derivative.py | 206 + .../third_party/mujoco_warp/_src/forward.py | 1037 +++++ .../mujoco_warp/_src/forward_test.py | 285 ++ .../third_party/mujoco_warp/_src/inverse.py | 147 + .../mujoco_warp/_src/inverse_test.py | 169 + .../mjx/third_party/mujoco_warp/_src/io.py | 1598 +++++++ .../third_party/mujoco_warp/_src/io_test.py | 354 ++ .../third_party/mujoco_warp/_src/jax_test.py | 101 + .../mjx/third_party/mujoco_warp/_src/math.py | 289 ++ .../third_party/mujoco_warp/_src/math_test.py | 131 + .../third_party/mujoco_warp/_src/passive.py | 617 +++ .../mujoco_warp/_src/passive_test.py | 144 + .../mjx/third_party/mujoco_warp/_src/ray.py | 903 ++++ .../third_party/mujoco_warp/_src/ray_test.py | 312 ++ .../third_party/mujoco_warp/_src/sensor.py | 2113 +++++++++ .../mujoco_warp/_src/sensor_test.py | 411 ++ .../third_party/mujoco_warp/_src/smooth.py | 3128 +++++++++++++ .../mujoco_warp/_src/smooth_test.py | 398 ++ .../third_party/mujoco_warp/_src/solver.py | 2578 +++++++++++ .../mujoco_warp/_src/solver_test.py | 344 ++ .../third_party/mujoco_warp/_src/support.py | 554 +++ .../mujoco_warp/_src/support_test.py | 123 + .../third_party/mujoco_warp/_src/test_util.py | 267 ++ .../mjx/third_party/mujoco_warp/_src/types.py | 1667 +++++++ .../third_party/mujoco_warp/_src/util_misc.py | 603 +++ .../mujoco_warp/_src/util_misc_test.py | 573 +++ .../third_party/mujoco_warp/_src/warp_util.py | 196 + .../test_data/actuation/actuation.xml | 40 + .../test_data/actuation/actuators.xml | 28 + .../test_data/actuation/adhesion.xml | 21 + .../test_data/actuation/muscle.xml | 22 + .../test_data/actuation/position.xml | 40 + .../mujoco_warp/test_data/actuation/site.xml | 31 + .../test_data/actuation/slidercrank.xml | 15 + .../actuation/tendon_force_limit.xml | 26 + .../mujoco_warp/test_data/collision.xml | 43 + .../test_data/collision_sdf/bolt.py | 70 + .../test_data/collision_sdf/nut.py | 65 + .../test_data/collision_sdf/nutbolt.xml | 56 + .../test_data/collision_sdf/scene.xml | 38 + .../test_data/collision_sdf/utils.py | 86 + .../mujoco_warp/test_data/constraints.xml | 114 + .../mujoco_warp/test_data/flex/cloth.xml | 45 + .../mujoco_warp/test_data/flex/floppy.xml | 45 + .../mujoco_warp/test_data/flex/mannequin.xml | 174 + .../mujoco_warp/test_data/flex/scene.xml | 41 + .../mujoco_warp/test_data/hfield/hfield.xml | 84 + .../test_data/humanoid/humanoid.xml | 252 + .../test_data/meshes/dodecahedron.stl | Bin 0 -> 1884 bytes .../test_data/meshes/tetrahedron.stl | Bin 0 -> 284 bytes .../mujoco_warp/test_data/pendula.xml | 167 + .../third_party/mujoco_warp/test_data/ray.xml | 21 + .../mujoco_warp/test_data/tendon/armature.xml | 19 + .../mujoco_warp/test_data/tendon/damping.xml | 20 + .../mujoco_warp/test_data/tendon/fixed.xml | 29 + .../test_data/tendon/fixed_site.xml | 35 + .../test_data/tendon/pulley_fixed_site.xml | 36 + .../test_data/tendon/pulley_site.xml | 46 + .../test_data/tendon/pulley_site_fixed.xml | 32 + .../test_data/tendon/pulley_wrap.xml | 42 + .../mujoco_warp/test_data/tendon/site.xml | 44 + .../test_data/tendon/site_fixed.xml | 31 + .../test_data/tendon/tendon_limit.xml | 29 + .../mujoco_warp/test_data/tendon/wrap.xml | 40 + .../warp/jax_experimental/__init__.py | 16 + .../warp/jax_experimental/custom_call.py | 363 ++ .../third_party/warp/jax_experimental/ffi.py | 804 ++++ .../warp/jax_experimental/xla_ffi.py | 615 +++ mjx/mujoco/mjx/viewer.py | 38 +- mjx/mujoco/mjx/warp/__init__.py | 70 + mjx/mujoco/mjx/warp/collision_driver.py | 498 ++ mjx/mujoco/mjx/warp/collision_driver_test.py | 107 + mjx/mujoco/mjx/warp/ffi.py | 350 ++ mjx/mujoco/mjx/warp/forward.py | 4123 +++++++++++++++++ mjx/mujoco/mjx/warp/forward_test.py | 274 ++ mjx/mujoco/mjx/warp/smooth.py | 291 ++ mjx/mujoco/mjx/warp/smooth_test.py | 229 + mjx/mujoco/mjx/warp/test_util.py | 204 + mjx/mujoco/mjx/warp/testspeed.py | 427 ++ mjx/mujoco/mjx/warp/types.py | 1602 +++++++ mjx/pyproject.toml | 7 + mjx/requirements.txt | 72 +- 112 files changed, 43209 insertions(+), 198 deletions(-) create mode 100644 mjx/mujoco/mjx/_src/dataclasses_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/__init__.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/broadphase_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/jax_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/math_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/support_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/test_util.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc_test.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/warp_util.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuation.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuators.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/adhesion.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/muscle.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/position.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/site.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/slidercrank.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/tendon_force_limit.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/bolt.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nut.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nutbolt.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/scene.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/utils.py create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/constraints.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/cloth.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/floppy.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/mannequin.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/scene.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/hfield/hfield.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/humanoid/humanoid.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/meshes/dodecahedron.stl create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/meshes/tetrahedron.stl create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/pendula.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/ray.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/armature.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/damping.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed_site.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_fixed_site.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site_fixed.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_wrap.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site_fixed.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/tendon_limit.xml create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/wrap.xml create mode 100644 mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py create mode 100644 mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py create mode 100644 mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py create mode 100644 mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py create mode 100644 mjx/mujoco/mjx/warp/__init__.py create mode 100644 mjx/mujoco/mjx/warp/collision_driver.py create mode 100644 mjx/mujoco/mjx/warp/collision_driver_test.py create mode 100644 mjx/mujoco/mjx/warp/ffi.py create mode 100644 mjx/mujoco/mjx/warp/forward.py create mode 100644 mjx/mujoco/mjx/warp/forward_test.py create mode 100644 mjx/mujoco/mjx/warp/smooth.py create mode 100644 mjx/mujoco/mjx/warp/smooth_test.py create mode 100644 mjx/mujoco/mjx/warp/test_util.py create mode 100644 mjx/mujoco/mjx/warp/testspeed.py create mode 100644 mjx/mujoco/mjx/warp/types.py diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index b2daab71..ac1a88e5 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -281,11 +281,12 @@ jobs: working-directory: mjx run: source ${{ matrix.tmpdir }}/venv/bin/activate && - pip install --require-hashes -r requirements.txt && + pip install --require-hashes -r requirements.txt pip install --no-index dist/mujoco_mjx-*.whl - name: Test MJX if: ${{ runner.os != 'Windows' }} shell: bash + working-directory: mjx run: source ${{ matrix.tmpdir }}/venv/bin/activate && pytest -n auto -v -k 'not IntegrationTest' --pyargs mujoco.mjx diff --git a/doc/changelog.rst b/doc/changelog.rst index 2f492926..4c82b923 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -47,6 +47,9 @@ Version 3.3.4 (July 8, 2025) 3. In the mjSpec C API, directly setting an element's name using :ref:`mjs_setString` has been replaced with a new function :ref:`mjs_setName` which allows checking for naming collisions at set-time rather than compile-time, for earlier catching of errors. Relatedly, the ``name`` attribute has been removed from all mjs elements. + 4. For MJX, the ``mjx.Option`` dataclass now has private and public fields similar to ``mjx.Model`` and + ``mjx.Data``. Some fields are no longer publicly available due to differences in the + underlying implementations of this data structure. General ^^^^^^^ @@ -66,6 +69,13 @@ Documentation 8. Added missing item documentation and clarified the nature of breaking changes in the 3.3.3 changelog. See items 3 and 4 below. +MJX +^^^ +- Add Warp as a backend implementation for MJX. The implementation can be specified via + ``mjx.put_model(m, impl='warp')`` and ``mjx.make_data(m, impl='warp')``. The warp implementation requires + a CUDA device, Python 3.12, and `warp-lang` to be installed. This feature is available in "beta" and + some bugs are expected. + Version 3.3.3 (June 10, 2025) ----------------------------- diff --git a/mjx/cuda_requirements.txt b/mjx/cuda_requirements.txt index 3791dff1..4ea1c845 100644 --- a/mjx/cuda_requirements.txt +++ b/mjx/cuda_requirements.txt @@ -1,15 +1,18 @@ -jax[cuda12_local]==0.4.34; python_version >= '3.10' \ - --hash=sha256:b957ca1fc91f7343f91a186af9f19c7f342c946f95a8c11c7f1e5cdfe2e58d9e -jax[cuda12_local]==0.4.30; python_version == '3.9' \ - --hash=sha256:289b30ae03b52f7f4baf6ef082a9f4e3e29c1080e22d13512c5ecf02d5f1a55b -jax-cuda12-plugin==0.4.34; python_version >= '3.10' \ - --hash=sha256:d035ea72bd9b8a65a6ea621bca1affdd33127fa3a52e7bded7692670d360adab \ - --hash=sha256:db988b7ba5063483a936ddbf162f04d1b4412e0d64340f11788c7bbc877e8a43 \ - --hash=sha256:e23721d1654b311b47cd6b35768520284bd036f8c7e6b11600143258b4a0409a \ - --hash=sha256:b2099a4407225122ff76f6dcdc8dbdae47e6f29343bdfd21460ad337dc34a209 +jax-cuda12-plugin==0.5.3; python_version >= '3.10' \ + --hash=sha256:6171aed2f4b3bdd5fc13782de1072c6a634fce13731b75d0cb0a6ab8f4e6e650 \ + --hash=sha256:ba2555967f9b6c381c8b4ef9fb03d05bc55ec25ecfee5cfe45c5ace34f7d4152 \ + --hash=sha256:298d2d768f1029b74a0b1d01270e549349d2c37dc07658796542cda967eb7bd3 \ + --hash=sha256:aaa704a5ef547595d022db1c1e4878a0677116412a9360c115d67ff4b64e1596 \ + --hash=sha256:c2517a7c2186f8708894696e26cf96ebd60b7879ceca398b2c46abb28d2c96c8 \ + --hash=sha256:2030cf1208ce4ea70ee56cac61ddd239f9798695fc39bb7739c50a25d6e9da44 \ + --hash=sha256:21fec1b56c98783ea0569b747a56751f1f9ff2187b48acc11c700d3bfc5e1a31 \ + --hash=sha256:1862595b2b6d815679d11e0e889e523185ee54a46d46e022689f70fc4554dd91 \ + --hash=sha256:6d43677f22f3be9544a205216cd6dac591335b1d9bbbed018cd17dbb1f3f4def \ + --hash=sha256:5bb9ea0e68d72d44e57e4cb6a58a1a729fe3fe32e964f71e398d8a25c2103b19 jax-cuda12-plugin==0.4.30; python_version == '3.9' \ --hash=sha256:d8d196241b9253ecb1144a4409b5deacbb9771624f097b2bbf025da3c7d8f4f8 -jax-cuda12-pjrt==0.4.34; python_version >= '3.10' \ - --hash=sha256:0c7cc98f962cc7fc8e0a5ea6331b42a0cee516f202f1c3019f6aa5cd9530cca0 +jax-cuda12-pjrt==0.5.3; python_version >= '3.10' \ + --hash=sha256:04ee111eaf5fc2692978ad4a5c84d5925e42eb05c1701849ba3a53f6515400cc \ + --hash=sha256:c5378306568ba0c81b230a779dd3194c9dd10339ab6360ae80928108d37e7f75 jax-cuda12-pjrt==0.4.30; python_version == '3.9' \ --hash=sha256:895d0198ad99638fcaf976c47592e2a543eef79ea15fabd24a402d055390c328 diff --git a/mjx/mujoco/mjx/_src/collision_driver.py b/mjx/mujoco/mjx/_src/collision_driver.py index 2f45bca8..158be7a2 100644 --- a/mjx/mujoco/mjx/_src/collision_driver.py +++ b/mjx/mujoco/mjx/_src/collision_driver.py @@ -76,6 +76,7 @@ from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import GeomType from mujoco.mjx._src.types import Model from mujoco.mjx._src.types import ModelJAX +from mujoco.mjx._src.types import OptionJAX # pylint: enable=g-importing-member import numpy as np @@ -342,6 +343,16 @@ def _numeric(m: Union[Model, mujoco.MjModel], name: str) -> int: def make_condim(m: Union[Model, mujoco.MjModel]) -> np.ndarray: """Returns the dims of the contacts for a Model.""" + if isinstance(m, mujoco.MjModel): + sdf_initpoints = m.opt.sdf_initpoints + elif isinstance(m.opt._impl, OptionJAX): + sdf_initpoints = m.opt._impl.sdf_initpoints + else: + raise ValueError( + 'make_condim requires mujoco.MjModel or mjx.Model with JAX backend' + ' implementation.' + ) + if m.opt.disableflags & DisableBit.CONTACT: return np.empty(0, dtype=int) @@ -364,7 +375,7 @@ def make_condim(m: Union[Model, mujoco.MjModel]) -> np.ndarray: condim_counts = {} for k, v in group_counts.items(): if k.types[1] == mujoco.mjtGeom.mjGEOM_SDF: - ncon = m.opt.sdf_initpoints + ncon = sdf_initpoints else: func = _COLLISION_FUNC[k.types] ncon = func.ncon # pytype: disable=attribute-error diff --git a/mjx/mujoco/mjx/_src/constraint.py b/mjx/mujoco/mjx/_src/constraint.py index 9b41daa7..410783f6 100644 --- a/mjx/mujoco/mjx/_src/constraint.py +++ b/mjx/mujoco/mjx/_src/constraint.py @@ -35,6 +35,7 @@ from mujoco.mjx._src.types import JointType from mujoco.mjx._src.types import Model from mujoco.mjx._src.types import ModelJAX from mujoco.mjx._src.types import ObjType +from mujoco.mjx._src.types import OptionJAX # pylint: enable=g-importing-member import numpy as np @@ -496,7 +497,11 @@ def _efc_contact_frictionless(m: Model, d: Data) -> Optional[_Efc]: def _efc_contact_pyramidal(m: Model, d: Data, condim: int) -> Optional[_Efc]: """Calculates constraint rows for frictional pyramidal contacts.""" - if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + if ( + not isinstance(m._impl, ModelJAX) + or not isinstance(d._impl, DataJAX) + or not isinstance(m.opt._impl, OptionJAX) + ): raise ValueError( '_efc_contact_pyramidal requires JAX backend implementation.' ) @@ -545,7 +550,11 @@ def _efc_contact_pyramidal(m: Model, d: Data, condim: int) -> Optional[_Efc]: def _efc_contact_elliptic(m: Model, d: Data, condim: int) -> Optional[_Efc]: """Calculates constraint rows for frictional elliptic contacts.""" - if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + if ( + not isinstance(m._impl, ModelJAX) + or not isinstance(d._impl, DataJAX) + or not isinstance(m.opt._impl, OptionJAX) + ): raise ValueError( '_efc_contact_elliptic requires JAX backend implementation.' ) diff --git a/mjx/mujoco/mjx/_src/dataclasses.py b/mjx/mujoco/mjx/_src/dataclasses.py index b8bdb758..e82d6047 100644 --- a/mjx/mujoco/mjx/_src/dataclasses.py +++ b/mjx/mujoco/mjx/_src/dataclasses.py @@ -62,11 +62,19 @@ def dataclass(clz: _T, register_as_pytree: bool) -> _T: def iterate_clz_with_keys(x): def to_meta(field, obj): val = getattr(obj, field.name) - # numpy arrays are not hashable so return raw bytes instead if isinstance(val, np.ndarray): + # numpy arrays are not hashable so return raw bytes instead return (val.tobytes(), val.dtype, val.shape) - else: - return val + if typing.get_origin(field.type) == tuple: + # variadic tuples of numpy arrays + type_args = typing.get_args(field.type) + if ( + len(type_args) == 2 + and type_args[0] == np.ndarray + and type_args[1] == ... + ): + return tuple((v.tobytes(), v.dtype, v.shape) for v in val) + return val def to_data(field, obj): return (jax.tree_util.GetAttrKey(field.name), getattr(obj, field.name)) @@ -81,8 +89,20 @@ def dataclass(clz: _T, register_as_pytree: bool) -> _T: if field.type is np.ndarray: arr = np.frombuffer(meta[0], dtype=meta[1]).reshape(meta[2]) return (field.name, arr) - else: - return (field.name, meta) + if typing.get_origin(field.type) == tuple: + type_args = typing.get_args(field.type) + if ( + len(type_args) == 2 + and type_args[0] == np.ndarray + and type_args[1] == ... + ): + return ( + field.name, + tuple( + np.frombuffer(m[0], dtype=m[1]).reshape(m[2]) for m in meta + ), + ) + return (field.name, meta) from_data = lambda field, meta: (field.name, meta) diff --git a/mjx/mujoco/mjx/_src/dataclasses_test.py b/mjx/mujoco/mjx/_src/dataclasses_test.py new file mode 100644 index 00000000..b5727941 --- /dev/null +++ b/mjx/mujoco/mjx/_src/dataclasses_test.py @@ -0,0 +1,69 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Tests for custom PyTreeNode object.""" + +from absl.testing import absltest +import jax +from jax import numpy as jp +from mujoco.mjx._src import dataclasses +import numpy as np + + +class Obj(dataclasses.PyTreeNode): + a: int + b: np.ndarray + c: tuple[int, ...] + d: tuple[np.ndarray, ...] + e: jax.Array + f: tuple[jax.Array, ...] + + +class DataclassesTest(absltest.TestCase): + + def test_pytree_structure(self): + obj = Obj( + a=1, + b=np.array([1, 2, 3]), + c=(4, 5, 6), + d=(np.array([7, 8]), np.array([9, 10])), + e=jax.numpy.array([11, 12]), + f=(jax.numpy.array([13, 14]), jax.numpy.array([15, 16])), + ) + + data, meta = jax.tree_util.tree_flatten_with_path(obj) + + # data fields + self.assertLen(data, 3) + self.assertEqual(data[0][0][0].name, 'e') + np.testing.assert_array_equal(data[0][1], jp.array([11, 12])) + self.assertEqual(data[1][0][0].name, 'f') + np.testing.assert_array_equal(data[1][1], jp.array([13, 14])) + self.assertEqual(data[2][0][0].name, 'f') + np.testing.assert_array_equal(data[2][1], jp.array([15, 16])) + + # meta fields + unflattened_meta = meta.unflatten([x[1] for x in data]) + self.assertEqual(unflattened_meta.a, 1) + np.testing.assert_array_equal(unflattened_meta.b, np.array([1, 2, 3])) + self.assertEqual(unflattened_meta.c, (4, 5, 6)) + np.testing.assert_array_equal(unflattened_meta.d[0], np.array([7, 8])) + np.testing.assert_array_equal(unflattened_meta.d[1], np.array([9, 10])) + + # ensure hashable meta + hash(meta) + + +if __name__ == '__main__': + absltest.main() diff --git a/mjx/mujoco/mjx/_src/derivative.py b/mjx/mujoco/mjx/_src/derivative.py index 4e3d6d68..cd56ad06 100644 --- a/mjx/mujoco/mjx/_src/derivative.py +++ b/mjx/mujoco/mjx/_src/derivative.py @@ -21,14 +21,24 @@ from jax import numpy as jp # pylint: disable=g-importing-member from mujoco.mjx._src.types import BiasType from mujoco.mjx._src.types import Data +from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import DynType from mujoco.mjx._src.types import GainType from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import ModelJAX +from mujoco.mjx._src.types import OptionJAX +# pylint: enable=g-importing-member def deriv_smooth_vel(m: Model, d: Data) -> Optional[jax.Array]: """Analytical derivative of smooth forces w.r.t. velocities.""" + if ( + not isinstance(m._impl, ModelJAX) + or not isinstance(d._impl, DataJAX) + or not isinstance(m.opt._impl, OptionJAX) + ): + raise ValueError('deriv_smooth_vel requires JAX MJX implementation.') qderiv = None @@ -53,7 +63,7 @@ def deriv_smooth_vel(m: Model, d: Data) -> Optional[jax.Array]: if m.ntendon: qderiv -= d._impl.ten_J.T @ jp.diag(m.tendon_damping) @ d._impl.ten_J # TODO(robotics-simulation): fluid drag model - if m.opt.has_fluid_params: # pytype: disable=attribute-error + if m.opt._impl.has_fluid_params: # pytype: disable=attribute-error raise NotImplementedError('fluid drag not supported for implicitfast') # TODO(team): rne derivative diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py index c760b0ff..bb01727d 100644 --- a/mjx/mujoco/mjx/_src/forward.py +++ b/mjx/mujoco/mjx/_src/forward.py @@ -37,12 +37,14 @@ from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import DynType from mujoco.mjx._src.types import GainType +from mujoco.mjx._src.types import Impl from mujoco.mjx._src.types import IntegratorType from mujoco.mjx._src.types import JointType from mujoco.mjx._src.types import Model from mujoco.mjx._src.types import ModelJAX from mujoco.mjx._src.types import TrnType # pylint: enable=g-importing-member +import mujoco.mjx.warp as mjxw import numpy as np # RK4 tableau @@ -423,6 +425,10 @@ def implicit(m: Model, d: Data) -> Data: @named_scope def forward(m: Model, d: Data) -> Data: """Forward dynamics.""" + if m.impl == Impl.WARP and d.impl == Impl.WARP and mjxw.WARP_INSTALLED: + from mujoco.mjx.warp import forward as mjxw_forward # pylint: disable=g-import-not-at-top # pytype: disable=import-error + return mjxw_forward.forward(m, d) + if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): raise ValueError('forward requires JAX backend implementation.') @@ -446,6 +452,10 @@ def forward(m: Model, d: Data) -> Data: @named_scope def step(m: Model, d: Data) -> Data: """Advance simulation.""" + if m.impl == Impl.WARP and d.impl == Impl.WARP and mjxw.WARP_INSTALLED: + from mujoco.mjx.warp import forward as mjxw_forward # pylint: disable=g-import-not-at-top # pytype: disable=import-error + return mjxw_forward.step(m, d) + d = forward(m, d) if m.opt.integrator == IntegratorType.EULER: diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index c690011c..fe2dabd6 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -22,23 +22,31 @@ import warnings import jax from jax import numpy as jp +from jax.extend import backend import mujoco from mujoco.mjx._src import collision_driver from mujoco.mjx._src import constraint from mujoco.mjx._src import mesh from mujoco.mjx._src import support from mujoco.mjx._src import types +import mujoco.mjx.warp as mjxw +# pylint: disable=g-importing-member +from mujoco.mjx.warp import mjwp_types +from mujoco.mjx.warp import mujoco_warp as mjwp +from mujoco.mjx.warp import warp as wp +# pylint: enable=g-importing-member import numpy as np import scipy +def has_cuda_gpu_device() -> bool: + return 'cuda' in backend.backends() + + def _is_cuda_gpu_device(device: jax.Device) -> bool: - try: - cuda_devices = jax.devices('cuda') - except RuntimeError: - logging.info('No CUDA GPU devices found in jax.devices("cuda").') + if not has_cuda_gpu_device(): return False - return device in cuda_devices + return device in jax.devices('cuda') def _resolve_impl( @@ -46,24 +54,19 @@ def _resolve_impl( ) -> types.Impl: """Pick a default implementation based on the device specified.""" if _is_cuda_gpu_device(device): - # TODO(btaba): Remove flag once Warp is ready to launch. - mjx_warp_enabled = os.environ.get('MJX_WARP_ENABLED', 'f').lower() == 'true' - if mjx_warp_enabled: + # TODO(btaba): Remove flag once Warp is ready for GPU default. + mjx_gpu_default_warp = ( + os.environ.get('MJX_GPU_DEFAULT_WARP', 'f').lower() == 'true' + ) + if mjx_gpu_default_warp and mjxw.WARP_INSTALLED: logging.debug('Picking default implementation: Warp.') return types.Impl.WARP - logging.info('MJX Warp is disabled via MJX_WARP_ENABLED=false.') if device.platform in ('gpu', 'tpu'): logging.debug('Picking default implementation: JAX.') return types.Impl.JAX if device.platform == 'cpu': - mjx_c_default = ( - os.environ.get('MJX_C_DEFAULT_ENABLED', 'f').lower() == 'true' - ) - if mjx_c_default: - logging.debug('Picking default implementation: C.') - return types.Impl.C return types.Impl.JAX raise ValueError(f'Unsupported device: {device}') @@ -115,11 +118,9 @@ def _check_impl_device_compatibility( 'Warp implementation requires a CUDA GPU device, got ' f'{device}.' ) - - mjx_warp_enabled = os.environ.get('MJX_WARP_ENABLED', 'f').lower() == 'true' - if not mjx_warp_enabled: - raise AssertionError( - 'Warp implementation is disabled via MJX_WARP_ENABLED=false.' + if not mjxw.WARP_INSTALLED: + raise RuntimeError( + 'Warp is not installed. Cannot use Warp implementation of MJX.' ) is_cpu_device = device.platform == 'cpu' @@ -165,10 +166,53 @@ def _strip_weak_type(tree): return jax.tree_util.tree_map(f, tree) +def _wp_to_np_type(wp_field: Any, name: str = '') -> Any: + """Converts a warp type to an MJX compatible numpy type.""" + if hasattr(wp_field, '_is_batched'): + wp_field.strides = wp_field.strides[1:] + wp_field.shape = wp_field.shape[1:] + # warp scalars + wp_dtype = type(wp_field) + if wp_dtype in wp.types.warp_type_to_np_dtype: + return wp.types.warp_type_to_np_dtype[wp_dtype](wp_field) + + # warp arrays + if isinstance(wp_field, wp.types.array): + return wp_field.numpy() + + # static + static_types = (bool, int, float, np.bool, np.int32, np.int64, + np.float32, np.float64) # fmt: skip + is_static = lambda x: isinstance(x, static_types) + if is_static(wp_field): + return wp_field + + # tuples + if isinstance(wp_field, tuple) and len(wp_field) == 0: + return () + if isinstance(wp_field, tuple) and isinstance(wp_field[0], wp.types.array): + return tuple(f.numpy() for f in wp_field) + if isinstance(wp_field, tuple) and isinstance( + wp_field[0], mjwp_types.TileSet + ): + return tuple( + mjxw.types.TileSet(wp_field[i].adr.numpy(), wp_field[i].size) + for i in range(len(wp_field)) + ) + if isinstance(wp_field, mjwp_types.BlockDim): + return mjxw.types.BlockDim(**wp_field.__dict__) + if isinstance(wp_field, tuple) and is_static(wp_field[0]): + return wp_field + + raise NotImplementedError( + f'Field {name} has unsupported type {type(wp_field)}.' + ) + + def _put_option( o: mujoco.MjOption, impl: types.Impl, - impl_fields: Optional[dict[str, Any]] = None, + impl_fields: Optional[Dict[str, Any]] = None, ) -> types.Option: """Returns mjx.Option given mujoco.MjOption.""" if o.integrator not in set(types.IntegratorType): @@ -187,32 +231,56 @@ def _put_option( if o.enableflags & 2**i and 2 ** i not in set(types.EnableBit): raise NotImplementedError(f'{mujoco.mjtEnableBit(2**i)}') - fields = {f.name: getattr(o, f.name, None) for f in types.Option.fields()} + fields = { + f.name: getattr(o, f.name, None) + for f in types.Option.fields() + if f.name != '_impl' + } fields['integrator'] = types.IntegratorType(o.integrator) fields['cone'] = types.ConeType(o.cone) - fields['jacobian'] = types.JacobianType(o.jacobian) fields['solver'] = types.SolverType(o.solver) fields['disableflags'] = types.DisableBit(o.disableflags) fields['enableflags'] = types.EnableBit(o.enableflags) + fields['jacobian'] = types.JacobianType(o.jacobian) + + option_obj = { + types.Impl.C: types.OptionC, + types.Impl.JAX: types.OptionJAX, + types.Impl.WARP: mjxw.types.OptionWarp, + }[impl] + private_fields = { + f.name: getattr(o, f.name, None) for f in option_obj.fields() + } + impl_fields = impl_fields or {} + impl_fields = {**private_fields, **impl_fields} if impl == types.Impl.JAX: has_fluid_params = o.density > 0 or o.viscosity > 0 or o.wind.any() implicitfast = o.integrator == mujoco.mjtIntegrator.mjINT_IMPLICITFAST if implicitfast and has_fluid_params: raise NotImplementedError('implicitfast not implemented for fluid drag.') - fields['has_fluid_params'] = has_fluid_params - return types.OptionJAX(**fields, **(impl_fields or {})) + impl_fields['has_fluid_params'] = has_fluid_params + return types.Option(**fields, _impl=types.OptionJAX(**impl_fields)) if impl == types.Impl.C: - c_field_keys = types.OptionC.__annotations__.keys() - fields.keys() - c_fields = {k: getattr(o, k, None) for k in c_field_keys} - return types.OptionC(**fields, **c_fields, **(impl_fields or {})) + return types.Option(**fields, _impl=types.OptionC(**impl_fields)) + + if impl == types.Impl.WARP: + impl_fields = {k: _wp_to_np_type(v) for k, v in impl_fields.items()} + return types.Option(**fields, _impl=mjxw.types.OptionWarp(**impl_fields)) raise NotImplementedError(f'Unsupported implementation: {impl}') -def _put_statistic(s: mujoco.MjStatistic) -> types.Statistic: +def _put_statistic( + s: mujoco.MjStatistic, impl: types.Impl +) -> Union[types.Statistic, types.StatisticWarp]: """Puts mujoco.MjStatistic onto a device, resulting in mjx.Statistic.""" + if impl == types.Impl.WARP: + fields = { + f.name: getattr(s, f.name, None) for f in types.StatisticWarp.fields() + } + return types.StatisticWarp(**fields) return types.Statistic( meaninertia=s.meaninertia, meanmass=s.meanmass, @@ -305,7 +373,7 @@ def _put_model_jax( fields = {f: getattr(m, f) for f in mj_field_names} fields['cam_mat0'] = fields['cam_mat0'].reshape((-1, 3, 3)) fields['opt'] = _put_option(m.opt, types.Impl.JAX) - fields['stat'] = _put_statistic(m.stat) + fields['stat'] = _put_statistic(m.stat, types.Impl.JAX) fields_jax = {} fields_jax['dof_hasfrictionloss'] = fields['dof_frictionloss'] > 0 @@ -362,7 +430,7 @@ def _put_model_c( fields = {f: getattr(m, f) for f in mj_field_names} fields['cam_mat0'] = fields['cam_mat0'].reshape((-1, 3, 3)) fields['opt'] = _put_option(m.opt, impl=types.Impl.C) - fields['stat'] = _put_statistic(m.stat) + fields['stat'] = _put_statistic(m.stat, impl=types.Impl.C) c_impl_keys = ( types.ModelC.__annotations__.keys() - types.Model.__annotations__.keys() @@ -377,6 +445,51 @@ def _put_model_c( return _strip_weak_type(model) +def _put_model_warp( + m: mujoco.MjModel, + device: Optional[jax.Device] = None, +) -> types.Model: + """Puts mujoco.MjModel onto a device, resulting in mjx.Model.""" + if not mjxw.WARP_INSTALLED: + raise RuntimeError('Warp not installed.') + + with wp.ScopedDevice('cpu'): # pylint: disable=undefined-variable + mw = mjwp.put_model(m) # pylint: disable=undefined-variable + mw.opt.graph_conditional = False + + fields = {f.name for f in types.Model.fields() if f.name != '_impl'} + fields = {f: getattr(m, f) for f in fields} + # Grab MJW private Option fields, and assume that public MjOption fields are + # directly compatible with MJXW. + option_keys = {f.name for f in mjxw.types.OptionWarp.fields()} - { + f.name for f in types.Option.fields() + } + private_options = {k: getattr(mw.opt, k) for k in option_keys} + fields['opt'] = _put_option(m.opt, types.Impl.WARP, private_options) + fields['stat'] = _put_statistic(m.stat, types.Impl.WARP) + + # Use MJW fields directly instead of MjModel, so that shape and dtype are + # always compatible with MJXW (e.g. cam_mat0/geom_aabb). + for k in fields: + if not hasattr(mw, k) or k in ('stat', 'opt'): + continue + field = _wp_to_np_type(getattr(mw, k), k) + fields[k] = field + + impl_fields = {} + for k in mjxw.types.ModelWarp.__annotations__.keys(): + field = _wp_to_np_type(getattr(mw, k), k) + impl_fields[k] = field + + model = types.Model( + **fields, + _impl=mjxw.types.ModelWarp(**impl_fields), + ) + + model = jax.device_put(model, device=device) + return _strip_weak_type(model) + + def put_model( m: mujoco.MjModel, device: Optional[jax.Device] = None, @@ -416,7 +529,7 @@ def put_model( elif impl == types.Impl.C: return _put_model_c(m, device) elif impl == types.Impl.WARP: - raise NotImplementedError('Warp implementation not implemented yet.') + return _put_model_warp(m, device) else: raise ValueError(f'Unsupported implementation: {impl}') @@ -719,11 +832,81 @@ def _make_data_c( return d +def _get_nested_attr(obj: Any, attr_name: str, split: str) -> Any: + """Returns the nested attribute from an object.""" + for part in attr_name.split(split): + obj = getattr(obj, part) + return obj + + +def _make_data_warp( + m: Union[types.Model, mujoco.MjModel], + device: Optional[jax.Device] = None, + nconmax: int = -1, + njmax: int = -1, +) -> types.Data: + """Allocate and initialize Data for the Warp implementation.""" + if not isinstance(m, mujoco.MjModel): + raise ValueError( + 'make_data for warp, only supports a mujoco.MjModel input, got' + f' {type(m)}.' + ) + + if not mjxw.WARP_INSTALLED: + raise RuntimeError('Warp is not installed.') + + with wp.ScopedDevice('cpu'): # pylint: disable=undefined-variable + dw = mjwp.make_data(m, nworld=1, nconmax=nconmax, njmax=njmax) # pylint: disable=undefined-variable + + fields = _make_data_public_fields(m) + for k in fields: + if k in {'userdata', 'plugin_state'}: + continue + if not hasattr(dw, k): + raise ValueError(f'Public data field {k} not found in Warp data.') + field = _wp_to_np_type(getattr(dw, k)) + if mjxw.types.BATCH_DIM['Data'][k]: + field = field.reshape(field.shape[1:]) + fields[k] = field + + impl_fields = {} + for k in mjxw.types.DataWarp.__annotations__.keys(): + field = _get_nested_attr(dw, k, split='__') + field = _wp_to_np_type(field) + if mjxw.types.BATCH_DIM['Data'][k]: + field = field.reshape(field.shape[1:]) + impl_fields[k] = field + + data = types.Data( + qpos=m.qpos0.astype(np.float32), + eq_active=m.eq_active0.astype(bool), + **fields, + _impl=mjxw.types.DataWarp(**impl_fields), + ) + + data = jax.device_put(data, device=device) + + with wp.ScopedDevice('cuda:0'): # pylint: disable=undefined-variable + # Warm-up the warp kernel cache. + # TODO(robotics-simulation): remove this warmup compilation once warp + # stops unloading modules during XLA graph capture for tile kernels. + # pylint: disable=undefined-variable + dw = mjwp.make_data(m, nworld=1) + mw = mjwp.put_model(m) + _ = mjwp.step(mw, dw) + # pylint: enable=undefined-variable + del dw, mw + + return data + + def make_data( m: Union[types.Model, mujoco.MjModel], device: Optional[jax.Device] = None, impl: Optional[Union[str, types.Impl]] = None, _full_compat: bool = False, # pylint: disable=invalid-name + nconmax: int = -1, + njmax: int = -1, ) -> types.Data: """Allocate and initialize Data. @@ -734,6 +917,8 @@ def make_data( _full_compat: put all fields onto device irrespective of MJX support This is an experimental feature. Avoid using it for now. If using this flag, also use _full_compat for put_model. + nconmax: maximum number of contacts to allocate for warp + njmax: maximum number of constraints to allocate for warp Returns: an initialized mjx.Data placed on device @@ -764,6 +949,8 @@ def make_data( return _make_data_jax(m, device) elif impl == types.Impl.C: return _make_data_c(m, device) + elif impl == types.Impl.WARP: + return _make_data_warp(m, device, nconmax, njmax) raise NotImplementedError( f'make_data for implementation "{impl}" not implemented yet.' @@ -1053,6 +1240,8 @@ def put_data( d: mujoco.MjData, device: Optional[jax.Device] = None, impl: Optional[Union[str, types.Impl]] = None, + nconmax: int = -1, + njmax: int = -1, _full_compat: bool = False, # pylint: disable=invalid-name ) -> types.Data: """Puts mujoco.MjData onto a device, resulting in mjx.Data. @@ -1062,6 +1251,8 @@ def put_data( d: the data to put on device device: which device to use - if unspecified picks the default device impl: implementation to use ('jax', 'warp') + nconmax: maximum number of contacts to allocate for warp + njmax: maximum number of constraints to allocate for warp _full_compat: put all MjModel fields onto device irrespective of MJX support This is an experimental feature. Avoid using it for now. If using this flag, also use _full_compat for put_model. @@ -1069,6 +1260,7 @@ def put_data( Returns: an mjx.Data placed on device """ + del nconmax, njmax if _full_compat: warnings.warn( 'mjx.put_data(..., _full_compat=True) is deprecated. Use' @@ -1084,6 +1276,8 @@ def put_data( elif impl == types.Impl.C: return _put_data_c(m, d, device) + # TODO(robotics-team): implement put_data_warp + raise NotImplementedError( f'put_data for implementation "{impl}" not implemented yet.' ) @@ -1099,6 +1293,86 @@ def _get_contact(c: mujoco._structs._MjContactList, cx: types.Contact): getattr(c, field.name)[:] = value +def _get_data_into_warp( + result: Union[mujoco.MjData, List[mujoco.MjData]], + m: mujoco.MjModel, + d: types.Data, +): + """Gets mjx.Data from a device into an existing mujoco.MjData or list.""" + batched = isinstance(result, list) + d = jax.device_get(d) + batch_size = d.qpos.shape[0] if batched else 1 + + for i in range(batch_size): + d_i = ( + jax.tree.map_with_path( + lambda path, x, i=i: x[i] + if path[-1].name not in mjxw.types.DATA_NON_VMAP + else x, + d, + ) + if batched + else d + ) + result_i = result[i] if batched else result + ncon = d_i._impl.ncon[0] + nefc = int(d_i._impl.nefc[0]) + # nj = int(d_i._impl.nj[0]) + nj = 0 # TODO(btaba): add nj back + + if ncon != result_i.ncon or nefc != result_i.nefc or nj != result_i.nJ: + mujoco._functions._realloc_con_efc(result_i, ncon=ncon, nefc=nefc, nJ=nj) # pylint: disable=protected-access + + all_fields = types.Data.fields() + mjxw.types.DataWarp.fields() + for field in all_fields: + if field.name not in mujoco.MjData.__dict__.keys(): + continue + + # TODO(btaba): contact + # TODO(btaba): actuator_moment + + if hasattr(d_i._impl, field.name): + value = getattr(d_i._impl, field.name) + else: + value = getattr(d_i, field.name) + + if field.name in ('ne', 'nl', 'nf'): + value = value[0] + elif field.name in ('nefc', 'ncon'): + value = {'nefc': nefc, 'ncon': ncon}[field.name] + elif field.name.endswith('xmat') or field.name == 'ximat': + value = value.reshape((-1, 9)) + # elif field.name == 'efc_J': # TODO(btaba): add this back + # elif field.name.startswith('efc_'): # TODO(btaba): add this back + # TODO(btaba): qM, qLD, qLDiagInv + + if field.name in ( + 'actuator_moment', + 'contact', + 'efc_J', + 'qM', + 'qLD', + 'qLDiagInv', + ): + continue + if field.name.startswith('efc_'): + continue + + if isinstance(value, np.ndarray) and value.shape: + result_field = getattr(result_i, field.name) + if result_field.shape != value.shape: + raise ValueError( + f'Input field {field.name} has shape {value.shape}, but output' + f' has shape {result_field.shape}' + ) + result_field[:] = value + else: + setattr(result_i, field.name, value) + + # TODO(btaba): add M back + # mujoco.mj_factorM(m, result_i) + + def _get_data_into( result: Union[mujoco.MjData, List[mujoco.MjData]], m: mujoco.MjModel, @@ -1250,6 +1524,9 @@ def get_data_into( # TODO(stunya): Split out _get_data_into once codepaths diverge enough. return _get_data_into(result, m, d) + if d.impl == types.Impl.WARP: + return _get_data_into_warp(result, m, d) + raise NotImplementedError( f'get_data_into for implementation "{d.impl}" not implemented yet.' ) diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 3edef119..b2fc6e11 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -24,12 +24,13 @@ import mujoco from mujoco import mjx from mujoco.mjx._src import io as mjx_io from mujoco.mjx._src import test_util - # pylint: disable=g-importing-member from mujoco.mjx._src.types import ConeType from mujoco.mjx._src.types import Impl from mujoco.mjx._src.types import JacobianType # pylint: enable=g-importing-member +import mujoco.mjx.warp as mjxw +from mujoco.mjx.warp import types as mjxw_types import numpy as np @@ -110,14 +111,31 @@ _SIMPLE_BODY = """ """ +def _get_name_from_path(path: jax.tree_util.KeyPath) -> str: + """Returns a flattened name from a jax.tree_util.KeyPath.""" + if any(isinstance(p, jax.tree_util.SequenceKey) for p in path): + is_seq_key = [isinstance(p, jax.tree_util.SequenceKey) for p in path] + path = path[: is_seq_key.index(True)] + assert all(isinstance(p, jax.tree_util.GetAttrKey) for p in path) + path = [p for p in path if p.name != '_impl'] + attr = '__'.join(p.name for p in path) + return attr + + class ModelIOTest(parameterized.TestCase): """IO tests for mjx.Model.""" @parameterized.product( xml=(_MULTIPLE_CONVEX_OBJECTS, _MULTIPLE_CONSTRAINTS), - impl=('jax', 'c'), + impl=('jax', 'c', 'warp'), ) + @mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'true'}) def test_put_model(self, xml, impl): + if impl == 'warp' and not mjxw.WARP_INSTALLED: + self.skipTest('Warp not installed.') + if impl == 'warp' and not mjx_io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device available.') + m = mujoco.MjModel.from_xml_string(xml) mx = mjx.put_model(m, impl=impl) @@ -146,9 +164,14 @@ class ModelIOTest(parameterized.TestCase): self.assertFalse(hasattr(mx, 'bvh_aabb')) elif impl == 'c': # Options specific to C are populated. - self.assertEqual(mx.opt.apirate, m.opt.apirate) + self.assertEqual(mx.opt._impl.apirate, m.opt.apirate) # Fields private to C backend impl are populated. self.assertTrue(hasattr(mx._impl, 'bvh_aabb')) + elif impl == 'warp': + # Options specific to Warp are populated. + self.assertTrue(hasattr(mx.opt._impl, 'ls_parallel')) + # Fields private to Warp backend impl are populated. + self.assertTrue(hasattr(mx._impl, 'nxn_geom_pair')) np.testing.assert_allclose(mx.body_parentid, m.body_parentid) np.testing.assert_allclose(mx.geom_type, m.geom_type) @@ -170,7 +193,7 @@ class ModelIOTest(parameterized.TestCase): np.testing.assert_equal(mx.wrap_type, m.wrap_type) np.testing.assert_equal(mx.wrap_objid, m.wrap_objid) - np.testing.assert_equal(mx.wrap_prm, m.wrap_prm) + np.testing.assert_almost_equal(mx.wrap_prm, m.wrap_prm) def test_fluid_params(self): """Test that has_fluid_params is set when fluid params are present.""" @@ -180,7 +203,7 @@ class ModelIOTest(parameterized.TestCase): ), impl='jax', ) - self.assertTrue(m.opt.has_fluid_params) + self.assertTrue(m.opt._impl.has_fluid_params) def test_implicit_not_implemented(self): """Test that MJX guards against models with unimplemented features.""" @@ -276,6 +299,29 @@ class ModelIOTest(parameterized.TestCase): with self.assertRaises(NotImplementedError): mjx.put_model(m, impl='jax') + def test_put_model_warp_has_expected_shapes(self): + """Tests that put_model produces expected shapes for MuJoCo Warp.""" + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp not installed.') + if not mjx_io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device available.') + + m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) + mx = mjx.put_model(m, impl='warp') + + def check_ndim(path, x): + k = _get_name_from_path(path) + if k not in mjxw_types.NDIM['Model']: + return + is_batched = mjxw_types.BATCH_DIM['Model'][k] + expected_ndim = mjxw_types.NDIM['Model'][k] - is_batched + if not hasattr(x, 'ndim'): + return + msg = f'Field {k} has ndim {x.ndim} but expected {expected_ndim}' + self.assertEqual(x.ndim, expected_ndim, msg) + + _ = jax.tree.map_with_path(check_ndim, mx) + class DataIOTest(parameterized.TestCase): """IO tests for mjx.Data.""" @@ -366,6 +412,17 @@ class DataIOTest(parameterized.TestCase): self.assertEqual(d._impl.light_xpos.shape, (m.nlight, 3)) self.assertEqual(d._impl.bvh_active.shape, (m.nbvh,)) + @mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'true'}) + def test_make_data_warp(self): + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp is not installed.') + if not mjx_io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device.') + m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONVEX_OBJECTS) + d = mjx.make_data(m, impl='warp', nconmax=9, njmax=11) + self.assertEqual(d._impl.contact__dist.shape[0], 9) + self.assertEqual(d._impl.efc__J.shape[0], 11) + @parameterized.parameters('jax', 'c') def test_put_data(self, impl: str): """Test that put_data puts the correct data for dense and sparse.""" @@ -676,6 +733,47 @@ class DataIOTest(parameterized.TestCase): np.testing.assert_allclose(res_mj[0], res, rtol=1e-3, atol=1e-3) + def test_make_data_warp_has_expected_shapes(self): + """Tests that make_data produces expected shapes for MuJoCo Warp.""" + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp is not installed.') + if not mjx_io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device.') + + m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) + dx = mjx.make_data(m, impl='warp') + + def check_ndim(path, x): + k = _get_name_from_path(path) + if k not in mjxw_types.NDIM['Data']: + return + is_batched = mjxw_types.BATCH_DIM['Data'][k] + expected_ndim = mjxw_types.NDIM['Data'][k] - is_batched + if not hasattr(x, 'ndim'): + return + msg = f'Field {k} has ndim {x.ndim} but expected {expected_ndim}' + self.assertEqual(x.ndim, expected_ndim, msg) + + _ = jax.tree.map_with_path(check_ndim, dx) + + @parameterized.parameters('jax', 'warp') + def test_data_slice(self, impl): + """Tests that slice on Data works as expected.""" + if impl == 'warp' and not mjxw.WARP_INSTALLED: + self.skipTest('Warp is not installed.') + if impl == 'warp' and not mjx_io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device.') + + m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) + dx = jax.vmap(lambda x: mjx.make_data(m, impl=impl))(jp.arange(10)) + + self.assertEqual(dx.qpos.shape, (10, m.nq)) + self.assertEqual(dx[0].qpos.shape, (m.nq,)) + + if impl == 'warp': + self.assertEqual(dx._impl.contact__dist.shape, (dx._impl.nconmax,)) + self.assertEqual(dx[0]._impl.contact__dist.shape, (dx._impl.nconmax,)) + class FullCompatTest(parameterized.TestCase): """Tests for the _full_compat flag.""" @@ -711,7 +809,7 @@ _DEVICE_TEST_CASES = [ # (device_type_str, impl_str, # (expected_device, expected_impl))) # No backend specified. - ('cpu', None, ('cpu', Impl.C)), + ('cpu', None, ('cpu', Impl.JAX)), ('gpu-notnvidia', None, ('gpu', Impl.JAX)), ('gpu-nvidia', None, ('gpu', Impl.WARP)), ('tpu', None, ('tpu', Impl.JAX)), @@ -739,7 +837,7 @@ _DEFAULT_DEVICE_TEST_CASES = [ # (jax.default_device, impl_str, # (expected_device, expected_impl)) # No backend impl specified. - ('cpu', None, ('cpu', Impl.C)), + ('cpu', None, ('cpu', Impl.JAX)), ('gpu-notnvidia', None, ('gpu', Impl.JAX)), ('gpu-nvidia', None, ('gpu', Impl.WARP)), ('tpu', None, ('tpu', Impl.JAX)), @@ -790,6 +888,9 @@ class ResolveImplAndDeviceTest(parameterized.TestCase): # Patch jax.devices for the entire test class using enter_context self.mock_jax_devices = self.enter_context(mock.patch('jax.devices')) + self.mock_jax_backends = self.enter_context( + mock.patch('jax.extend.backend.backends') + ) self.mock_default_backend = self.enter_context( mock.patch('jax.default_backend') ) @@ -797,9 +898,7 @@ class ResolveImplAndDeviceTest(parameterized.TestCase): @parameterized.named_parameters( (f'{str(args[0])}_{str(args[1])}', *args) for args in _DEVICE_TEST_CASES ) - @mock.patch.dict( - os.environ, {'MJX_WARP_ENABLED': 'true', 'MJX_C_DEFAULT_ENABLED': 'true'} - ) + @mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'true'}) def test_resolve_with_device( self, device_type_str, @@ -819,7 +918,7 @@ class ResolveImplAndDeviceTest(parameterized.TestCase): if backend == 'cpu': return [self.mock_cpu] elif backend == 'gpu': - if 'nvidia' in device_type_str: + if device_type_str == 'gpu-nvidia': return [self.mock_nvidia_gpu] return [self.mock_other_gpu] elif backend == 'tpu': @@ -828,9 +927,19 @@ class ResolveImplAndDeviceTest(parameterized.TestCase): return [self.mock_nvidia_gpu] raise AssertionError('Should not be called.') - self.mock_jax_devices.side_effect = devices_side_effect + def backends_side_effect(): + if device_type_str == 'gpu-nvidia': + return ['cuda', 'cpu'] + if 'tpu' in device_type_str: + return ['tpu', 'cpu'] + if 'gpu' in device_type_str: + return ['gpu', 'cpu'] + return ['cpu'] + + self.mock_jax_backends.side_effect = backends_side_effect + expected_device, expected_impl = expected if expected_impl == 'error': with self.assertRaises(AssertionError): @@ -839,6 +948,14 @@ class ResolveImplAndDeviceTest(parameterized.TestCase): ) return + if impl_str == 'warp' and not mjxw.WARP_INSTALLED: + with self.assertRaisesRegex(RuntimeError, 'not installed'): + mjx_io._resolve_impl_and_device(impl=impl_str, device=input_device) + return + + if impl_str is None and not mjxw.WARP_INSTALLED: + expected_impl = Impl.JAX + actual_impl, actual_device = ( mjx_io._resolve_impl_and_device( impl=impl_str, device=input_device @@ -853,9 +970,7 @@ class ResolveImplAndDeviceTest(parameterized.TestCase): (f'{str(args[0])}_{str(args[1])}', *args) for args in _DEFAULT_DEVICE_TEST_CASES ) - @mock.patch.dict( - os.environ, {'MJX_WARP_ENABLED': 'true', 'MJX_C_DEFAULT_ENABLED': 'true'} - ) + @mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'true'}) def test_resolve_without_device( self, default_device_str, @@ -886,8 +1001,8 @@ class ResolveImplAndDeviceTest(parameterized.TestCase): if backend == 'cuda': raise RuntimeError('cuda backend not supported') raise AssertionError('jax.devices error') - self.mock_jax_devices.side_effect = devices_side_effect + default_device_side_effect_str = { 'cpu': 'cpu', 'gpu-nvidia': 'gpu', @@ -898,6 +1013,17 @@ class ResolveImplAndDeviceTest(parameterized.TestCase): lambda: default_device_side_effect_str ) + def backends_side_effect(): + if default_device_str == 'gpu-nvidia': + return ['cuda', 'cpu'] + if 'tpu' in default_device_str: + return ['tpu', 'cpu'] + if 'gpu' in default_device_str: + return ['gpu', 'cpu'] + return ['cpu'] + + self.mock_jax_backends.side_effect = backends_side_effect + expected_device, expected_impl = expected if ( expected_impl == 'error' @@ -905,31 +1031,36 @@ class ResolveImplAndDeviceTest(parameterized.TestCase): and impl_str == 'warp' ): with self.assertRaisesRegex(RuntimeError, 'cuda backend not supported'): - mjx_io._resolve_impl_and_device( - impl=impl_str, device=None - ) + mjx_io._resolve_impl_and_device(impl=impl_str, device=None) return if expected_impl == 'error': with self.assertRaises(AssertionError): - mjx_io._resolve_impl_and_device( - impl=impl_str, device=None - ) + mjx_io._resolve_impl_and_device(impl=impl_str, device=None) return - actual_impl, actual_device = ( - mjx_io._resolve_impl_and_device( - impl=impl_str, device=None - ) + if impl_str == 'warp' and not mjxw.WARP_INSTALLED: + with self.assertRaises(RuntimeError): + mjx_io._resolve_impl_and_device(impl=impl_str, device=None) + return + + if impl_str is None and not mjxw.WARP_INSTALLED: + expected_impl = Impl.JAX + + actual_impl, actual_device = mjx_io._resolve_impl_and_device( + impl=impl_str, device=None ) self.assertEqual(actual_impl, expected_impl) self.assertIsNotNone(actual_device) self.assertEqual(actual_device.platform, expected_device) - @mock.patch.dict(os.environ, {'MJX_WARP_ENABLED': 'false'}) + @mock.patch.dict(os.environ, {'MJX_GPU_DEFAULT_WARP': 'false'}) def test_resolve_warp_disabled(self): - """Tests behavior when MJX_WARP_ENABLED is false.""" + """Tests behavior when MJX_GPU_DEFAULT_WARP is false.""" + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp is not installed.') + self.mock_jax_devices.side_effect = lambda backend=None: ( [self.mock_nvidia_gpu, self.mock_cpu] if backend is None @@ -951,64 +1082,6 @@ class ResolveImplAndDeviceTest(parameterized.TestCase): self.assertEqual(impl, Impl.JAX) self.assertEqual(device.platform, 'gpu') - # Requesting warp explicitly should fail since it is disabled. - with self.assertRaises(AssertionError): - mjx_io._resolve_impl_and_device( - impl='warp', device=self.mock_nvidia_gpu - ) - with self.assertRaises(AssertionError): - mjx_io._resolve_impl_and_device(impl='warp', device=None) - - @mock.patch.dict(os.environ, {'MJX_C_DEFAULT_ENABLED': 'false'}) - def test_resolve_c_disabled(self): - """Tests behavior when MJX_C_DEFAULT_ENABLED is false.""" - # Users expect that CPU defaults to the JAX impl. But in the future, it will - # default to the C backend implementation. This test checks that - # MJX_C_DEFAULT_ENABLED=false defaults to the old behavior, until the - # migration to MJEP-15 is complete. - self.mock_jax_devices.side_effect = lambda backend=None: ([self.mock_cpu]) - self.mock_default_backend.side_effect = lambda: 'cpu' - - # Default to JAX instead of C on CPU. - impl, device = mjx_io._resolve_impl_and_device( - impl=None, device=None - ) - self.assertEqual(impl, Impl.JAX) - self.assertEqual(device.platform, 'cpu') - - # Specifing CPU should still choose JAX. - impl, device = mjx_io._resolve_impl_and_device( - impl=None, device=self.mock_cpu - ) - self.assertEqual(impl, Impl.JAX) - self.assertEqual(device.platform, 'cpu') - - # Specifying C should choose C! - impl, device = mjx_io._resolve_impl_and_device( - impl='c', device=None - ) - self.assertEqual(impl, Impl.C) - self.assertEqual(device.platform, 'cpu') - - impl, device = mjx_io._resolve_impl_and_device( - impl='c', device=self.mock_cpu - ) - self.assertEqual(impl, Impl.C) - self.assertEqual(device.platform, 'cpu') - - def test_flex_jax(self): - with self.assertRaises(NotImplementedError): - m = mujoco.MjModel.from_xml_string(""" - - - - - - - - """) - mjx.put_model(m, impl='jax') - if __name__ == '__main__': absltest.main() diff --git a/mjx/mujoco/mjx/_src/passive.py b/mjx/mujoco/mjx/_src/passive.py index b4a8b674..f7acdcd2 100644 --- a/mjx/mujoco/mjx/_src/passive.py +++ b/mjx/mujoco/mjx/_src/passive.py @@ -116,7 +116,11 @@ def _fluid(m: Model, d: Data) -> jax.Array: def passive(m: Model, d: Data) -> Data: """Adds all passive forces.""" - if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + if ( + not isinstance(m._impl, ModelJAX) + or not isinstance(d._impl, DataJAX) + or not isinstance(m.opt._impl, OptionJAX) + ): raise ValueError('passive requires JAX backend implementation.') if m.opt.disableflags & DisableBit.PASSIVE: @@ -130,7 +134,7 @@ def passive(m: Model, d: Data) -> Data: # add gravcomp unless added via actuators qfrc_passive += qfrc_gravcomp * (1 - m.jnt_actgravcomp[m.dof_jntid]) - if m.opt.has_fluid_params: # pytype: disable=attribute-error + if m.opt._impl.has_fluid_params: # pytype: disable=attribute-error qfrc_passive += _fluid(m, d) d = d.replace(qfrc_passive=qfrc_passive, qfrc_gravcomp=qfrc_gravcomp) @@ -147,6 +151,11 @@ def _inertia_box_fluid_model( cvel: jax.Array, ) -> Tuple[jax.Array, jax.Array]: """Fluid forces based on inertia-box approximation.""" + if not isinstance(m.opt._impl, OptionJAX): + raise ValueError( + '_inertia_box_fluid_model requires JAX backend implementation.' + ) + box = jp.repeat(inertia[None, :], 3, axis=0) box *= jp.ones((3, 3)) - 2 * jp.eye(3) box = 6.0 * jp.clip(jp.sum(box, axis=-1), a_min=1e-12) diff --git a/mjx/mujoco/mjx/_src/solver.py b/mjx/mujoco/mjx/_src/solver.py index 90c1add8..f57dbd1c 100644 --- a/mjx/mujoco/mjx/_src/solver.py +++ b/mjx/mujoco/mjx/_src/solver.py @@ -28,6 +28,7 @@ from mujoco.mjx._src.types import DataJAX from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import Model from mujoco.mjx._src.types import ModelJAX +from mujoco.mjx._src.types import OptionJAX from mujoco.mjx._src.types import SolverType # pylint: enable=g-importing-member @@ -75,7 +76,9 @@ class Context(PyTreeNode): @classmethod def create(cls, m: Model, d: Data, grad: bool = True) -> 'Context': - if not isinstance(d._impl, DataJAX): + if not isinstance(d._impl, DataJAX) or not isinstance( + m.opt._impl, OptionJAX + ): raise ValueError( 'Constraint context requires JAX backend implementation.' ) @@ -430,7 +433,11 @@ def _linesearch(m: Model, d: Data, ctx: Context) -> Context: Returns: updated context with new qacc, Ma, Jaref """ - if not isinstance(m._impl, ModelJAX) or not isinstance(d._impl, DataJAX): + if ( + not isinstance(m._impl, ModelJAX) + or not isinstance(d._impl, DataJAX) + or not isinstance(m.opt._impl, OptionJAX) + ): raise ValueError('_lineasearch requires JAX backend implementation.') smag = math.norm(ctx.search) * m.stat.meaninertia * max(1, m.nv) @@ -549,6 +556,8 @@ def _linesearch(m: Model, d: Data, ctx: Context) -> Context: def solve(m: Model, d: Data) -> Data: """Finds forces that satisfy constraints using conjugate gradient descent.""" + if not isinstance(m.opt._impl, OptionJAX): + raise ValueError('solve requires JAX backend implementation.') def cond(ctx: Context) -> jax.Array: improvement = _rescale(m, ctx.prev_cost - ctx.cost) diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index ae21fa2e..a71f934c 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -21,6 +21,7 @@ import warnings import jax import mujoco from mujoco.mjx._src.dataclasses import PyTreeNode # pylint: disable=g-importing-member +from mujoco.mjx.warp import types as mjxw_types import numpy as np @@ -465,48 +466,65 @@ class Statistic(PyTreeNode): center: jax.Array -class Option(PyTreeNode): - """Physics options.""" # fmt: skip - timestep: jax.Array - impratio: jax.Array - tolerance: jax.Array - ls_tolerance: jax.Array - gravity: jax.Array - wind: jax.Array - magnetic: jax.Array - density: jax.Array - viscosity: jax.Array +class StatisticWarp(mjxw_types.StatisticWarp, Statistic): + """Warp-specific model statistics.""" + + # NB: StatisticWarp type annotations may not match those on Statistic. + pass + + +class OptionJAX(PyTreeNode): + """JAX-specific option.""" + o_margin: jax.Array o_solref: jax.Array o_solimp: jax.Array o_friction: jax.Array - integrator: IntegratorType - cone: ConeType - jacobian: JacobianType - solver: SolverType - iterations: int - ls_iterations: int - disableflags: DisableBit - enableflags: int disableactuator: int sdf_initpoints: int + has_fluid_params: bool -class OptionC(Option): +class OptionC(PyTreeNode): """C-specific option.""" + o_margin: jax.Array + o_solref: jax.Array + o_solimp: jax.Array + o_friction: jax.Array + disableactuator: int + sdf_initpoints: int + has_fluid_params: bool apirate: jax.Array noslip_tolerance: jax.Array ccd_tolerance: jax.Array noslip_iterations: int ccd_iterations: int sdf_iterations: int + sdf_initpoints: int -class OptionJAX(Option): - """JAX-specific option.""" +class Option(PyTreeNode): + """Physics options.""" - has_fluid_params: bool + iterations: int + ls_iterations: int + tolerance: jax.Array + ls_tolerance: jax.Array + impratio: jax.Array + gravity: jax.Array + density: jax.Array + viscosity: jax.Array + magnetic: jax.Array + wind: jax.Array + jacobian: JacobianType + cone: ConeType + disableflags: DisableBit + enableflags: int + integrator: IntegratorType + solver: SolverType + timestep: jax.Array + _impl: Union[OptionJAX, OptionC, mjxw_types.OptionWarp] class ModelC(PyTreeNode): @@ -640,7 +658,7 @@ class Model(PyTreeNode): nsensordata: int npluginstate: int opt: Option - stat: Statistic + stat: Union[Statistic, StatisticWarp] qpos0: jax.Array qpos_spring: jax.Array body_parentid: np.ndarray @@ -743,8 +761,8 @@ class Model(PyTreeNode): light_pos: jax.Array light_dir: jax.Array light_poscom0: jax.Array - light_pos0: np.ndarray - light_dir0: np.ndarray + light_pos0: jax.Array + light_dir0: jax.Array light_cutoff: jax.Array mesh_vertadr: np.ndarray mesh_vertnum: np.ndarray @@ -883,13 +901,14 @@ class Model(PyTreeNode): names: bytes signature: np.uint64 _sizes: jax.Array - _impl: Union[ModelC, ModelJAX] + _impl: Union[ModelC, ModelJAX, mjxw_types.ModelWarp] @property def impl(self) -> Impl: return { ModelC: Impl.C, ModelJAX: Impl.JAX, + mjxw_types.ModelWarp: Impl.WARP, }[type(self._impl)] def __getattr__(self, name: str): @@ -1132,13 +1151,14 @@ class Data(PyTreeNode): qacc_smooth: jax.Array qfrc_constraint: jax.Array qfrc_inverse: jax.Array - _impl: Union[DataC, DataJAX] + _impl: Union[DataC, DataJAX, mjxw_types.DataWarp] @property def impl(self) -> Impl: return { DataC: Impl.C, DataJAX: Impl.JAX, + mjxw_types.DataWarp: Impl.WARP, }[type(self._impl)] def __getattr__(self, name: str): @@ -1157,3 +1177,23 @@ class Data(PyTreeNode): f"'{type(self).__name__}' object has no attribute '{name}'" ) return val + + def __getitem__(self, key): + def get_name_from_path(path: jax.tree_util.KeyPath) -> str: + if any(isinstance(p, jax.tree_util.SequenceKey) for p in path): + is_seq_key = [isinstance(p, jax.tree_util.SequenceKey) for p in path] + path = path[: is_seq_key.index(True)] + assert all(isinstance(p, jax.tree_util.GetAttrKey) for p in path) + path = [p for p in path if p.name != '_impl'] + attr = '__'.join(p.name for p in path) + return attr + + if self.impl == Impl.WARP: + return jax.tree.map_with_path( + lambda path, x, k=key: x[k] + if get_name_from_path(path) not in mjxw_types.DATA_NON_VMAP + else x, + self, + ) + + return jax.tree.map(lambda x: x[key], self) diff --git a/mjx/mujoco/mjx/test_data/humanoid/humanoid.xml b/mjx/mujoco/mjx/test_data/humanoid/humanoid.xml index 0851b8ec..168310bc 100644 --- a/mjx/mujoco/mjx/test_data/humanoid/humanoid.xml +++ b/mjx/mujoco/mjx/test_data/humanoid/humanoid.xml @@ -252,5 +252,11 @@ -0.24 -0.007 -0.34 -1.76 -0.466 -0.0415 -0.08 -0.01 -0.37 -0.685 -0.35 -0.09 0.109 -0.067 -0.7 -0.05 0.12 0.16"/> + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py new file mode 100644 index 00000000..1b76ba88 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py @@ -0,0 +1,79 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Public API for MJWarp.""" + +# isort: off +from mujoco.mjx.third_party.mujoco_warp._src.forward import step as step +from mujoco.mjx.third_party.mujoco_warp._src.types import Model as Model +from mujoco.mjx.third_party.mujoco_warp._src.types import Data as Data +# isort: on + +from mujoco.mjx.third_party.mujoco_warp._src.collision_driver import collision as collision +from mujoco.mjx.third_party.mujoco_warp._src.collision_driver import nxn_broadphase as nxn_broadphase +from mujoco.mjx.third_party.mujoco_warp._src.collision_driver import sap_broadphase as sap_broadphase +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import primitive_narrowphase as primitive_narrowphase +from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import sdf_narrowphase as sdf_narrowphase +from mujoco.mjx.third_party.mujoco_warp._src.constraint import make_constraint as make_constraint +from mujoco.mjx.third_party.mujoco_warp._src.derivative import deriv_smooth_vel as deriv_smooth_vel +from mujoco.mjx.third_party.mujoco_warp._src.forward import euler as euler +from mujoco.mjx.third_party.mujoco_warp._src.forward import forward as forward +from mujoco.mjx.third_party.mujoco_warp._src.forward import fwd_acceleration as fwd_acceleration +from mujoco.mjx.third_party.mujoco_warp._src.forward import fwd_actuation as fwd_actuation +from mujoco.mjx.third_party.mujoco_warp._src.forward import fwd_position as fwd_position +from mujoco.mjx.third_party.mujoco_warp._src.forward import fwd_velocity as fwd_velocity +from mujoco.mjx.third_party.mujoco_warp._src.forward import implicit as implicit +from mujoco.mjx.third_party.mujoco_warp._src.forward import rungekutta4 as rungekutta4 +from mujoco.mjx.third_party.mujoco_warp._src.inverse import inverse as inverse +from mujoco.mjx.third_party.mujoco_warp._src.io import get_data_into as get_data_into +from mujoco.mjx.third_party.mujoco_warp._src.io import make_data as make_data +from mujoco.mjx.third_party.mujoco_warp._src.io import put_data as put_data +from mujoco.mjx.third_party.mujoco_warp._src.io import put_model as put_model +from mujoco.mjx.third_party.mujoco_warp._src.passive import passive as passive +from mujoco.mjx.third_party.mujoco_warp._src.ray import ray as ray +from mujoco.mjx.third_party.mujoco_warp._src.sensor import energy_pos as energy_pos +from mujoco.mjx.third_party.mujoco_warp._src.sensor import energy_vel as energy_vel +from mujoco.mjx.third_party.mujoco_warp._src.sensor import sensor_acc as sensor_acc +from mujoco.mjx.third_party.mujoco_warp._src.sensor import sensor_pos as sensor_pos +from mujoco.mjx.third_party.mujoco_warp._src.sensor import sensor_vel as sensor_vel +from mujoco.mjx.third_party.mujoco_warp._src.smooth import camlight as camlight +from mujoco.mjx.third_party.mujoco_warp._src.smooth import com_pos as com_pos +from mujoco.mjx.third_party.mujoco_warp._src.smooth import com_vel as com_vel +from mujoco.mjx.third_party.mujoco_warp._src.smooth import crb as crb +from mujoco.mjx.third_party.mujoco_warp._src.smooth import factor_m as factor_m +from mujoco.mjx.third_party.mujoco_warp._src.smooth import kinematics as kinematics +from mujoco.mjx.third_party.mujoco_warp._src.smooth import rne as rne +from mujoco.mjx.third_party.mujoco_warp._src.smooth import rne_postconstraint as rne_postconstraint +from mujoco.mjx.third_party.mujoco_warp._src.smooth import solve_m as solve_m +from mujoco.mjx.third_party.mujoco_warp._src.smooth import subtree_vel as subtree_vel +from mujoco.mjx.third_party.mujoco_warp._src.smooth import tendon as tendon +from mujoco.mjx.third_party.mujoco_warp._src.smooth import transmission as transmission +from mujoco.mjx.third_party.mujoco_warp._src.solver import solve as solve +from mujoco.mjx.third_party.mujoco_warp._src.support import contact_force as contact_force +from mujoco.mjx.third_party.mujoco_warp._src.support import mul_m as mul_m +from mujoco.mjx.third_party.mujoco_warp._src.support import xfrc_accumulate as xfrc_accumulate +from mujoco.mjx.third_party.mujoco_warp._src.test_util import benchmark as benchmark +from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseType as BroadphaseType +from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType as ConeType +from mujoco.mjx.third_party.mujoco_warp._src.types import Constraint as Constraint +from mujoco.mjx.third_party.mujoco_warp._src.types import Contact as Contact +from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit as DisableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import DynType as DynType +from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit as EnableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import JointType as JointType +from mujoco.mjx.third_party.mujoco_warp._src.types import Option as Option +from mujoco.mjx.third_party.mujoco_warp._src.types import SolverType as SolverType +from mujoco.mjx.third_party.mujoco_warp._src.types import Statistic as Statistic +from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType as TrnType diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/__init__.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/__init__.py new file mode 100644 index 00000000..2b276beb --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/__init__.py @@ -0,0 +1,14 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py new file mode 100644 index 00000000..a080679b --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py @@ -0,0 +1,218 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +from functools import lru_cache + +import warp as wp + + +@lru_cache(maxsize=None) +def create_blocked_cholesky_func(block_size: int): + @wp.func + def blocked_cholesky_func( + # In: + tid_block: int, + A: wp.array(dtype=float, ndim=2), + active_matrix_size: int, + # Out: + L: wp.array(dtype=float, ndim=2), + ): + """ + Computes the Cholesky factorization of a symmetric positive definite matrix A in blocks. + It returns a lower-triangular matrix L such that A = L L^T. + """ + + num_threads_per_block = wp.block_dim() + + # Round up active_matrix_size to next multiple of block_size + n = ((active_matrix_size + block_size - 1) // block_size) * block_size + + # Process the matrix in blocks along its leading dimension. + for k in range(0, n, block_size): + end = k + block_size + + # Load current diagonal block A[k:end, k:end] + # and update with contributions from previously computed blocks. + A_kk_tile = wp.tile_load(A, shape=(block_size, block_size), offset=(k, k), storage="shared") + # The following if pads the matrix if it is not divisible by block_size + if k + block_size > active_matrix_size or k + block_size > active_matrix_size: + num_tile_elements = block_size * block_size + num_iterations = (num_tile_elements + num_threads_per_block - 1) // num_threads_per_block + + for i in range(num_iterations): + linear_index = tid_block + i * num_threads_per_block + linear_index = linear_index % num_tile_elements + row = linear_index // block_size + col = linear_index % block_size + value = A_kk_tile[row, col] + if k + row >= active_matrix_size or k + col >= active_matrix_size: + value = wp.where(row == col, float(1), float(0)) + A_kk_tile[row, col] = value + + if k > 0: + for j in range(0, k, block_size): + L_block = wp.tile_load(L, shape=(block_size, block_size), offset=(k, j)) + L_block_T = wp.tile_transpose(L_block) + L_L_T_block = wp.tile_matmul(L_block, L_block_T) + A_kk_tile -= L_L_T_block + + # Compute the Cholesky factorization for the block + L_kk_tile = wp.tile_cholesky(A_kk_tile) + wp.tile_store(L, L_kk_tile, offset=(k, k)) + + # Process the blocks below the current block + for i in range(end, n, block_size): + A_ik_tile = wp.tile_load(A, shape=(block_size, block_size), offset=(i, k), storage="shared") + # The following if pads the matrix if it is not divisible by block_size + if i + block_size > active_matrix_size or k + block_size > active_matrix_size: + num_tile_elements = block_size * block_size + num_iterations = (num_tile_elements + num_threads_per_block - 1) // num_threads_per_block + + for ii in range(num_iterations): + linear_index = tid_block + ii * num_threads_per_block + linear_index = linear_index % num_tile_elements + row = linear_index // block_size + col = linear_index % block_size + value = A_ik_tile[row, col] + if i + row >= active_matrix_size or k + col >= active_matrix_size: + value = wp.where(i + row == k + col, float(1), float(0)) + A_ik_tile[row, col] = value + + if k > 0: + for j in range(0, k, block_size): + L_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(i, j)) + L_2_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(k, j)) + L_T_tile = wp.tile_transpose(L_2_tile) + L_L_T_tile = wp.tile_matmul(L_tile, L_T_tile) + A_ik_tile -= L_L_T_tile + + t = wp.tile_transpose(A_ik_tile) + tmp = wp.tile_lower_solve(L_kk_tile, t) + sol_tile = wp.tile_transpose(tmp) + + wp.tile_store(L, sol_tile, offset=(i, k)) + + return blocked_cholesky_func + + +@lru_cache(maxsize=None) +def create_blocked_cholesky_solve_func(block_size: int): + @wp.func + def blocked_cholesky_solve_func( + # In: + tid_block: int, + L: wp.array(dtype=float, ndim=2), + b: wp.array(dtype=float, ndim=2), + tmp: wp.array(dtype=float, ndim=2), + active_matrix_size: int, + # Out: + x: wp.array(dtype=float, ndim=2), + ): + """ + Solves A x = b given the Cholesky factor L (A = L L^T) using + blocked forward and backward substitution. + """ + + num_threads_per_block = wp.block_dim() + + # Round up active_matrix_size to next multiple of block_size + n = ((active_matrix_size + block_size - 1) // block_size) * block_size + + # Forward substitution: solve L y = b + for i in range(0, n, block_size): + i_end = i + block_size + rhs_tile = wp.tile_load(b, shape=(block_size, 1), offset=(i, 0)) + if i > 0: + for j in range(0, i, block_size): + L_block = wp.tile_load(L, shape=(block_size, block_size), offset=(i, j)) + y_block = wp.tile_load(tmp, shape=(block_size, 1), offset=(j, 0)) + Ly_block = wp.tile_matmul(L_block, y_block) + rhs_tile -= Ly_block + L_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(i, i)) + + # The following if pads the matrix if it is not divisible by block_size + if i + block_size > active_matrix_size: + num_tile_elements = block_size * block_size + num_iterations = (num_tile_elements + num_threads_per_block - 1) // num_threads_per_block + + for ii in range(num_iterations): + linear_index = tid_block + ii * num_threads_per_block + linear_index = linear_index % num_tile_elements + row = linear_index // block_size + col = linear_index % block_size + value = L_tile[row, col] + if i + row >= active_matrix_size or i + col >= active_matrix_size: + value = wp.where(row == col, float(1), float(0)) + L_tile[row, col] = value + + # Handle rhs + num_tile_elements = block_size + num_iterations = (num_tile_elements + num_threads_per_block - 1) // num_threads_per_block + + for ii in range(num_iterations): + linear_index = tid_block + ii * num_threads_per_block + linear_index = linear_index % num_tile_elements + value = rhs_tile[linear_index, 0] + if i + linear_index >= active_matrix_size: + value = float(0) + rhs_tile[linear_index, 0] = value + + y_tile = wp.tile_lower_solve(L_tile, rhs_tile) + wp.tile_store(tmp, y_tile, offset=(i, 0)) + + # Backward substitution: solve L^T x = y + for i in range(n - block_size, -1, -block_size): + i_end = i + block_size + rhs_tile = wp.tile_load(tmp, shape=(block_size, 1), offset=(i, 0)) + if i_end < n: + for j in range(i_end, n, block_size): + L_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(j, i)) + L_T_tile = wp.tile_transpose(L_tile) + x_tile = wp.tile_load(x, shape=(block_size, 1), offset=(j, 0)) + L_T_x_tile = wp.tile_matmul(L_T_tile, x_tile) + rhs_tile -= L_T_x_tile + L_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(i, i)) + + # The following if pads the matrix if it is not divisible by block_size + if i + block_size > active_matrix_size: + num_tile_elements = block_size * block_size + num_iterations = (num_tile_elements + num_threads_per_block - 1) // num_threads_per_block + + for ii in range(num_iterations): + linear_index = tid_block + ii * num_threads_per_block + linear_index = linear_index % num_tile_elements + row = linear_index // block_size + col = linear_index % block_size + value = L_tile[row, col] + if i + row >= active_matrix_size or i + col >= active_matrix_size: + value = wp.where(row == col, float(1), float(0)) + L_tile[row, col] = value + + # Handle rhs + num_tile_elements = block_size + num_iterations = (num_tile_elements + num_threads_per_block - 1) // num_threads_per_block + + for ii in range(num_iterations): + linear_index = tid_block + ii * num_threads_per_block + linear_index = linear_index % num_tile_elements + value = rhs_tile[linear_index, 0] + if i + linear_index >= active_matrix_size: + value = float(0) + rhs_tile[linear_index, 0] = value + + x_tile = wp.tile_upper_solve(wp.tile_transpose(L_tile), rhs_tile) + wp.tile_store(x, x_tile, offset=(i, 0)) + + return blocked_cholesky_solve_func diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/broadphase_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/broadphase_test.py new file mode 100644 index 00000000..57a7b439 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/broadphase_test.py @@ -0,0 +1,335 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Tests for broadphase functions.""" + +import numpy as np +import warp as wp +from absl.testing import absltest +from absl.testing import parameterized + +import mujoco_warp as mjwarp + +from mujoco.mjx.third_party.mujoco_warp._src import collision_driver +from mujoco.mjx.third_party.mujoco_warp._src import test_util +from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseFilter +from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseType + + +def broadphase_caller(m, d): + if m.opt.broadphase == int(BroadphaseType.NXN): + collision_driver.nxn_broadphase(m, d) + else: + collision_driver.sap_broadphase(m, d) + + +class BroadphaseTest(parameterized.TestCase): + # filter combinations + plane_sphere = BroadphaseFilter.PLANE.value | BroadphaseFilter.SPHERE.value + plane_aabb = BroadphaseFilter.PLANE.value | BroadphaseFilter.AABB.value + plane_obb = BroadphaseFilter.PLANE | BroadphaseFilter.OBB.value + plane_sphere_aabb = plane_sphere | BroadphaseFilter.AABB.value + plane_sphere_obb = plane_sphere | BroadphaseFilter.OBB.value + plane_sphere_aabb_obb = plane_sphere_aabb | BroadphaseFilter.OBB.value + + @parameterized.product( + broadphase=list(BroadphaseType), + filter=[plane_sphere, plane_aabb, plane_obb, plane_sphere_aabb, plane_sphere_obb, plane_sphere_aabb_obb], + ) + def test_broadphase(self, broadphase, filter): + """Tests collision broadphase algorithms.""" + + _XML = """ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + """ + + # one world and zero collisions + mjm, _, m, d0 = test_util.fixture(xml=_XML, keyframe=0) + + m.opt.broadphase = broadphase + m.opt.broadphase_filter = filter + + broadphase_caller(m, d0) + np.testing.assert_allclose(d0.ncollision.numpy()[0], 0) + + # one world and one collision + _, mjd1, _, d1 = test_util.fixture(xml=_XML, keyframe=1) + broadphase_caller(m, d1) + + np.testing.assert_allclose(d1.ncollision.numpy()[0], 1) + np.testing.assert_allclose(d1.collision_pair.numpy()[0][0], 0) + np.testing.assert_allclose(d1.collision_pair.numpy()[0][1], 1) + + # one world and three collisions + _, mjd2, _, d2 = test_util.fixture(xml=_XML, keyframe=2) + broadphase_caller(m, d2) + + ncollision = d2.ncollision.numpy()[0] + np.testing.assert_allclose(ncollision, 3) + + collision_pairs = [[0, 1], [0, 2], [1, 2]] + for i in range(ncollision): + self.assertTrue([d2.collision_pair.numpy()[i][0], d2.collision_pair.numpy()[i][1]] in collision_pairs) + + # two worlds and four collisions + d3 = mjwarp.make_data(mjm, nworld=2, nconmax=512, njmax=512) + d3.geom_xpos = wp.array( + np.vstack([np.expand_dims(mjd1.geom_xpos, axis=0), np.expand_dims(mjd2.geom_xpos, axis=0)]), + dtype=wp.vec3, + ) + d3.geom_xmat = wp.array( + np.vstack([np.expand_dims(mjd1.geom_xmat, axis=0), np.expand_dims(mjd2.geom_xmat, axis=0)]), + dtype=wp.mat33, + ) + broadphase_caller(m, d3) + + ncollision = d3.ncollision.numpy()[0] + np.testing.assert_allclose(ncollision, 4) + + collision_pairs = [[[0, 1]], [[0, 1], [0, 2], [1, 2]]] + worldids = [0, 1, 1, 1] + for i in range(ncollision): + worldid = d3.collision_worldid.numpy()[i] + self.assertTrue(worldid == worldids[i]) + self.assertTrue([d3.collision_pair.numpy()[i][0], d3.collision_pair.numpy()[i][1]] in collision_pairs[worldid]) + + # one world and zero collisions: contype and conaffinity incompatibility + mjm4, _, m4, d4 = test_util.fixture(xml=_XML, keyframe=1) + mjm4.geom_contype[:3] = 0 + m4 = mjwarp.put_model(mjm4) + + broadphase_caller(m4, d4) + np.testing.assert_allclose(d4.ncollision.numpy()[0], 0) + + # one world and one collision: geomtype ordering + _, _, _, d5 = test_util.fixture(xml=_XML, keyframe=3) + broadphase_caller(m, d5) + np.testing.assert_allclose(d5.ncollision.numpy()[0], 1) + np.testing.assert_allclose(d5.collision_pair.numpy()[0][0], 3) + np.testing.assert_allclose(d5.collision_pair.numpy()[0][1], 2) + + @parameterized.parameters((0, 0, 0), (0, 0.011, 1), (0.011, 0, 1), (0.00999, 0, 0), (0, 0.00999, 0), (0.00999, 0.00999, 0)) + def test_broadphase_margin(self, margin1, margin2, ncollision): + _MJCF = f""" + + + + + + + + + + + + + + + + """ + _, _, m, d = test_util.fixture(xml=_MJCF, keyframe=0) + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], ncollision) + + @parameterized.parameters(True, False) + def test_broadphase_filterparent(self, filterparent): + _MJCF = """ + + + + + + + + + + + + + + + + """ + _, _, m, d = test_util.fixture(xml=_MJCF, filterparent=filterparent, keyframe=0) + + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 0 if filterparent else 1) + + def test_broadphase_filter(self): + plane = BroadphaseFilter.PLANE.value + sphere = BroadphaseFilter.SPHERE.value + aabb = BroadphaseFilter.AABB.value + obb = BroadphaseFilter.OBB.value + plane_sphere = plane | sphere + plane_aabb = plane | aabb + plane_obb = plane | obb + + _PLANE_CAPSULE_CAPSULE = """ + + + """ + + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=0) + m.opt.broadphase_filter = plane_sphere + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 0) + + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=0) + m.opt.broadphase_filter = plane_aabb + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 0) + + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=0) + m.opt.broadphase_filter = plane_obb + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 0) + + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=1) + m.opt.broadphase_filter = plane_sphere + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 1) + + # note: collision_driver._plane_filter checks bounding sphere + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=1) + m.opt.broadphase_filter = plane + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 2) + + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=1) + m.opt.broadphase_filter = plane_sphere + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 1) + + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=1) + m.opt.broadphase_filter = plane_obb + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 1) + + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=2) + m.opt.broadphase_filter = plane_sphere + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 1) + + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=2) + m.opt.broadphase_filter = plane_aabb + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 0) + + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=2) + m.opt.broadphase_filter = plane_obb + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 0) + + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=3) + m.opt.broadphase_filter = plane_sphere + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 1) + + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=3) + m.opt.broadphase_filter = plane_aabb + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 1) + + _, _, m, d = test_util.fixture(xml=_PLANE_CAPSULE_CAPSULE, keyframe=3) + m.opt.broadphase_filter = plane_obb + broadphase_caller(m, d) + self.assertEqual(d.ncollision.numpy()[0], 0) + + # TODO(team): test margin + # TODO(team): test DisableBit.FILTERPARENT + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py new file mode 100644 index 00000000..8ddb801e --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py @@ -0,0 +1,495 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import ccd +from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk_legacy import epa_legacy +from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk_legacy import gjk_legacy +from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk_legacy import multicontact_legacy +from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_prism_vertex +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import _geom +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import contact_params +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import write_contact +from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame +from mujoco.mjx.third_party.mujoco_warp._src.math import upper_tri_index +from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXCONPAIR +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel + +# TODO(team): improve compile time to enable backward pass +wp.config.enable_backward = False + +MULTI_CONTACT_COUNT = 4 +mat3c = wp.types.matrix(shape=(MULTI_CONTACT_COUNT, 3), dtype=float) + +_CONVEX_COLLISION_PAIRS = [ + (GeomType.HFIELD.value, GeomType.SPHERE.value), + (GeomType.HFIELD.value, GeomType.CAPSULE.value), + (GeomType.HFIELD.value, GeomType.ELLIPSOID.value), + (GeomType.HFIELD.value, GeomType.CYLINDER.value), + (GeomType.HFIELD.value, GeomType.BOX.value), + (GeomType.HFIELD.value, GeomType.MESH.value), + (GeomType.SPHERE.value, GeomType.ELLIPSOID.value), + (GeomType.SPHERE.value, GeomType.MESH.value), + (GeomType.CAPSULE.value, GeomType.ELLIPSOID.value), + (GeomType.CAPSULE.value, GeomType.CYLINDER.value), + (GeomType.CAPSULE.value, GeomType.MESH.value), + (GeomType.ELLIPSOID.value, GeomType.ELLIPSOID.value), + (GeomType.ELLIPSOID.value, GeomType.CYLINDER.value), + (GeomType.ELLIPSOID.value, GeomType.BOX.value), + (GeomType.ELLIPSOID.value, GeomType.MESH.value), + (GeomType.CYLINDER.value, GeomType.CYLINDER.value), + (GeomType.CYLINDER.value, GeomType.BOX.value), + (GeomType.CYLINDER.value, GeomType.MESH.value), + (GeomType.BOX.value, GeomType.MESH.value), + (GeomType.MESH.value, GeomType.MESH.value), +] + + +def _check_convex_collision_pairs(): + prev_idx = -1 + for pair in _CONVEX_COLLISION_PAIRS: + idx = upper_trid_index(len(GeomType), pair[0], pair[1]) + if pair[1] < pair[0] or idx <= prev_idx: + return False + prev_idx = idx + return True + + +assert _check_convex_collision_pairs(), "_CONVEX_COLLISION_PAIRS is in invalid order." + + +@wp.func +def _max_contacts_height_field( + # Model: + ngeom: int, + geom_type: wp.array(dtype=int), + geompair2hfgeompair: wp.array(dtype=int), + # In: + g1: int, + g2: int, + worldid: int, + # Data out: + ncon_hfield_out: wp.array2d(dtype=int), +): + hfield = int(GeomType.HFIELD.value) + if geom_type[g1] == hfield or (geom_type[g2] == hfield): + geompairid = upper_tri_index(ngeom, g1, g2) + hfgeompairid = geompair2hfgeompair[geompairid] + hfncon = wp.atomic_add(ncon_hfield_out[worldid], hfgeompairid, 1) + if hfncon >= MJ_MAXCONPAIR: + return True + + return False + + +@cache_kernel +def ccd_kernel_builder( + default_gjk: bool, + geomtype1: int, + geomtype2: int, + gjk_iterations: int, + epa_iterations: int, + epa_exact_neg_distance: bool, + depth_extension: float, +): + # runs convex collision on a set of geom pairs to recover contact info + @nested_kernel + def ccd_kernel( + # Model: + ngeom: int, + geom_type: wp.array(dtype=int), + geom_condim: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_priority: wp.array(dtype=int), + geom_solmix: wp.array2d(dtype=float), + geom_solref: wp.array2d(dtype=wp.vec2), + geom_solimp: wp.array2d(dtype=vec5), + geom_size: wp.array2d(dtype=wp.vec3), + geom_friction: wp.array2d(dtype=wp.vec3), + geom_margin: wp.array2d(dtype=float), + geom_gap: wp.array2d(dtype=float), + hfield_adr: wp.array(dtype=int), + hfield_nrow: wp.array(dtype=int), + hfield_ncol: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + hfield_data: wp.array(dtype=float), + mesh_vertadr: wp.array(dtype=int), + mesh_vertnum: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_graphadr: wp.array(dtype=int), + mesh_graph: wp.array(dtype=int), + mesh_polynum: wp.array(dtype=int), + mesh_polyadr: wp.array(dtype=int), + mesh_polynormal: wp.array(dtype=wp.vec3), + mesh_polyvertadr: wp.array(dtype=int), + mesh_polyvertnum: wp.array(dtype=int), + mesh_polyvert: wp.array(dtype=int), + mesh_polymapadr: wp.array(dtype=int), + mesh_polymapnum: wp.array(dtype=int), + mesh_polymap: wp.array(dtype=int), + pair_dim: wp.array(dtype=int), + pair_solref: wp.array2d(dtype=wp.vec2), + pair_solreffriction: wp.array2d(dtype=wp.vec2), + pair_solimp: wp.array2d(dtype=vec5), + pair_margin: wp.array2d(dtype=float), + pair_gap: wp.array2d(dtype=float), + pair_friction: wp.array2d(dtype=vec5), + geompair2hfgeompair: wp.array(dtype=int), + # Data in: + nconmax_in: int, + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + collision_pair_in: wp.array(dtype=wp.vec2i), + collision_hftri_index_in: wp.array(dtype=int), + collision_pairid_in: wp.array(dtype=int), + collision_worldid_in: wp.array(dtype=int), + ncollision_in: wp.array(dtype=int), + epa_vert_in: wp.array2d(dtype=wp.vec3), + epa_vert1_in: wp.array2d(dtype=wp.vec3), + epa_vert2_in: wp.array2d(dtype=wp.vec3), + epa_vert_index1_in: wp.array2d(dtype=int), + epa_vert_index2_in: wp.array2d(dtype=int), + epa_face_in: wp.array2d(dtype=wp.vec3i), + epa_pr_in: wp.array2d(dtype=wp.vec3), + epa_norm2_in: wp.array2d(dtype=float), + epa_index_in: wp.array2d(dtype=int), + epa_map_in: wp.array2d(dtype=int), + epa_horizon_in: wp.array2d(dtype=int), + # Data out: + ncon_out: wp.array(dtype=int), + ncon_hfield_out: wp.array2d(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), + ): + tid = wp.tid() + if tid >= ncollision_in[0]: + return + + geoms = collision_pair_in[tid] + g1 = geoms[0] + g2 = geoms[1] + + if geom_type[g1] != geomtype1 or geom_type[g2] != geomtype2: + return + + worldid = collision_worldid_in[tid] + _, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params( + geom_condim, + geom_priority, + geom_solmix, + geom_solref, + geom_solimp, + geom_friction, + geom_margin, + geom_gap, + pair_dim, + pair_solref, + pair_solreffriction, + pair_solimp, + pair_margin, + pair_gap, + pair_friction, + collision_pair_in, + collision_pairid_in, + tid, + worldid, + ) + + hftri_index = collision_hftri_index_in[tid] + + geom1 = _geom( + geom_type, + geom_dataid, + geom_size, + hfield_adr, + hfield_nrow, + hfield_ncol, + hfield_size, + hfield_data, + mesh_vertadr, + mesh_vertnum, + mesh_vert, + mesh_graphadr, + mesh_graph, + mesh_polynum, + mesh_polyadr, + mesh_polynormal, + mesh_polyvertadr, + mesh_polyvertnum, + mesh_polyvert, + mesh_polymapadr, + mesh_polymapnum, + mesh_polymap, + geom_xpos_in, + geom_xmat_in, + worldid, + g1, + hftri_index, + ) + + geom2 = _geom( + geom_type, + geom_dataid, + geom_size, + hfield_adr, + hfield_nrow, + hfield_ncol, + hfield_size, + hfield_data, + mesh_vertadr, + mesh_vertnum, + mesh_vert, + mesh_graphadr, + mesh_graph, + mesh_polynum, + mesh_polyadr, + mesh_polynormal, + mesh_polyvertadr, + mesh_polyvertnum, + mesh_polyvert, + mesh_polymapadr, + mesh_polymapnum, + mesh_polymap, + geom_xpos_in, + geom_xmat_in, + worldid, + g2, + hftri_index, + ) + + points = mat3c() + + if default_gjk: + simplex, normal = gjk_legacy( + gjk_iterations, + geom1, + geom2, + geomtype1, + geomtype2, + ) + + depth, normal = epa_legacy( + epa_iterations, geom1, geom2, geomtype1, geomtype2, depth_extension, epa_exact_neg_distance, simplex, normal + ) + dist = -depth + + if (dist - margin) >= 0.0 or depth != depth: + return + sphere = int(GeomType.SPHERE.value) + ellipsoid = int(GeomType.ELLIPSOID.value) + if geom_type[g1] == sphere or geom_type[g1] == ellipsoid or geom_type[g2] == sphere or geom_type[g2] == ellipsoid: + count, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 1, 2, 1.0e-5) + else: + count, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 4, 8, 1.0e-3) + else: + x1 = geom1.pos + x2 = geom2.pos + + # find prism center for height field + if geomtype1 == int(GeomType.HFIELD.value): + x1 = wp.vec3(0.0, 0.0, 0.0) + for i in range(6): + x1 += hfield_prism_vertex(geom1.hfprism, i) + x1 = x1 / 6.0 + + dist, x1, x2 = ccd( + 1e-6, + 0.0, + gjk_iterations, + epa_iterations, + geom1, + geom2, + geomtype1, + geomtype2, + x1, + x2, + epa_vert_in[tid], + epa_vert1_in[tid], + epa_vert2_in[tid], + epa_vert_index1_in[tid], + epa_vert_index2_in[tid], + epa_face_in[tid], + epa_pr_in[tid], + epa_norm2_in[tid], + epa_index_in[tid], + epa_map_in[tid], + epa_horizon_in[tid], + ) + count = 0 + if dist < 0.0: + count = 1 + + points[0] = 0.5 * (x1 + x2) + normal = x1 - x2 + + frame = make_frame(normal) + for i in range(count): + # limit maximum number of contacts with height field + if _max_contacts_height_field(ngeom, geom_type, geompair2hfgeompair, g1, g2, worldid, ncon_hfield_out): + return + + write_contact( + nconmax_in, + dist, + points[i], + frame, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + return ccd_kernel + + +@event_scope +def convex_narrowphase(m: Model, d: Data): + """Runs narrowphase collision detection for convex geom pairs. + + This function handles collision detection for pairs of convex geometries that were + identified during the broadphase. It uses the Gilbert-Johnson-Keerthi (GJK) algorithm to + determine the distance between shapes and the Expanding Polytope Algorithm (EPA) to find + the penetration depth and contact normal for colliding pairs. + + The convex geom types handled by this function are SPHERE, CAPSULE, ELLIPSOID, CYLINDER, + BOX, MESH, HFIELD. + + To optimize performance, this function dynamically builds and launches a specialized + kernel for each type of convex collision pair present in the model, avoiding unnecessary + computations for non-existent pair types. + """ + for geom_pair in _CONVEX_COLLISION_PAIRS: + if m.geom_pair_type_count[upper_trid_index(len(GeomType), geom_pair[0], geom_pair[1])]: + wp.launch( + ccd_kernel_builder( + False, + geom_pair[0], + geom_pair[1], + m.opt.gjk_iterations, + m.opt.epa_iterations, + False, + 0.1, + ), + dim=d.nconmax, + inputs=[ + m.ngeom, + m.geom_type, + m.geom_condim, + m.geom_dataid, + m.geom_priority, + m.geom_solmix, + m.geom_solref, + m.geom_solimp, + m.geom_size, + m.geom_friction, + m.geom_margin, + m.geom_gap, + m.hfield_adr, + m.hfield_nrow, + m.hfield_ncol, + m.hfield_size, + m.hfield_data, + m.mesh_vertadr, + m.mesh_vertnum, + m.mesh_vert, + m.mesh_graphadr, + m.mesh_graph, + m.mesh_polynum, + m.mesh_polyadr, + m.mesh_polynormal, + m.mesh_polyvertadr, + m.mesh_polyvertnum, + m.mesh_polyvert, + m.mesh_polymapadr, + m.mesh_polymapnum, + m.mesh_polymap, + m.pair_dim, + m.pair_solref, + m.pair_solreffriction, + m.pair_solimp, + m.pair_margin, + m.pair_gap, + m.pair_friction, + m.geompair2hfgeompair, + d.nconmax, + d.geom_xpos, + d.geom_xmat, + d.collision_pair, + d.collision_hftri_index, + d.collision_pairid, + d.collision_worldid, + d.ncollision, + d.epa_vert, + d.epa_vert1, + d.epa_vert2, + d.epa_vert_index1, + d.epa_vert_index2, + d.epa_face, + d.epa_pr, + d.epa_norm2, + d.epa_index, + d.epa_map, + d.epa_horizon, + ], + outputs=[ + d.ncon, + d.ncon_hfield, + d.contact.dist, + d.contact.pos, + d.contact.frame, + d.contact.includemargin, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + d.contact.dim, + d.contact.geom, + d.contact.worldid, + ], + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py new file mode 100644 index 00000000..4cdd85c2 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py @@ -0,0 +1,767 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +from typing import Any + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src.collision_convex import convex_narrowphase +from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_midphase +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import primitive_narrowphase +from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import sdf_narrowphase +from mujoco.mjx.third_party.mujoco_warp._src.math import upper_tri_index +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseFilter +from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseType +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope + +wp.set_module_options({"enable_backward": False}) + + +@wp.kernel +def _zero_collision_arrays( + # Data in: + nworld_in: int, + # In: + hfield_geom_pair_in: int, + # Data out: + ncon_out: wp.array(dtype=int), + ncon_hfield_out: wp.array(dtype=int), # kernel_analyzer: ignore + collision_hftri_index_out: wp.array(dtype=int), + ncollision_out: wp.array(dtype=int), +): + tid = wp.tid() + + if tid == 0: + # Zero the single collision counter + ncollision_out[0] = 0 + ncon_out[0] = 0 + + if tid < hfield_geom_pair_in * nworld_in: + ncon_hfield_out[tid] = 0 + + # Zero collision pair indices + collision_hftri_index_out[tid] = 0 + + +@wp.func +def _plane_filter( + size1: float, size2: float, margin1: float, margin2: float, xpos1: wp.vec3, xpos2: wp.vec3, xmat1: wp.mat33, xmat2: wp.mat33 +) -> bool: + if size1 == 0.0: + # geom1 is a plane + dist = wp.dot(xpos2 - xpos1, wp.vec3(xmat1[0, 2], xmat1[1, 2], xmat1[2, 2])) + return dist <= size2 + wp.max(margin1, margin2) + elif size2 == 0.0: + # geom2 is a plane + dist = wp.dot(xpos1 - xpos2, wp.vec3(xmat2[0, 2], xmat2[1, 2], xmat2[2, 2])) + return dist <= size1 + wp.max(margin1, margin2) + + return True + + +@wp.func +def _sphere_filter(size1: float, size2: float, margin1: float, margin2: float, xpos1: wp.vec3, xpos2: wp.vec3) -> bool: + bound = size1 + size2 + wp.max(margin1, margin2) + dif = xpos2 - xpos1 + dist_sq = wp.dot(dif, dif) + return dist_sq <= bound * bound + + +# TODO(team): improve performance by precomputing bounding box +@wp.func +def _aabb_filter( + # In: + center1: wp.vec3, + center2: wp.vec3, + size1: wp.vec3, + size2: wp.vec3, + margin1: float, + margin2: float, + xpos1: wp.vec3, + xpos2: wp.vec3, + xmat1: wp.mat33, + xmat2: wp.mat33, +) -> bool: + """Axis aligned boxes collision. + + references: see Ericson, Real-time Collision Detection section 4.2. + filterBox: filter contact based on global AABBs. + """ + center1 = xmat1 @ center1 + xpos1 + center2 = xmat2 @ center2 + xpos2 + + margin = wp.max(margin1, margin2) + + max_x1 = -MJ_MAXVAL + max_y1 = -MJ_MAXVAL + max_z1 = -MJ_MAXVAL + min_x1 = MJ_MAXVAL + min_y1 = MJ_MAXVAL + min_z1 = MJ_MAXVAL + + max_x2 = -MJ_MAXVAL + max_y2 = -MJ_MAXVAL + max_z2 = -MJ_MAXVAL + min_x2 = MJ_MAXVAL + min_y2 = MJ_MAXVAL + min_z2 = MJ_MAXVAL + + sign = wp.vec2(-1.0, 1.0) + + for i in range(2): + for j in range(2): + for k in range(2): + corner1 = wp.vec3(sign[i] * size1[0], sign[j] * size1[1], sign[k] * size1[2]) + pos1 = xmat1 @ corner1 + + corner2 = wp.vec3(sign[i] * size2[0], sign[j] * size2[1], sign[k] * size2[2]) + pos2 = xmat2 @ corner2 + + if pos1[0] > max_x1: + max_x1 = pos1[0] + + if pos1[1] > max_y1: + max_y1 = pos1[1] + + if pos1[2] > max_z1: + max_z1 = pos1[2] + + if pos1[0] < min_x1: + min_x1 = pos1[0] + + if pos1[1] < min_y1: + min_y1 = pos1[1] + + if pos1[2] < min_z1: + min_z1 = pos1[2] + + if pos2[0] > max_x2: + max_x2 = pos2[0] + + if pos2[1] > max_y2: + max_y2 = pos2[1] + + if pos2[2] > max_z2: + max_z2 = pos2[2] + + if pos2[0] < min_x2: + min_x2 = pos2[0] + + if pos2[1] < min_y2: + min_y2 = pos2[1] + + if pos2[2] < min_z2: + min_z2 = pos2[2] + + if center1[0] + max_x1 + margin < center2[0] + min_x2: + return False + if center1[1] + max_y1 + margin < center2[1] + min_y2: + return False + if center1[2] + max_z1 + margin < center2[2] + min_z2: + return False + if center2[0] + max_x2 + margin < center1[0] + min_x1: + return False + if center2[1] + max_y2 + margin < center1[1] + min_y1: + return False + if center2[2] + max_z2 + margin < center1[2] + min_z1: + return False + + return True + + +mat23 = wp.types.matrix(shape=(2, 3), dtype=float) +mat63 = wp.types.matrix(shape=(6, 3), dtype=float) + + +# TODO(team): improve performance by precomputing bounding box +@wp.func +def _obb_filter( + # In: + center1: wp.vec3, + center2: wp.vec3, + size1: wp.vec3, + size2: wp.vec3, + margin1: float, + margin2: float, + xpos1: wp.vec3, + xpos2: wp.vec3, + xmat1: wp.mat33, + xmat2: wp.mat33, +) -> bool: + """Oriented bounding boxes collision (see Gottschalk et al.), see mj_collideOBB.""" + margin = wp.max(margin1, margin2) + + xcenter = mat23() + normal = mat63() + proj = wp.vec2() + radius = wp.vec2() + + # compute centers in local coordinates + xcenter[0] = xmat1 @ center1 + xpos1 + xcenter[1] = xmat2 @ center2 + xpos2 + + # compute normals in global coordinates + normal[0] = wp.vec3(xmat1[0, 0], xmat1[1, 0], xmat1[2, 0]) + normal[1] = wp.vec3(xmat1[0, 1], xmat1[1, 1], xmat1[2, 1]) + normal[2] = wp.vec3(xmat1[0, 2], xmat1[1, 2], xmat1[2, 2]) + normal[3] = wp.vec3(xmat2[0, 0], xmat2[1, 0], xmat2[2, 0]) + normal[4] = wp.vec3(xmat2[0, 1], xmat2[1, 1], xmat2[2, 1]) + normal[5] = wp.vec3(xmat2[0, 2], xmat2[1, 2], xmat2[2, 2]) + + # check intersections + for j in range(2): + for k in range(3): + for i in range(2): + proj[i] = wp.dot(xcenter[i], normal[3 * j + k]) + if i == 0: + size = size1 + else: + size = size2 + + # fmt: off + radius[i] = ( + wp.abs(size[0] * wp.dot(normal[3 * i + 0], normal[3 * j + k])) + + wp.abs(size[1] * wp.dot(normal[3 * i + 1], normal[3 * j + k])) + + wp.abs(size[2] * wp.dot(normal[3 * i + 2], normal[3 * j + k])) + ) + # fmt: on + if radius[0] + radius[1] + margin < wp.abs(proj[1] - proj[0]): + return False + + return True + + +@wp.func +def _broadphase_filter( + # Model: + opt_broadphase_filter: int, + geom_aabb: wp.array2d(dtype=wp.vec3), + geom_rbound: wp.array2d(dtype=float), + geom_margin: wp.array2d(dtype=float), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + geom1: int, + geom2: int, + worldid: int, +) -> bool: + # 1: plane + # 2: sphere + # 4: aabb + # 8: obb + + center1 = geom_aabb[geom1, 0] + center2 = geom_aabb[geom2, 0] + size1 = geom_aabb[geom1, 1] + size2 = geom_aabb[geom2, 1] + rbound1 = geom_rbound[worldid, geom1] + rbound2 = geom_rbound[worldid, geom2] + margin1 = geom_margin[worldid, geom1] + margin2 = geom_margin[worldid, geom2] + xpos1 = geom_xpos_in[worldid, geom1] + xpos2 = geom_xpos_in[worldid, geom2] + xmat1 = geom_xmat_in[worldid, geom1] + xmat2 = geom_xmat_in[worldid, geom2] + + if rbound1 == 0.0 or rbound2 == 0.0: + if opt_broadphase_filter & int(BroadphaseFilter.PLANE.value): + return _plane_filter(rbound1, rbound2, margin1, margin2, xpos1, xpos2, xmat1, xmat2) + else: + if opt_broadphase_filter & int(BroadphaseFilter.SPHERE.value): + if not _sphere_filter(rbound1, rbound2, margin1, margin2, xpos1, xpos2): + return False + if opt_broadphase_filter & int(BroadphaseFilter.AABB.value): + if not _aabb_filter(center1, center2, size1, size2, margin1, margin2, xpos1, xpos2, xmat1, xmat2): + return False + if opt_broadphase_filter & int(BroadphaseFilter.OBB.value): + if not _obb_filter(center1, center2, size1, size2, margin1, margin2, xpos1, xpos2, xmat1, xmat2): + return False + + return True + + +@wp.func +def _add_geom_pair( + # Model: + geom_type: wp.array(dtype=int), + nxn_pairid: wp.array(dtype=int), + # Data in: + nconmax_in: int, + # In: + geom1: int, + geom2: int, + worldid: int, + nxnid: int, + # Data out: + collision_pair_out: wp.array(dtype=wp.vec2i), + collision_hftri_index_out: wp.array(dtype=int), + collision_pairid_out: wp.array(dtype=int), + collision_worldid_out: wp.array(dtype=int), + ncollision_out: wp.array(dtype=int), +): + pairid = wp.atomic_add(ncollision_out, 0, 1) + + if pairid >= nconmax_in: + return + + type1 = geom_type[geom1] + type2 = geom_type[geom2] + + if type1 > type2: + pair = wp.vec2i(geom2, geom1) + else: + pair = wp.vec2i(geom1, geom2) + + collision_pair_out[pairid] = pair + collision_pairid_out[pairid] = nxn_pairid[nxnid] + collision_worldid_out[pairid] = worldid + + # Writing -1 to collision_hftri_index_out[pairid] signals + # hfield_midphase to generate a collision pair for every + # potentially colliding triangle + if type1 == int(GeomType.HFIELD.value) or type2 == int(GeomType.HFIELD.value): + collision_hftri_index_out[pairid] = -1 + + +@wp.func +def _binary_search(values: wp.array(dtype=Any), value: Any, lower: int, upper: int) -> int: + while lower < upper: + mid = (lower + upper) >> 1 + if values[mid] > value: + upper = mid + else: + lower = mid + 1 + + return upper + + +@wp.kernel +def _sap_project( + # Model: + geom_rbound: wp.array2d(dtype=float), + geom_margin: wp.array2d(dtype=float), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + # In: + direction_in: wp.vec3, + # Data out: + sap_projection_lower_out: wp.array2d(dtype=float), # kernel_analyzer: ignore + sap_projection_upper_out: wp.array2d(dtype=float), + sap_sort_index_out: wp.array2d(dtype=int), # kernel_analyzer: ignore +): + worldid, geomid = wp.tid() + + xpos = geom_xpos_in[worldid, geomid] + rbound = geom_rbound[worldid, geomid] + + if rbound == 0.0: + # geom is a plane + rbound = MJ_MAXVAL + + radius = rbound + geom_margin[worldid, geomid] + center = wp.dot(direction_in, xpos) + + sap_sort_index_out[worldid, geomid] = geomid + if not wp.isnan(center): + sap_projection_lower_out[worldid, geomid] = center - radius + sap_projection_upper_out[worldid, geomid] = center + radius + else: + sap_projection_lower_out[worldid, geomid] = MJ_MAXVAL + sap_projection_upper_out[worldid, geomid] = MJ_MAXVAL + + +@wp.kernel +def _sap_range( + # Model: + ngeom: int, + # Data in: + sap_projection_lower_in: wp.array2d(dtype=float), # kernel_analyzer: ignore + sap_projection_upper_in: wp.array2d(dtype=float), + sap_sort_index_in: wp.array2d(dtype=int), # kernel_analyzer: ignore + # Data out: + sap_range_out: wp.array2d(dtype=int), +): + worldid, geomid = wp.tid() + + # current bounding geom + idx = sap_sort_index_in[worldid, geomid] + + upper = sap_projection_upper_in[worldid, idx] + + limit = _binary_search(sap_projection_lower_in[worldid], upper, geomid + 1, ngeom) + limit = wp.min(ngeom - 1, limit) + + # range of geoms for the sweep and prune process + sap_range_out[worldid, geomid] = limit - geomid + + +@wp.kernel +def _sap_broadphase( + # Model: + ngeom: int, + opt_broadphase_filter: int, + geom_type: wp.array(dtype=int), + geom_aabb: wp.array2d(dtype=wp.vec3), + geom_rbound: wp.array2d(dtype=float), + geom_margin: wp.array2d(dtype=float), + nxn_pairid: wp.array(dtype=int), + # Data in: + nworld_in: int, + nconmax_in: int, + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + sap_sort_index_in: wp.array2d(dtype=int), # kernel_analyzer: ignore + sap_cumulative_sum_in: wp.array(dtype=int), # kernel_analyzer: ignore + # In: + nsweep_in: int, + # Data out: + collision_pair_out: wp.array(dtype=wp.vec2i), + collision_hftri_index_out: wp.array(dtype=int), + collision_pairid_out: wp.array(dtype=int), + collision_worldid_out: wp.array(dtype=int), + ncollision_out: wp.array(dtype=int), +): + worldgeomid = wp.tid() + + nworldgeom = nworld_in * ngeom + nworkpackages = sap_cumulative_sum_in[nworldgeom - 1] + + while worldgeomid < nworkpackages: + # binary search to find current and next geom pair indices + i = _binary_search(sap_cumulative_sum_in, worldgeomid, 0, nworldgeom) + j = i + worldgeomid + 1 + + if i > 0: + j -= sap_cumulative_sum_in[i - 1] + + worldid = i // ngeom + i = i % ngeom + j = j % ngeom + + # get geom indices and swap if necessary + geom1 = sap_sort_index_in[worldid, i] + geom2 = sap_sort_index_in[worldid, j] + + # find linear index of (geom1, geom2) in upper triangular nxn_pairid + if geom2 < geom1: + idx = upper_tri_index(ngeom, geom2, geom1) + else: + idx = upper_tri_index(ngeom, geom1, geom2) + + if nxn_pairid[idx] < -1: + worldgeomid += nsweep_in + continue + + if _broadphase_filter( + opt_broadphase_filter, geom_aabb, geom_rbound, geom_margin, geom_xpos_in, geom_xmat_in, geom1, geom2, worldid + ): + _add_geom_pair( + geom_type, + nxn_pairid, + nconmax_in, + geom1, + geom2, + worldid, + idx, + collision_pair_out, + collision_hftri_index_out, + collision_pairid_out, + collision_worldid_out, + ncollision_out, + ) + + worldgeomid += nsweep_in + + +def _segmented_sort(tile_size: int): + @wp.kernel + def segmented_sort( + # Data in: + sap_projection_lower_in: wp.array2d(dtype=float), # kernel_analyzer: ignore + sap_sort_index_in: wp.array2d(dtype=int), # kernel_analyzer: ignore + ): + worldid = wp.tid() + + # Load input into shared memory + keys = wp.tile_load(sap_projection_lower_in[worldid], shape=tile_size, storage="shared") + values = wp.tile_load(sap_sort_index_in[worldid], shape=tile_size, storage="shared") + + # Perform in-place sorting + wp.tile_sort(keys, values) + + # Store sorted shared memory into output arrays + wp.tile_store(sap_projection_lower_in[worldid], keys) + wp.tile_store(sap_sort_index_in[worldid], values) + + return segmented_sort + + +@event_scope +def sap_broadphase(m: Model, d: Data): + """Runs broadphase collision detection using a sweep-and-prune (SAP) algorithm. + + This method is more efficient than the N-squared approach for large numbers of + objects. It works by projecting the bounding spheres of all geoms onto a + single axis and sorting them. It then sweeps along the axis, only checking + for overlaps between geoms whose projections are close to each other. + + For each potentially colliding pair identified by the sweep, a more precise + bounding sphere check is performed. If this check passes, the pair is added + to the collision arrays in `d` for the narrowphase stage. + + Two sorting strategies are supported, controlled by `m.opt.broadphase`: + - `SAP_TILE`: Uses a tile-based sort. + - `SAP_SEGMENTED`: Uses a segmented sort. + """ + + nworldgeom = d.nworld * m.ngeom + + # TODO(team): direction + + # random fixed direction + direction = wp.vec3(0.5935, 0.7790, 0.1235) + direction = wp.normalize(direction) + + wp.launch( + kernel=_sap_project, + dim=(d.nworld, m.ngeom), + inputs=[ + m.geom_rbound, + m.geom_margin, + d.geom_xpos, + direction, + ], + outputs=[ + d.sap_projection_lower.reshape((-1, m.ngeom)), + d.sap_projection_upper, + d.sap_sort_index.reshape((-1, m.ngeom)), + ], + ) + + if m.opt.broadphase == int(BroadphaseType.SAP_TILE): + wp.launch_tiled( + kernel=_segmented_sort(m.ngeom), + dim=(d.nworld), + inputs=[d.sap_projection_lower.reshape((-1, m.ngeom)), d.sap_sort_index.reshape((-1, m.ngeom))], + block_dim=m.block_dim.segmented_sort, + ) + else: + wp.utils.segmented_sort_pairs( + d.sap_projection_lower.reshape((-1, m.ngeom)), + d.sap_sort_index.reshape((-1, m.ngeom)), + nworldgeom, + d.sap_segment_index.reshape(-1), + ) + + wp.launch( + kernel=_sap_range, + dim=(d.nworld, m.ngeom), + inputs=[ + m.ngeom, + d.sap_projection_lower.reshape((-1, m.ngeom)), + d.sap_projection_upper, + d.sap_sort_index.reshape((-1, m.ngeom)), + ], + outputs=[ + d.sap_range, + ], + ) + + # scan is used for load balancing among the threads + wp.utils.array_scan(d.sap_range.reshape(-1), d.sap_cumulative_sum.reshape(-1), True) + + # estimate number of overlap checks + # assumes each geom has 5 other geoms (batched over all worlds) + nsweep = 5 * nworldgeom + wp.launch( + kernel=_sap_broadphase, + dim=nsweep, + inputs=[ + m.ngeom, + m.opt.broadphase_filter, + m.geom_type, + m.geom_aabb, + m.geom_rbound, + m.geom_margin, + m.nxn_pairid, + d.nworld, + d.nconmax, + d.geom_xpos, + d.geom_xmat, + d.sap_sort_index.reshape((-1, m.ngeom)), + d.sap_cumulative_sum.reshape(-1), + nsweep, + ], + outputs=[ + d.collision_pair, + d.collision_hftri_index, + d.collision_pairid, + d.collision_worldid, + d.ncollision, + ], + ) + + +@wp.kernel +def _nxn_broadphase( + # Model: + opt_broadphase_filter: int, + geom_type: wp.array(dtype=int), + geom_aabb: wp.array2d(dtype=wp.vec3), + geom_rbound: wp.array2d(dtype=float), + geom_margin: wp.array2d(dtype=float), + nxn_geom_pair: wp.array(dtype=wp.vec2i), + nxn_pairid: wp.array(dtype=int), + # Data in: + nconmax_in: int, + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # Data out: + collision_pair_out: wp.array(dtype=wp.vec2i), + collision_hftri_index_out: wp.array(dtype=int), + collision_pairid_out: wp.array(dtype=int), + collision_worldid_out: wp.array(dtype=int), + ncollision_out: wp.array(dtype=int), +): + worldid, elementid = wp.tid() + + geom = nxn_geom_pair[elementid] + geom1 = geom[0] + geom2 = geom[1] + + if _broadphase_filter( + opt_broadphase_filter, geom_aabb, geom_rbound, geom_margin, geom_xpos_in, geom_xmat_in, geom1, geom2, worldid + ): + _add_geom_pair( + geom_type, + nxn_pairid, + nconmax_in, + geom1, + geom2, + worldid, + elementid, + collision_pair_out, + collision_hftri_index_out, + collision_pairid_out, + collision_worldid_out, + ncollision_out, + ) + + +@event_scope +def nxn_broadphase(m: Model, d: Data): + """Runs broadphase collision detection using a brute-force N-squared approach. + + This function iterates through a pre-filtered list of all possible geometry pairs and + performs a quick bounding sphere check to identify potential collisions. + + For each pair that passes the sphere check, it populates the collision arrays in `d` + (`d.collision_pair`, `d.collision_pairid`, etc.), which are then consumed by the + narrowphase. + + The initial list of pairs is filtered at model creation time to exclude pairs based on + `contype`/`conaffinity`, parent-child relationships, and explicit `` tags. + """ + + wp.launch( + _nxn_broadphase, + dim=(d.nworld, m.nxn_geom_pair_filtered.shape[0]), + inputs=[ + m.opt.broadphase_filter, + m.geom_type, + m.geom_aabb, + m.geom_rbound, + m.geom_margin, + m.nxn_geom_pair_filtered, + m.nxn_pairid_filtered, + d.nconmax, + d.geom_xpos, + d.geom_xmat, + ], + outputs=[ + d.collision_pair, + d.collision_hftri_index, + d.collision_pairid, + d.collision_worldid, + d.ncollision, + ], + ) + + +def _narrowphase(m, d): + # Process heightfield collisions + if m.nhfield > 0: + hfield_midphase(m, d) + + # TODO(team): we should reject far-away contacts in the narrowphase instead of constraint + # partitioning because we can move some pressure of the atomics + convex_narrowphase(m, d) + primitive_narrowphase(m, d) + + if m.has_sdf_geom: + sdf_narrowphase(m, d) + + +@event_scope +def collision(m: Model, d: Data): + """Runs the full collision detection pipeline. + + This function orchestrates the broadphase and narrowphase collision detection stages. It + first identifies potential collision pairs using a broadphase algorithm (either N-squared + or Sweep-and-Prune, based on `m.opt.broadphase`). Then, for each potential pair, it + performs narrowphase collision detection to compute detailed contact information like + distance, position, and frame. + + The results are used to populate the `d.contact` array, and the total number of contacts + is stored in `d.ncon`. If `d.ncon` is larger than `d.nconmax` then an overflow has + occurred and the remaining contacts will be skipped. If this happens, raise the `nconmax` + parameter in `io.make_data` or `io.put_data`. + + This function will do nothing except zero out arrays if collision detection is disabled + via `m.opt.disableflags` or if `d.nconmax` is 0. + """ + + # zero collision-related arrays + wp.launch( + _zero_collision_arrays, + dim=d.nconmax, + inputs=[ + d.nworld, + d.ncon_hfield.shape[1], + d.ncon, + d.ncon_hfield.reshape(-1), + d.collision_hftri_index, + d.ncollision, + ], + ) + + if d.nconmax == 0 or m.opt.disableflags & (DisableBit.CONSTRAINT | DisableBit.CONTACT): + return + + if m.opt.broadphase == int(BroadphaseType.NXN): + nxn_broadphase(m, d) + else: + sap_broadphase(m, d) + + if m.opt.graph_conditional: + wp.capture_if(condition=d.ncollision, on_true=_narrowphase, m=m, d=d) + else: + _narrowphase(m, d) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py new file mode 100644 index 00000000..130121d4 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py @@ -0,0 +1,870 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Tests the collision driver.""" + +import mujoco +import numpy as np +import warp as wp +from absl.testing import absltest +from absl.testing import parameterized + +import mujoco_warp as mjwarp +from mujoco.mjx.third_party.mujoco_warp.test_data.collision_sdf.utils import register_sdf_plugins + +from mujoco.mjx.third_party.mujoco_warp._src import test_util +from mujoco.mjx.third_party.mujoco_warp._src import types + + +class CollisionTest(parameterized.TestCase): + """Tests the collision contact functions.""" + + _SDF_SDF = { + "_NUT_NUT": """ + + + + + + + + + + + + + +""", + "NUT_BOLT": """ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +""", + } + _FIXTURES = { + "box_plane": """ + + + + + + + + + + """, + "box_box_vf": """ + + + + + + + + + + + + """, + "box_box_vf_flat": """ + + + + + + + + + + + + """, + "box_box_ee": """ + + + + + + + + + + + + """, + "box_box_ee_deep": """ + + + + + + + + + + + + """, + "plane_sphere": """ + + + + + + + + + + """, + "plane_ellipsoid": """ + + + + + + + + + + """, + "plane_capsule": """ + + + + + + + + + + """, + "convex_convex": """ + + + + + + + + + + + + + + + + """, + "capsule_capsule": """ + + + + + + + + + + + + + """, + "sphere_sphere": """ + + + + + + + + + + + + + """, + "sphere_capsule": """ + + + + + + + + + + + + + """, + "sphere_cylinder_corner": """ + + + + + + + + + + + + + + """, + "sphere_cylinder_cap": """ + + + + + + + + + + + + + + """, + "sphere_cylinder_side": """ + + + + + + + + + + + + + + """, + "plane_cylinder_1": """ + + + + + + + + + + """, + "plane_cylinder_2": """ + + + + + + + + + + """, + "plane_cylinder_3": """ + + + + + + + + + + """, + "mesh_plane_simple": """ + + + + + + + + + + + + + """, + "mesh_plane_complex": """ + + + + + + + + + + + + + """, + "sphere_box_shallow": """ + + + + + + + + + + """, + "sphere_box_deep": """ + + + + + + + + + + """, + "capsule_box_edge": """ + + + + + + + + + + """, + "capsule_box_corner": """ + + + + + + + + + + """, + "capsule_box_face_tip": """ + + + + + + + + + + """, + "capsule_box_face_flat": """ + + + + + + + + + + """, + } + + @classmethod + def setUpClass(cls): + register_sdf_plugins(mjwarp._src.collision_sdf) + + @parameterized.parameters(_SDF_SDF.keys()) + def test_sdf_collision(self, fixture): + """Tests collisions with different geometries.""" + + mjm, mjd, m, d = test_util.fixture(xml=self._SDF_SDF[fixture], qpos0=True) + + mujoco.mj_collision(mjm, mjd) + mjwarp.collision(m, d) + for i in range(min(mjd.ncon, d.ncon.numpy()[0])): + actual_dist = mjd.contact.dist[i] + actual_pos = mjd.contact.pos[i] + actual_frame = mjd.contact.frame[i][0:3] + result = False + test_dist = d.contact.dist.numpy()[i] + test_pos = d.contact.pos.numpy()[i, :] + test_frame = d.contact.frame.numpy()[i].flatten()[0:3] + check_dist = np.allclose(actual_dist, test_dist, rtol=5e-2, atol=1.0e-1) + check_frame = np.allclose(actual_frame, test_frame, rtol=5e-2, atol=1.0e-1) + check_pos = np.allclose(actual_pos, test_pos, rtol=5e-2, atol=1.0e-1) + result = check_dist + np.testing.assert_equal(result, True, f"Contact {i} not found in Gjk results") + + @parameterized.parameters(_FIXTURES.keys()) + def test_collision(self, fixture): + """Tests collisions with different geometries.""" + mjm, mjd, m, d = test_util.fixture(xml=self._FIXTURES[fixture], qpos0=True) + + # Exempt GJK collisions from exact contact count check + # because GJK generates more contacts + allow_different_contact_count = False + + mujoco.mj_collision(mjm, mjd) + mjwarp.collision(m, d) + + self.assertGreater(d.ncon.numpy()[0], 0) + self.assertGreater(mjd.ncon, 0) + + for i in range(mjd.ncon): + actual_dist = mjd.contact.dist[i] + actual_pos = mjd.contact.pos[i] + actual_frame = mjd.contact.frame[i] + result = False + for j in range(d.ncon.numpy()[0]): + test_dist = d.contact.dist.numpy()[j] + test_pos = d.contact.pos.numpy()[j, :] + test_frame = d.contact.frame.numpy()[j].flatten() + check_dist = np.allclose(actual_dist, test_dist, rtol=5e-2, atol=1.0e-2) + check_pos = np.allclose(actual_pos, test_pos, rtol=5e-2, atol=1.0e-2) + check_frame = np.allclose(actual_frame, test_frame, rtol=5e-2, atol=1.0e-2) + if check_dist and check_pos and check_frame: + result = True + break + np.testing.assert_equal(result, True, f"Contact {i} not found in Gjk results") + + if not allow_different_contact_count: + self.assertEqual(d.ncon.numpy()[0], mjd.ncon) + + _HFIELD_FIXTURES = { + "hfield_box": """ + + + + + + + + + + + + + """, + } + + @parameterized.parameters(_HFIELD_FIXTURES.keys()) + def test_hfield_collision(self, fixture): + """Tests hfield collision with different geometries.""" + mjm, mjd, m, d = test_util.fixture(xml=self._HFIELD_FIXTURES[fixture]) + + mujoco.mj_collision(mjm, mjd) + mjwarp.collision(m, d) + + self.assertEqual(mjd.ncon > 0, d.ncon.numpy()[0] > 0, "If MJ collides, MJW should too") + + def test_contact_exclude(self): + """Tests contact exclude.""" + _, _, m, _ = test_util.fixture( + xml=""" + + + + + + + + + + + + + + + + + + + + """ + ) + self.assertEqual(m.nxn_geom_pair.numpy().shape[0], 3) + np.testing.assert_equal(m.nxn_pairid.numpy(), np.array([-2, -1, -1])) + + def test_contact_pair(self): + """Tests contact pair.""" + # no pairs + _, _, m, _ = test_util.fixture( + xml=""" + + + + + + + + + """ + ) + self.assertTrue((m.nxn_pairid.numpy() == -1).all()) + + # 1 pair + _, _, m, d = test_util.fixture( + xml=""" + + + + + + + + + + + + + + + + """, + qpos0=True, + ) + self.assertTrue((m.nxn_pairid.numpy() == 0).all()) + + for arr in ( + d.ncon, + d.contact.includemargin, + d.contact.dim, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + ): + arr.zero_() + + mjwarp.collision(m, d) + + self.assertEqual(d.ncon.numpy()[0], 1) + self.assertEqual(d.contact.includemargin.numpy()[0], -1) + self.assertEqual(d.contact.dim.numpy()[0], 6) + np.testing.assert_allclose(d.contact.friction.numpy()[0], np.array([5, 4, 3, 2, 1])) + np.testing.assert_allclose(d.contact.solref.numpy()[0], np.array([-0.25, -0.5])) + np.testing.assert_allclose(d.contact.solreffriction.numpy()[0], np.array([2.0, 4.0])) + np.testing.assert_allclose(d.contact.solimp.numpy()[0], np.array([0.1, 0.2, 0.3, 0.4, 0.5])) + + # 1 pair: override contype and conaffinity + _, _, m, d = test_util.fixture( + xml=""" + + + + + + + + + + + + + + + + """, + qpos0=True, + ) + self.assertTrue((m.nxn_pairid.numpy() == 0).all()) + + for arr in ( + d.ncon, + d.contact.includemargin, + d.contact.dim, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + ): + arr.zero_() + + mjwarp.collision(m, d) + + self.assertEqual(d.ncon.numpy()[0], 1) + self.assertEqual(d.contact.includemargin.numpy()[0], -1) + self.assertEqual(d.contact.dim.numpy()[0], 6) + np.testing.assert_allclose(d.contact.friction.numpy()[0], np.array([5, 4, 3, 2, 1])) + np.testing.assert_allclose(d.contact.solref.numpy()[0], np.array([-0.25, -0.5])) + np.testing.assert_allclose(d.contact.solreffriction.numpy()[0], np.array([2.0, 4.0])) + np.testing.assert_allclose(d.contact.solimp.numpy()[0], np.array([0.1, 0.2, 0.3, 0.4, 0.5])) + + # 1 pair: override exclude + _, _, m, d = test_util.fixture( + xml=""" + + + + + + + + + + + + + + + + + """, + qpos0=True, + ) + self.assertTrue((m.nxn_pairid.numpy() == 0).all()) + + for arr in ( + d.ncon, + d.contact.includemargin, + d.contact.dim, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + ): + arr.zero_() + + mjwarp.collision(m, d) + + self.assertEqual(d.ncon.numpy()[0], 1) + self.assertEqual(d.contact.includemargin.numpy()[0], -1) + self.assertEqual(d.contact.dim.numpy()[0], 6) + np.testing.assert_allclose(d.contact.friction.numpy()[0], np.array([5, 4, 3, 2, 1])) + np.testing.assert_allclose(d.contact.solref.numpy()[0], np.array([-0.25, -0.5])) + np.testing.assert_allclose(d.contact.solreffriction.numpy()[0], np.array([2.0, 4.0])) + np.testing.assert_allclose(d.contact.solimp.numpy()[0], np.array([0.1, 0.2, 0.3, 0.4, 0.5])) + + # 1 pair 1 exclude + _, _, m, d = test_util.fixture( + xml=""" + + + + + + + + + + + + + + + + + + + + + """, + qpos0=True, + ) + np.testing.assert_equal(m.nxn_pairid.numpy(), np.array([-2, -1, 0])) + + for arr in ( + d.ncon, + d.contact.includemargin, + d.contact.dim, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + ): + arr.zero_() + + mjwarp.collision(m, d) + + self.assertEqual(d.ncon.numpy()[0], 2) + self.assertEqual(d.contact.includemargin.numpy()[1], -1) + self.assertEqual(d.contact.dim.numpy()[1], 6) + np.testing.assert_allclose(d.contact.friction.numpy()[1], np.array([5, 4, 3, 2, 1])) + np.testing.assert_allclose(d.contact.solref.numpy()[1], np.array([-0.25, -0.5])) + np.testing.assert_allclose(d.contact.solreffriction.numpy()[1], np.array([2.0, 4.0])) + np.testing.assert_allclose(d.contact.solimp.numpy()[1], np.array([0.1, 0.2, 0.3, 0.4, 0.5])) + + # TODO(team): test sap_broadphase + + @parameterized.parameters( + (True, True), + (True, False), + (False, True), + (False, False), + ) + def test_collision_disableflags(self, constraint, contact): + """Tests collision disableflags.""" + mjm, mjd, m, d = test_util.fixture( + "humanoid/humanoid.xml", + keyframe=0, + constraint=constraint, + contact=contact, + kick=False, + ) + + mujoco.mj_collision(mjm, mjd) + mjwarp.collision(m, d) + + self.assertEqual(d.ncon.numpy()[0], mjd.ncon) + + def test_hfield_maxconpair(self): + _XML = f""" + + + + + + + + + + + + + + + + """ + + _, _, m, d = test_util.fixture(xml=_XML, keyframe=0) + + mjwarp.collision(m, d) + + np.testing.assert_equal(d.ncon.numpy()[0], types.MJ_MAXCONPAIR) + + def test_min_friction(self): + _, _, _, d = test_util.fixture( + xml=""" + + + + + + + + + + + + + + + + """, + keyframe=0, + ) + + self.assertEqual(d.ncon.numpy()[0], 1) + np.testing.assert_allclose(d.contact.friction.numpy()[0], types.MJ_MINMU) + + # TODO(team): test contact parameter mixing + + +if __name__ == "__main__": + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py new file mode 100644 index 00000000..8128262a --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py @@ -0,0 +1,1361 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_prism_vertex +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import Geom +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType + +# TODO(team): improve compile time to enable backward pass +wp.config.enable_backward = False + +FLOAT_MIN = -1e30 +FLOAT_MAX = 1e30 +MJ_MINVAL2 = MJ_MINVAL * MJ_MINVAL + +mat43 = wp.types.matrix(shape=(4, 3), dtype=float) + + +@wp.struct +class GJKResult: + dist: float + x1: wp.vec3 + x2: wp.vec3 + dim: int + simplex: mat43 + simplex1: mat43 + simplex2: mat43 + simplex_index1: wp.vec4i + simplex_index2: wp.vec4i + + +@wp.struct +class Polytope: + status: int + + # vertices in polytope + vert: wp.array(dtype=wp.vec3) + vert1: wp.array(dtype=wp.vec3) + vert2: wp.array(dtype=wp.vec3) + vert_index1: wp.array(dtype=int) + vert_index2: wp.array(dtype=int) + nvert: int + + # faces in polytope + face: wp.array(dtype=wp.vec3i) + face_pr: wp.array(dtype=wp.vec3) + face_norm2: wp.array(dtype=float) + face_index: wp.array(dtype=int) + nface: int + + # TODO(kbayes): look into if a linear map actually improves performance + face_map: wp.array(dtype=int) + nmap: int + + # edges that make up the horizon when adding new vertices to polytope + horizon: wp.array(dtype=int) + nhorizon: int + + +@wp.func +def _support(geom: Geom, geomtype: int, dir: wp.vec3): + cached_index = -1 + vertex_index = -1 + local_dir = wp.transpose(geom.rot) @ dir + if geomtype == int(GeomType.SPHERE.value): + support_pt = geom.pos + geom.size[0] * dir + elif geomtype == int(GeomType.BOX.value): + tmp = wp.sign(local_dir) + res = wp.cw_mul(tmp, geom.size) + support_pt = geom.rot @ res + geom.pos + vertex_index = 0 + if tmp[0] > 0: + vertex_index += 1 + if tmp[1] > 0: + vertex_index += 2 + if tmp[2] > 0: + vertex_index += 4 + elif geomtype == int(GeomType.CAPSULE.value): + res = local_dir * geom.size[0] + # add cylinder contribution + res[2] += wp.sign(local_dir[2]) * geom.size[1] + support_pt = geom.rot @ res + geom.pos + elif geomtype == int(GeomType.ELLIPSOID.value): + res = wp.cw_mul(local_dir, geom.size) + res = wp.normalize(res) + # transform to ellipsoid + res = wp.cw_mul(res, geom.size) + support_pt = geom.rot @ res + geom.pos + elif geomtype == int(GeomType.CYLINDER.value): + res = wp.vec3(0.0, 0.0, 0.0) + # set result in XY plane: support on circle + d = wp.sqrt(local_dir[0] * local_dir[0] + local_dir[1] * local_dir[1]) + if d > MJ_MINVAL: + scl = geom.size[0] / d + res[0] = local_dir[0] * scl + res[1] = local_dir[1] * scl + # set result in Z direction + res[2] = wp.sign(local_dir[2]) * geom.size[1] + support_pt = geom.rot @ res + geom.pos + elif geomtype == int(GeomType.MESH.value): + max_dist = float(FLOAT_MIN) + if geom.graphadr == -1 or geom.vertnum < 10: + if geom.index > -1: + cached_index = geom.index + max_dist = wp.dot(geom.vert[geom.index], local_dir) + support_pt = geom.vert[geom.index] + # exhaustive search over all vertices + for i in range(geom.vertnum): + vert = geom.vert[geom.vertadr + i] + dist = wp.dot(vert, local_dir) + if dist > max_dist: + max_dist = dist + support_pt = vert + cached_index = geom.vertadr + i + vertex_index = cached_index - geom.vertadr + else: + numvert = geom.graph[geom.graphadr] + vert_edgeadr = geom.graphadr + 2 + vert_globalid = geom.graphadr + 2 + numvert + edge_localid = geom.graphadr + 2 + 2 * numvert + # hillclimb until no change + prev = int(-1) + imax = int(0) + if geom.index > -1: + imax = geom.index + cached_index = geom.index + + while True: + prev = int(imax) + i = int(geom.graph[vert_edgeadr + imax]) + while geom.graph[edge_localid + i] >= 0: + subidx = geom.graph[edge_localid + i] + idx = geom.graph[vert_globalid + subidx] + dist = wp.dot(local_dir, geom.vert[geom.vertadr + idx]) + if dist > max_dist: + max_dist = dist + imax = int(subidx) + i += int(1) + if imax == prev: + break + cached_index = imax + imax = geom.graph[vert_globalid + imax] + vertex_index = imax + support_pt = geom.vert[geom.vertadr + imax] + + support_pt = geom.rot @ support_pt + geom.pos + elif geomtype == int(GeomType.HFIELD.value): + max_dist = float(FLOAT_MIN) + for i in range(6): + vert = hfield_prism_vertex(geom.hfprism, i) + dist = wp.dot(vert, local_dir) + if dist > max_dist: + max_dist = dist + support_pt = vert + support_pt = geom.rot @ support_pt + geom.pos + + return support_pt, cached_index, vertex_index + + +@wp.func +def _attach_face(pt: Polytope, idx: int, v1: int, v2: int, v3: int): + # out of memory, returning 0 will force EPA to return early without contact + if pt.nface == pt.face.shape[0]: + return 0.0 + + # compute witness point v + r, ret = _project_origin_plane(pt.vert[v3], pt.vert[v2], pt.vert[v1]) + if ret: + return 0.0 + + face = wp.vec3i(v1, v2, v3) + pt.face[idx] = face + pt.face_pr[idx] = r + + pt.face_norm2[idx] = wp.dot(r, r) + pt.face_index[idx] = -1 + return pt.face_norm2[idx] + + +@wp.func +def _epa_support(pt: Polytope, idx: int, geom1: Geom, geom2: Geom, geom1_type: int, geom2_type: int, dir: wp.vec3): + s1, index1, vertex_index1 = _support(geom1, geom1_type, dir) + s2, index2, vertex_index2 = _support(geom2, geom2_type, -dir) + + pt.vert[idx] = s1 - s2 + pt.vert1[idx] = s1 + pt.vert2[idx] = s2 + pt.vert_index1[idx] = vertex_index1 + pt.vert_index2[idx] = vertex_index2 + return index1, index2 + + +@wp.func +def _linear_combine(n: int, coefs: wp.vec4, mat: mat43): + v = wp.vec3(0.0) + if n == 1: + v = coefs[0] * mat[0] + elif n == 2: + v = coefs[0] * mat[0] + coefs[1] * mat[1] + elif n == 3: + v = coefs[0] * mat[0] + coefs[1] * mat[1] + coefs[2] * mat[2] + else: + v = coefs[0] * mat[0] + coefs[1] * mat[1] + coefs[2] * mat[2] + coefs[3] * mat[3] + return v + + +@wp.func +def _almost_equal(v1: wp.vec3, v2: wp.vec3): + return wp.abs(v1[0] - v2[0]) < MJ_MINVAL and wp.abs(v1[1] - v2[1]) < MJ_MINVAL and wp.abs(v1[2] - v2[2]) < MJ_MINVAL + + +@wp.func +def _subdistance(n: int, simplex: mat43): + if n == 4: + return _S3D(simplex[0], simplex[1], simplex[2], simplex[3]) + if n == 3: + coordinates3 = _S2D(simplex[0], simplex[1], simplex[2]) + return wp.vec4(coordinates3[0], coordinates3[1], coordinates3[2], 0.0) + if n == 2: + coordinates2 = _S1D(simplex[0], simplex[1]) + return wp.vec4(coordinates2[0], coordinates2[1], 0.0, 0.0) + return wp.vec4(1.0, 0.0, 0.0, 0.0) + + +@wp.func +def _det3(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3): + return wp.dot(v1, wp.cross(v2, v3)) + + +@wp.func +def _same_sign(a: float, b: float): + if a > 0 and b > 0: + return 1 + if a < 0 and b < 0: + return -1 + return 0 + + +@wp.func +def _project_origin_line(v1: wp.vec3, v2: wp.vec3): + diff = v2 - v1 + scl = -(wp.dot(v2, diff) / wp.dot(diff, diff)) + return v2 + scl * diff + + +@wp.func +def _project_origin_plane(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3): + z = wp.vec3(0.0) + diff21 = v2 - v1 + diff31 = v3 - v1 + diff32 = v3 - v2 + + # n = (v1 - v2) x (v3 - v2) + n = wp.cross(diff32, diff21) + nv = wp.dot(n, v2) + nn = wp.dot(n, n) + if nn == 0: + return z, 1 + if nv != 0 and nn > MJ_MINVAL: + v = (nv / nn) * n + return v, 0 + + # n = (v2 - v1) x (v3 - v1) + n = wp.cross(diff21, diff31) + nv = wp.dot(n, v1) + nn = wp.dot(n, n) + if nn == 0: + return z, 1 + if nv != 0 and nn > MJ_MINVAL: + v = (nv / nn) * n + return v, 0 + + # n = (v1 - v3) x (v2 - v3) + n = wp.cross(diff31, diff32) + nv = wp.dot(n, v3) + nn = wp.dot(n, n) + v = (nv / nn) * n + return v, 0 + + +@wp.func +def _S3D(s1: wp.vec3, s2: wp.vec3, s3: wp.vec3, s4: wp.vec3): + # [[ s1_x, s2_x, s3_x, s4_x ], + # [ s1_y, s2_y, s3_y, s4_y ], + # [ s1_z, s2_z, s3_z, s4_z ], + # [ 1, 1, 1, 1 ]] + # we want to solve M*lambda = P, where P = [p_x, p_y, p_z, 1] with [p_x, p_y, p_z] is the + # origin projected onto the simplex + + # compute cofactors to find det(M) + C41 = -_det3(s2, s3, s4) + C42 = _det3(s1, s3, s4) + C43 = -_det3(s1, s2, s4) + C44 = _det3(s1, s2, s3) + + # NOTE: m_det = 6*SignVol(simplex) with C4i corresponding to the volume of the 3-simplex + # with vertices {s1, s2, s3, 0} - si + m_det = C41 + C42 + C43 + C44 + + comp1 = _same_sign(m_det, C41) + comp2 = _same_sign(m_det, C42) + comp3 = _same_sign(m_det, C43) + comp4 = _same_sign(m_det, C44) + + # if all signs are the same then the origin is inside the simplex + if comp1 and comp2 and comp3 and comp4: + return wp.vec4(C41 / m_det, C42 / m_det, C43 / m_det, C44 / m_det) + + # find the smallest distance, and use the corresponding barycentric coordinates + coordinates = wp.vec4(0.0, 0.0, 0.0, 0.0) + dmin = FLOAT_MAX + + if not comp1: + subcoord = _S2D(s2, s3, s4) + x = subcoord[0] * s2 + subcoord[1] * s3 + subcoord[2] * s4 + d = wp.dot(x, x) + coordinates[0] = 0.0 + coordinates[1] = subcoord[0] + coordinates[2] = subcoord[1] + coordinates[3] = subcoord[2] + dmin = d + + if not comp2: + subcoord = _S2D(s1, s3, s4) + x = subcoord[0] * s1 + subcoord[1] * s3 + subcoord[2] * s4 + d = wp.dot(x, x) + if d < dmin: + coordinates[0] = subcoord[0] + coordinates[1] = 0.0 + coordinates[2] = subcoord[1] + coordinates[3] = subcoord[2] + dmin = d + + if not comp3: + subcoord = _S2D(s1, s2, s4) + x = subcoord[0] * s1 + subcoord[1] * s2 + subcoord[2] * s4 + d = wp.dot(x, x) + if d < dmin: + coordinates[0] = subcoord[0] + coordinates[1] = subcoord[1] + coordinates[2] = 0.0 + coordinates[3] = subcoord[2] + dmin = d + + if not comp4: + subcoord = _S2D(s1, s2, s3) + x = subcoord[0] * s1 + subcoord[1] * s2 + subcoord[2] * s3 + d = wp.dot(x, x) + if d < dmin: + coordinates[0] = subcoord[0] + coordinates[1] = subcoord[1] + coordinates[2] = subcoord[2] + coordinates[3] = 0.0 + return coordinates + + +@wp.func +def _S2D(s1: wp.vec3, s2: wp.vec3, s3: wp.vec3): + # project origin onto affine hull of the simplex + p_o, ret = _project_origin_plane(s1, s2, s3) + if ret: + v = _S1D(s1, s2) + return wp.vec3(v[0], v[1], 0.0) + + # Below are the minors M_i4 of the matrix M given by + # [[ s1_x, s2_x, s3_x, s4_x ], + # [ s1_y, s2_y, s3_y, s4_y ], + # [ s1_z, s2_z, s3_z, s4_z ], + # [ 1, 1, 1, 1 ]] + M_14 = s2[1] * s3[2] - s2[2] * s3[1] - s1[1] * s3[2] + s1[2] * s3[1] + s1[1] * s2[2] - s1[2] * s2[1] + M_24 = s2[0] * s3[2] - s2[2] * s3[0] - s1[0] * s3[2] + s1[2] * s3[0] + s1[0] * s2[2] - s1[2] * s2[0] + M_34 = s2[0] * s3[1] - s2[1] * s3[0] - s1[0] * s3[1] + s1[1] * s3[0] + s1[0] * s2[1] - s1[1] * s2[0] + + # exclude the axis with the largest projection of the simplex using the computed minors + M_max = 0.0 + s1_2D = wp.vec2(0.0) + s2_2D = wp.vec2(0.0) + s3_2D = wp.vec2(0.0) + p_o_2D = wp.vec2(0.0) + + mu1 = wp.abs(M_14) + mu2 = wp.abs(M_24) + mu3 = wp.abs(M_34) + + if mu1 >= mu2 and mu1 >= mu3: + M_max = M_14 + s1_2D[0] = s1[1] + s1_2D[1] = s1[2] + + s2_2D[0] = s2[1] + s2_2D[1] = s2[2] + + s3_2D[0] = s3[1] + s3_2D[1] = s3[2] + + p_o_2D[0] = p_o[1] + p_o_2D[1] = p_o[2] + elif mu2 >= mu3: + M_max = M_24 + s1_2D[0] = s1[0] + s1_2D[1] = s1[2] + + s2_2D[0] = s2[0] + s2_2D[1] = s2[2] + + s3_2D[0] = s3[0] + s3_2D[1] = s3[2] + + p_o_2D[0] = p_o[0] + p_o_2D[1] = p_o[2] + else: + M_max = M_34 + s1_2D[0] = s1[0] + s1_2D[1] = s1[1] + + s2_2D[0] = s2[0] + s2_2D[1] = s2[1] + + s3_2D[0] = s3[0] + s3_2D[1] = s3[1] + + p_o_2D[0] = p_o[0] + p_o_2D[1] = p_o[1] + + # compute the cofactors C3i of the following matrix: + # [[ s1_2D[0] - p_o_2D[0], s2_2D[0] - p_o_2D[0], s3_2D[0] - p_o_2D[0] ], + # [ s1_2D[1] - p_o_2D[1], s2_2D[1] - p_o_2D[1], s3_2D[1] - p_o_2D[1] ], + # [ 1, 1, 1 ]] + + # C31 corresponds to the signed area of 2-simplex: (p_o_2D, s2_2D, s3_2D) + C31 = ( + p_o_2D[0] * s2_2D[1] + + p_o_2D[1] * s3_2D[0] + + s2_2D[0] * s3_2D[1] + - p_o_2D[0] * s3_2D[1] + - p_o_2D[1] * s2_2D[0] + - s3_2D[0] * s2_2D[1] + ) + + # C32 corresponds to the signed area of 2-simplex: (_po_2D, s1_2D, s3_2D) + C32 = ( + p_o_2D[0] * s3_2D[1] + + p_o_2D[1] * s1_2D[0] + + s3_2D[0] * s1_2D[1] + - p_o_2D[0] * s1_2D[1] + - p_o_2D[1] * s3_2D[0] + - s1_2D[0] * s3_2D[1] + ) + + # C33 corresponds to the signed area of 2-simplex: (p_o_2D, s1_2D, s2_2D) + C33 = ( + p_o_2D[0] * s1_2D[1] + + p_o_2D[1] * s2_2D[0] + + s1_2D[0] * s2_2D[1] + - p_o_2D[0] * s2_2D[1] + - p_o_2D[1] * s1_2D[0] + - s2_2D[0] * s1_2D[1] + ) + + comp1 = _same_sign(M_max, C31) + comp2 = _same_sign(M_max, C32) + comp3 = _same_sign(M_max, C33) + + # all the same sign, p_o is inside the 2-simplex + if comp1 and comp2 and comp3: + return wp.vec3(C31 / M_max, C32 / M_max, C33 / M_max) + + # find the smallest distance, and use the corresponding barycentric coordinates + dmin = FLOAT_MAX + coordinates = wp.vec3(0.0, 0.0, 0.0) + + if not comp1: + subcoord = _S1D(s2, s3) + x = subcoord[0] * s2 + subcoord[1] * s3 + d = wp.dot(x, x) + coordinates[0] = 0.0 + coordinates[1] = subcoord[0] + coordinates[2] = subcoord[1] + dmin = d + + if not comp2: + subcoord = _S1D(s1, s3) + x = subcoord[0] * s1 + subcoord[1] * s3 + d = wp.dot(x, x) + if d < dmin: + coordinates[0] = subcoord[0] + coordinates[1] = 0.0 + coordinates[2] = subcoord[1] + dmin = d + + if not comp3: + subcoord = _S1D(s1, s2) + x = subcoord[0] * s1 + subcoord[1] * s2 + d = wp.dot(x, x) + if d < dmin: + coordinates[0] = subcoord[0] + coordinates[1] = subcoord[1] + coordinates[2] = 0.0 + return coordinates + + +@wp.func +def _S1D(s1: wp.vec3, s2: wp.vec3): + # find projection of origin onto the 1-simplex: + p_o = _project_origin_line(s1, s2) + + # find the axis with the largest projection "shadow" of the simplex + mu_max = 0.0 + index = 0 + for i in range(3): + mu = s1[i] - s2[i] + if wp.abs(mu) >= wp.abs(mu_max): + mu_max = mu + index = i + + C1 = p_o[index] - s2[index] + C2 = s1[index] - p_o[index] + + # inside the simplex + if _same_sign(mu_max, C1) and _same_sign(mu_max, C2): + return wp.vec2(C1 / mu_max, C2 / mu_max) + return wp.vec2(0.0, 1.0) + + +@wp.func +def _gjk( + # In: + tolerance: float, + gjk_iterations: int, + geom1: Geom, + geom2: Geom, + x1_0: wp.vec3, + x2_0: wp.vec3, + geomtype1: int, + geomtype2: int, + cutoff: float, +): + """Find distance within a tolerance between two geoms.""" + cutoff2 = cutoff * cutoff + simplex = mat43() + simplex1 = mat43() + simplex2 = mat43() + simplex_index1 = wp.vec4i() + simplex_index2 = wp.vec4i() + n = int(0) + coordinates = wp.vec4() # barycentric coordinates + epsilon = 0.5 * tolerance * tolerance + + # set initial guess + x_k = x1_0 - x2_0 + + for k in range(gjk_iterations): + xnorm = wp.dot(x_k, x_k) + # TODO(kbayes): determine new constant here + if xnorm < 1e-12: + break + dir_neg = x_k / wp.sqrt(xnorm) + + # compute the kth support point + s1_k, i1, vertex_index1 = _support(geom1, geomtype1, -dir_neg) + s2_k, i2, vertex_index2 = _support(geom2, geomtype2, dir_neg) + geom1.index = i1 + geom2.index = i2 + simplex1[n] = s1_k + simplex2[n] = s2_k + simplex_index1[n] = vertex_index1 + simplex_index2[n] = vertex_index2 + simplex[n] = s1_k - s2_k + + if cutoff == 0.0: + if wp.dot(x_k, simplex[n]) > 0: + result = GJKResult() + result.dim = 0 + result.dist = FLOAT_MAX + return result + elif cutoff < FLOAT_MAX: + vs = wp.dot(x_k, simplex[n]) + vv = wp.dot(x_k, x_k) + if wp.dot(x_k, simplex[n]) > 0 and (vs * vs / vv) >= cutoff2: + result = GJKResult() + result.dim = 0 + result.dist = FLOAT_MAX + return result + + # stopping criteria using the Frank-Wolfe duality gap given by + # |f(x_k) - f(x_min)|^2 <= < grad f(x_k), (x_k - simplex[n]) > + if wp.dot(x_k, x_k - simplex[n]) < epsilon: + break + + # run the distance subalgorithm to compute the barycentric coordinates + # of the closest point to the origin in the simplex + coordinates = _subdistance(n + 1, simplex) + + # remove vertices from the simplex no longer needed + n = int(0) + for i in range(4): + if coordinates[i] == 0: + continue + + simplex[n] = simplex[i] + simplex1[n] = simplex1[i] + simplex2[n] = simplex2[i] + simplex_index1[n] = simplex_index1[i] + simplex_index2[n] = simplex_index2[i] + coordinates[n] = coordinates[i] + n += int(1) + + # SHOULD NOT OCCUR + if n < 1: + break + + # get the next iteration of x_k + x_next = _linear_combine(n, coordinates, simplex) + + # x_k has converged to minimum + if _almost_equal(x_next, x_k): + break + + # copy next iteration into x_k + x_k = x_next + + # we have a tetrahedron containing the origin so return early + if n == 4: + break + + result = GJKResult() + + # compute the approximate witness points + result.x1 = _linear_combine(n, coordinates, simplex1) + result.x2 = _linear_combine(n, coordinates, simplex2) + result.dist = wp.norm_l2(x_k) + + result.dim = n + result.simplex1 = simplex1 + result.simplex2 = simplex2 + result.simplex_index1 = simplex_index1 + result.simplex_index2 = simplex_index2 + result.simplex = simplex + return result + + +@wp.func +def _same_side(p0: wp.vec3, p1: wp.vec3, p2: wp.vec3, p3: wp.vec3): + n = wp.cross(p1 - p0, p2 - p0) + dot1 = wp.dot(n, p3 - p0) + dot2 = wp.dot(n, -p0) + if dot1 > 0 and dot2 > 0: + return 1 + if dot1 < 0 and dot2 < 0: + return 1 + return 0 + + +@wp.func +def _test_tetra(p0: wp.vec3, p1: wp.vec3, p2: wp.vec3, p3: wp.vec3): + return _same_side(p0, p1, p2, p3) and _same_side(p1, p2, p3, p0) and _same_side(p2, p3, p0, p1) and _same_side(p3, p0, p1, p2) + + +@wp.func +def _tri_affine_coord(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, p: wp.vec3): + # compute minors as in S2D + M_14 = v2[1] * v3[2] - v2[2] * v3[1] - v1[1] * v3[2] + v1[2] * v3[1] + v1[1] * v2[2] - v1[2] * v2[1] + M_24 = v2[0] * v3[2] - v2[2] * v3[0] - v1[0] * v3[2] + v1[2] * v3[0] + v1[0] * v2[2] - v1[2] * v2[0] + M_34 = v2[0] * v3[1] - v2[1] * v3[0] - v1[0] * v3[1] + v1[1] * v3[0] + v1[0] * v2[1] - v1[1] * v2[0] + + # exclude one of the axes with the largest projection + # of the simplex using the computed minors + M_max = 0.0 + x = 0 + y = 0 + + mu1 = wp.abs(M_14) + mu2 = wp.abs(M_24) + mu3 = wp.abs(M_34) + + if mu1 >= mu2 and mu1 >= mu3: + M_max = M_14 + x = 1 + y = 2 + elif mu2 >= mu3: + M_max = M_24 + x = 0 + y = 2 + else: + M_max = M_34 + x = 0 + y = 1 + + # C31 corresponds to the signed area of 2-simplex: (v, s2, s3) + C31 = p[x] * v2[y] + p[y] * v3[x] + v2[x] * v3[y] - p[x] * v3[y] - p[y] * v2[x] - v3[x] * v2[y] + + # C32 corresponds to the signed area of 2-simplex: (v, s1, s3) + C32 = p[x] * v3[y] + p[y] * v1[x] + v3[x] * v1[y] - p[x] * v1[y] - p[y] * v3[x] - v1[x] * v3[y] + + # C33 corresponds to the signed area of 2-simplex: (v, s1, s2) + C33 = p[x] * v1[y] + p[y] * v2[x] + v1[x] * v2[y] - p[x] * v2[y] - p[y] * v1[x] - v2[x] * v1[y] + + # compute affine coordinates + return wp.vec3(C31 / M_max, C32 / M_max, C33 / M_max) + + +@wp.func +def _tri_point_intersect(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, p: wp.vec3): + coordinates = _tri_affine_coord(v1, v2, v3, p) + l1 = coordinates[0] + l2 = coordinates[1] + l3 = coordinates[2] + + if l1 < 0 or l2 < 0 or l3 < 0: + return False + + pr = wp.vec3() + pr[0] = v1[0] * l1 + v2[0] * l2 + v3[0] * l3 + pr[1] = v1[1] * l1 + v2[1] * l2 + v3[1] * l3 + pr[2] = v1[2] * l1 + v2[2] * l2 + v3[2] * l3 + return wp.norm_l2(pr - p) < MJ_MINVAL + + +@wp.func +def _replace_simplex3(pt: Polytope, v1: int, v2: int, v3: int): + result = GJKResult() + + # reset GJK simplex + simplex = mat43() + simplex[0] = pt.vert[v1] + simplex[1] = pt.vert[v2] + simplex[2] = pt.vert[v3] + + simplex1 = mat43() + simplex1[0] = pt.vert1[v1] + simplex1[1] = pt.vert1[v2] + simplex1[2] = pt.vert1[v3] + + simplex2 = mat43() + simplex2[0] = pt.vert2[v1] + simplex2[1] = pt.vert2[v2] + simplex2[2] = pt.vert2[v3] + + simplex_index1 = wp.vec4i() + simplex_index1[0] = pt.vert_index1[v1] + simplex_index1[1] = pt.vert_index1[v2] + simplex_index1[2] = pt.vert_index1[v3] + + simplex_index2 = wp.vec4i() + simplex_index2[0] = pt.vert_index2[v1] + simplex_index2[1] = pt.vert_index2[v2] + simplex_index2[2] = pt.vert_index2[v3] + + result.simplex = simplex + result.simplex1 = simplex1 + result.simplex2 = simplex2 + result.simplex_index1 = simplex_index1 + result.simplex_index2 = simplex_index2 + + return result + + +@wp.func +def _rotmat(axis: wp.vec3): + n = wp.norm_l2(axis) + u1 = axis[0] / n + u2 = axis[1] / n + u3 = axis[2] / n + + sin = 0.86602540378 # sin(120 deg) + cos = -0.5 # cos(120 deg) + R = wp.mat33() + R[0, 0] = cos + u1 * u1 * (1.0 - cos) + R[0, 1] = u1 * u2 * (1.0 - cos) - u3 * sin + R[0, 2] = u1 * u3 * (1.0 - cos) + u2 * sin + R[1, 0] = u2 * u1 * (1.0 - cos) + u3 * sin + R[1, 1] = cos + u2 * u2 * (1.0 - cos) + R[1, 2] = u2 * u3 * (1.0 - cos) - u1 * sin + R[2, 0] = u1 * u3 * (1.0 - cos) - u2 * sin + R[2, 1] = u2 * u3 * (1.0 - cos) + u1 * sin + R[2, 2] = cos + u3 * u3 * (1.0 - cos) + return R + + +@wp.func +def _ray_triangle(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, v4: wp.vec3, v5: wp.vec3): + vol1 = _det3(v3 - v1, v4 - v1, v2 - v1) + vol2 = _det3(v4 - v1, v5 - v1, v2 - v1) + vol3 = _det3(v5 - v1, v3 - v1, v2 - v1) + + if vol1 >= 0 and vol2 >= 0 and vol3 >= 0: + return 1 + if vol1 <= 0 and vol2 <= 0 and vol3 <= 0: + return -1 + return 0 + + +@wp.func +def _add_edge(pt: Polytope, e1: int, e2: int): + n = pt.nhorizon + + if n < 0: + return -1 + + for i in range(n): + old_e1 = pt.horizon[2 * i + 0] + old_e2 = pt.horizon[2 * i + 1] + if (old_e1 == e1 and old_e2 == e2) or (old_e1 == e2 and old_e2 == e1): + pt.horizon[2 * i + 0] = pt.horizon[2 * (n - 1) + 0] + pt.horizon[2 * i + 1] = pt.horizon[2 * (n - 1) + 1] + return n - 1 + + # out of memory, force EPA to return early without contact + if n > pt.horizon.shape[0] - 2: + return -1 + + pt.horizon[2 * n + 0] = e1 + pt.horizon[2 * n + 1] = e2 + return n + 1 + + +@wp.func +def _delete_face(pt: Polytope, face_id: int): + index = pt.face_index[face_id] + # delete from map + if index >= 0: + last_face = pt.face_map[pt.nmap - 1] + pt.face_map[index] = last_face + pt.face_index[last_face] = index + pt.nmap -= 1 + # mark face as deleted from polytope + pt.face_index[face_id] = -2 + return pt.nmap + + +@wp.func +def _epa_witness(pt: Polytope, face_idx: int): + # compute affine coordinates for witness points on plane defined by face + v1 = pt.vert[pt.face[face_idx][0]] + v2 = pt.vert[pt.face[face_idx][1]] + v3 = pt.vert[pt.face[face_idx][2]] + + coordinates = _tri_affine_coord(v1, v2, v3, pt.face_pr[face_idx]) + l1 = coordinates[0] + l2 = coordinates[1] + l3 = coordinates[2] + + # face on geom 1 + v1 = pt.vert1[pt.face[face_idx][0]] + v2 = pt.vert1[pt.face[face_idx][1]] + v3 = pt.vert1[pt.face[face_idx][2]] + x1 = wp.vec3() + x1[0] = v1[0] * l1 + v2[0] * l2 + v3[0] * l3 + x1[1] = v1[1] * l1 + v2[1] * l2 + v3[1] * l3 + x1[2] = v1[2] * l1 + v2[2] * l2 + v3[2] * l3 + + # face on geom 2 + v1 = pt.vert2[pt.face[face_idx][0]] + v2 = pt.vert2[pt.face[face_idx][1]] + v3 = pt.vert2[pt.face[face_idx][2]] + x2 = wp.vec3() + x2[0] = v1[0] * l1 + v2[0] * l2 + v3[0] * l3 + x2[1] = v1[1] * l1 + v2[1] * l2 + v3[1] * l3 + x2[2] = v1[2] * l1 + v2[2] * l2 + v3[2] * l3 + + return x1, x2 + + +@wp.func +def _polytope2( + # In: + pt: Polytope, + dist: float, + simplex: mat43, + simplex1: mat43, + simplex2: mat43, + simplex_index1: wp.vec4i, + simplex_index2: wp.vec4i, + geom1: Geom, + geom2: Geom, + geomtype1: int, + geomtype2: int, +): + """Create polytope for EPA given a 1-simplex from GJK""" + diff = simplex[1] - simplex[0] + + # find component with smallest magnitude (so cross product is largest) + value = FLOAT_MAX + index = 0 + for i in range(3): + if wp.abs(diff[i]) < value: + value = wp.abs(diff[i]) + index = i + + # cross product with best coordinate axis + e = wp.vec(0.0, 0.0, 0.0) + e[index] = 1.0 + d1 = wp.cross(e, diff) + + # rotate around the line segment to get three more points spaced 120 degrees apart + R = _rotmat(diff) + d2 = R @ d1 + d3 = R @ d2 + + # save vertices and get indices for each one + pt.vert[0] = simplex[0] + pt.vert[1] = simplex[1] + + pt.vert1[0] = simplex1[0] + pt.vert1[1] = simplex1[1] + + pt.vert_index1[0] = simplex_index1[0] + pt.vert_index1[1] = simplex_index1[1] + + pt.vert2[0] = simplex2[0] + pt.vert2[1] = simplex2[1] + + pt.vert_index2[0] = simplex_index2[0] + pt.vert_index2[1] = simplex_index2[1] + + _epa_support(pt, 2, geom1, geom2, geomtype1, geomtype2, d1 / wp.norm_l2(d1)) + _epa_support(pt, 3, geom1, geom2, geomtype1, geomtype2, d2 / wp.norm_l2(d2)) + _epa_support(pt, 4, geom1, geom2, geomtype1, geomtype2, d3 / wp.norm_l2(d3)) + + # build hexahedron + if _attach_face(pt, 0, 0, 2, 3) < MJ_MINVAL: + pt.status = -1 + return pt, _replace_simplex3(pt, 0, 2, 3) + + if _attach_face(pt, 1, 0, 4, 2) < MJ_MINVAL2: + pt.status = -1 + return pt, _replace_simplex3(pt, 0, 4, 2) + + if _attach_face(pt, 2, 0, 3, 4) < MJ_MINVAL2: + pt.status = -1 + return pt, _replace_simplex3(pt, 0, 3, 4) + + if _attach_face(pt, 3, 1, 3, 2) < MJ_MINVAL2: + pt.status = -1 + return pt, _replace_simplex3(pt, 1, 3, 2) + + if _attach_face(pt, 4, 1, 2, 4) < MJ_MINVAL2: + pt.status = -1 + return pt, _replace_simplex3(pt, 1, 2, 4) + + if _attach_face(pt, 5, 1, 4, 3) < MJ_MINVAL2: + pt.status = -1 + return pt, _replace_simplex3(pt, 1, 4, 3) + + # check hexahedron is convex + if not _ray_triangle(simplex[0], simplex[1], pt.vert[2], pt.vert[3], pt.vert[4]): + pt.status = 1 + return pt, GJKResult() + + # populate face map + for i in range(6): + pt.face_map[i] = i + pt.face_index[i] = i + + # set polytope counts + pt.nvert = 5 + pt.nface = 6 + pt.nmap = 6 + pt.status = 0 + return pt, GJKResult() + + +@wp.func +def _polytope3( + # In: + pt: Polytope, + dist: float, + simplex: mat43, + simplex1: mat43, + simplex2: mat43, + simplex_index1: wp.vec4i, + simplex_index2: wp.vec4i, + geom1: Geom, + geom2: Geom, + geomtype1: int, + geomtype2: int, +): + """Create polytope for EPA given a 2-simplex from GJK""" + # get normals in both directions + n = wp.cross(simplex[1] - simplex[0], simplex[2] - simplex[0]) + if wp.norm_l2(n) < MJ_MINVAL: + pt.status = 2 + return pt + + pt.vert[0] = simplex[0] + pt.vert[1] = simplex[1] + pt.vert[2] = simplex[2] + + pt.vert1[0] = simplex1[0] + pt.vert1[1] = simplex1[1] + pt.vert1[2] = simplex1[2] + + pt.vert_index1[0] = simplex_index1[0] + pt.vert_index1[1] = simplex_index1[1] + pt.vert_index1[2] = simplex_index1[2] + + pt.vert2[0] = simplex2[0] + pt.vert2[1] = simplex2[1] + pt.vert2[2] = simplex2[2] + + pt.vert_index2[0] = simplex_index2[0] + pt.vert_index2[1] = simplex_index2[1] + pt.vert_index2[2] = simplex_index2[2] + + _epa_support(pt, 3, geom1, geom2, geomtype1, geomtype2, -n) + _epa_support(pt, 4, geom1, geom2, geomtype1, geomtype2, n) + + v1 = simplex[0] + v2 = simplex[1] + v3 = simplex[2] + v4 = pt.vert[3] + v5 = pt.vert[4] + + # check that v4 is not contained in the 2-simplex + if _tri_point_intersect(v1, v2, v3, v4): + pt.status = 3 + return pt + + # check that v5 is not contained in the 2-simplex + if _tri_point_intersect(v1, v2, v3, v5): + pt.status = 4 + return pt + + # if origin does not lie on simplex then we need to check that the hexahedron contains the + # origin + if dist > 1e-5 and not _test_tetra(v1, v2, v3, v4) and not _test_tetra(v1, v2, v3, v5): + pt.status = 5 + return pt + + # create hexahedron for EPA + if _attach_face(pt, 0, 4, 0, 1) < MJ_MINVAL2: + pt.status = 6 + return pt + if _attach_face(pt, 1, 4, 2, 0) < MJ_MINVAL2: + pt.status = 7 + return pt + if _attach_face(pt, 2, 4, 1, 2) < MJ_MINVAL2: + pt.status = 8 + return pt + if _attach_face(pt, 3, 3, 1, 0) < MJ_MINVAL2: + pt.status = 9 + return pt + if _attach_face(pt, 4, 3, 0, 2) < MJ_MINVAL2: + pt.status = 10 + return pt + if _attach_face(pt, 5, 3, 2, 1) < MJ_MINVAL2: + pt.status = 11 + return pt + + # populate face map + for i in range(6): + pt.face_map[i] = i + pt.face_index[i] = i + + # set polytope counts + pt.nvert = 5 + pt.nface = 6 + pt.nmap = 6 + pt.status = 0 + return pt + + +@wp.func +def _polytope4( + # In: + pt: Polytope, + dist: float, + simplex: mat43, + simplex1: mat43, + simplex2: mat43, + simplex_index1: wp.vec4i, + simplex_index2: wp.vec4i, + geom1: Geom, + geom2: Geom, + geomtype1: int, + geomtype2: int, +): + """Create polytope for EPA given a 3-simplex from GJK""" + pt.vert[0] = simplex[0] + pt.vert[1] = simplex[1] + pt.vert[2] = simplex[2] + pt.vert[3] = simplex[3] + + pt.vert1[0] = simplex1[0] + pt.vert1[1] = simplex1[1] + pt.vert1[2] = simplex1[2] + pt.vert1[3] = simplex1[3] + + pt.vert_index1[0] = simplex_index1[0] + pt.vert_index1[1] = simplex_index1[1] + pt.vert_index1[2] = simplex_index1[2] + pt.vert_index1[3] = simplex_index1[3] + + pt.vert2[0] = simplex2[0] + pt.vert2[1] = simplex2[1] + pt.vert2[2] = simplex2[2] + pt.vert2[3] = simplex2[3] + + pt.vert_index2[0] = simplex_index2[0] + pt.vert_index2[1] = simplex_index2[1] + pt.vert_index2[2] = simplex_index2[2] + pt.vert_index2[3] = simplex_index2[3] + + # if the origin is on a face, replace the 3-simplex with a 2-simplex + if _attach_face(pt, 0, 0, 1, 2) < MJ_MINVAL2: + pt.status = -1 + return pt, _replace_simplex3(pt, 0, 1, 2) + + if _attach_face(pt, 1, 0, 3, 1) < MJ_MINVAL2: + pt.status = -1 + return pt, _replace_simplex3(pt, 0, 3, 1) + + if _attach_face(pt, 2, 0, 2, 3) < MJ_MINVAL2: + pt.status = -1 + return pt, _replace_simplex3(pt, 0, 2, 3) + + if _attach_face(pt, 3, 3, 2, 1) < MJ_MINVAL2: + pt.status = -1 + return pt, _replace_simplex3(pt, 3, 2, 1) + + if not _test_tetra(pt.vert[0], pt.vert[1], pt.vert[2], pt.vert[3]): + pt.status = 12 + return pt, GJKResult() + + # populate face map + for i in range(4): + pt.face_map[i] = i + pt.face_index[i] = i + + # set polytope counts + pt.nvert = 4 + pt.nface = 4 + pt.nmap = 4 + pt.status = 0 + return pt, GJKResult() + + +@wp.func +def _epa(tolerance2: float, epa_iterations: int, pt: Polytope, geom1: Geom, geom2: Geom, geomtype1: int, geomtype2: int): + """Recover penetration data from two geoms in contact given an initial polytope.""" + upper = FLOAT_MAX + upper2 = FLOAT_MAX + idx = int(-1) + pidx = int(-1) + + for k in range(epa_iterations): + pidx = int(idx) + idx = int(-1) + + # find the face closest to the origin (lower bound for penetration depth) + lower2 = float(FLOAT_MAX) + for i in range(pt.nmap): + face_idx = pt.face_map[i] + if pt.face_norm2[face_idx] < lower2: + idx = int(face_idx) + lower2 = float(pt.face_norm2[face_idx]) + + # face not valid, return previous face + if lower2 > upper2 or idx < 0: + idx = pidx + break + + # check if lower bound is 0 + if lower2 <= 0: + break + + # compute support point w from the closest face's normal + lower = wp.sqrt(lower2) + wi = pt.nvert + i1, i2 = _epa_support(pt, wi, geom1, geom2, geomtype1, geomtype2, pt.face_pr[idx] / lower) + geom1.index = i1 + geom2.index = i2 + pt.nvert += 1 + + # upper bound for kth iteration + upper_k = wp.dot(pt.face_pr[idx], pt.vert[wi]) / lower + if upper_k < upper: + upper = upper_k + upper2 = upper * upper + + if upper - lower < tolerance2: + break + + pt.nmap = _delete_face(pt, idx) + pt.nhorizon = _add_edge(pt, pt.face[idx][0], pt.face[idx][1]) + pt.nhorizon = _add_edge(pt, pt.face[idx][1], pt.face[idx][2]) + pt.nhorizon = _add_edge(pt, pt.face[idx][2], pt.face[idx][0]) + if pt.nhorizon == -1: + idx = -1 + break + + # compute horizon for w + for i in range(pt.nface): + if pt.face_index[i] == -2: + continue + + if wp.dot(pt.face_pr[i], pt.vert[wi]) - pt.face_norm2[i] > MJ_MINVAL: + pt.nmap = _delete_face(pt, i) + pt.nhorizon = _add_edge(pt, pt.face[i][0], pt.face[i][1]) + pt.nhorizon = _add_edge(pt, pt.face[i][1], pt.face[i][2]) + pt.nhorizon = _add_edge(pt, pt.face[i][2], pt.face[i][0]) + if pt.nhorizon == -1: + idx = -1 + break + + # insert w as new vertex and attach faces along the horizon + for i in range(pt.nhorizon): + dist2 = _attach_face(pt, pt.nface, wi, pt.horizon[2 * i + 0], pt.horizon[2 * i + 1]) + if dist2 == 0: + idx = -1 + break + + pt.nface += 1 + + # store face in map + if dist2 >= lower2 and dist2 <= upper2: + pt.face_map[pt.nmap] = pt.nface - 1 + pt.face_index[pt.nface - 1] = pt.nmap + pt.nmap += 1 + + # no face candidates left + if pt.nmap == 0 or idx == -1: + break + + # clear horizon + pt.nhorizon = 0 + + # return from valid face + if idx > -1: + x1, x2 = _epa_witness(pt, idx) + return -wp.sqrt(pt.face_norm2[idx]), x1, x2 + return 0.0, wp.vec3(), wp.vec3() + + +@wp.func +def ccd( + # In: + tolerance: float, + cutoff: float, + gjk_iterations: int, + epa_iterations: int, + geom1: Geom, + geom2: Geom, + geomtype1: int, + geomtype2: int, + x_1: wp.vec3, + x_2: wp.vec3, + vert: wp.array(dtype=wp.vec3), + vert1: wp.array(dtype=wp.vec3), + vert2: wp.array(dtype=wp.vec3), + vert_index1: wp.array(dtype=int), + vert_index2: wp.array(dtype=int), + face: wp.array(dtype=wp.vec3i), + face_pr: wp.array(dtype=wp.vec3), + face_norm2: wp.array(dtype=float), + face_index: wp.array(dtype=int), + face_map: wp.array(dtype=int), + horizon: wp.array(dtype=int), +): + """General convex collision detection via GJK/EPA.""" + result = _gjk(tolerance, gjk_iterations, geom1, geom2, x_1, x_2, geomtype1, geomtype2, cutoff) + + # no penetration depth to recover + if result.dist > tolerance or result.dim < 2: + return result.dist, result.x1, result.x2 + + pt = Polytope() + pt.nface = 0 + pt.nmap = 0 + pt.nvert = 0 + pt.nhorizon = 0 + pt.vert = vert + pt.vert1 = vert1 + pt.vert2 = vert2 + pt.vert_index1 = vert_index1 + pt.vert_index2 = vert_index2 + pt.face = face + pt.face_pr = face_pr + pt.face_norm2 = face_norm2 + pt.face_index = face_index + pt.face_map = face_map + pt.horizon = horizon + + if result.dim == 2: + pt, new_result = _polytope2( + pt, + result.dist, + result.simplex, + result.simplex1, + result.simplex2, + result.simplex_index1, + result.simplex_index2, + geom1, + geom2, + geomtype1, + geomtype2, + ) + if pt.status == -1: + result.simplex = new_result.simplex + result.simplex1 = new_result.simplex1 + result.simplex2 = new_result.simplex2 + result.simplex_index1 = new_result.simplex_index1 + result.simplex_index2 = new_result.simplex_index2 + result.dim = 3 + elif result.dim == 4: + pt, new_result = _polytope4( + pt, + result.dist, + result.simplex, + result.simplex1, + result.simplex2, + result.simplex_index1, + result.simplex_index2, + geom1, + geom2, + geomtype1, + geomtype2, + ) + if pt.status == -1: + result.simplex = new_result.simplex + result.simplex1 = new_result.simplex1 + result.simplex2 = new_result.simplex2 + result.simplex_index1 = new_result.simplex_index1 + result.simplex_index2 = new_result.simplex_index2 + result.dim = 3 + + # polytope2 and polytope4 may need to fallback here + if result.dim == 3: + pt = _polytope3( + pt, + result.dist, + result.simplex, + result.simplex1, + result.simplex2, + result.simplex_index1, + result.simplex_index2, + geom1, + geom2, + geomtype1, + geomtype2, + ) + + # origin on boundary (objects are not considered penetrating) + if pt.status: + return result.dist, result.x1, result.x2 + + return _epa(tolerance * tolerance, epa_iterations, pt, geom1, geom2, geomtype1, geomtype2) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py new file mode 100644 index 00000000..e8313be6 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py @@ -0,0 +1,714 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_prism_vertex +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import Geom +from mujoco.mjx.third_party.mujoco_warp._src.math import gjk_normalize +from mujoco.mjx.third_party.mujoco_warp._src.math import orthonormal +from mujoco.mjx.third_party.mujoco_warp._src.support import all_same +from mujoco.mjx.third_party.mujoco_warp._src.support import any_different +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType + +# TODO(team): improve compile time to enable backward pass +wp.config.enable_backward = False + +FLOAT_MIN = -1e30 +FLOAT_MAX = 1e30 +EPS_BEST_COUNT = 12 +MULTI_CONTACT_COUNT = 4 +MULTI_POLYGON_COUNT = 8 + +matc3 = wp.types.matrix(shape=(EPS_BEST_COUNT, 3), dtype=float) +vecc3 = wp.types.vector(EPS_BEST_COUNT * 3, dtype=float) + +# Matrix definition for the `tris` scratch space which is used to store the +# triangles of the polytope. Note that the first dimension is 2, as we need +# to store the previous and current polytope. But since Warp doesn't support +# 3D matrices yet, we use 2 * 3 * EPS_BEST_COUNT as the first dimension. +TRIS_DIM = 3 * EPS_BEST_COUNT +mat2c3 = wp.types.matrix(shape=(2 * TRIS_DIM, 3), dtype=float) +mat3p = wp.types.matrix(shape=(MULTI_POLYGON_COUNT, 3), dtype=float) +mat3c = wp.types.matrix(shape=(MULTI_CONTACT_COUNT, 3), dtype=float) +mat43 = wp.types.matrix(shape=(4, 3), dtype=float) + +vec6 = wp.types.vector(6, dtype=int) +VECI1 = vec6(0, 0, 0, 1, 1, 2) +VECI2 = vec6(1, 2, 3, 2, 3, 3) + + +@wp.func +def _gjk_support_geom(geom: Geom, geomtype: int, dir: wp.vec3): + local_dir = wp.transpose(geom.rot) @ dir + if geomtype == int(GeomType.SPHERE.value): + support_pt = geom.pos + geom.size[0] * dir + elif geomtype == int(GeomType.BOX.value): + res = wp.cw_mul(wp.sign(local_dir), geom.size) + support_pt = geom.rot @ res + geom.pos + elif geomtype == int(GeomType.CAPSULE.value): + res = local_dir * geom.size[0] + # add cylinder contribution + res[2] += wp.sign(local_dir[2]) * geom.size[1] + support_pt = geom.rot @ res + geom.pos + elif geomtype == int(GeomType.ELLIPSOID.value): + res = wp.cw_mul(local_dir, geom.size) + res = wp.normalize(res) + # transform to ellipsoid + res = wp.cw_mul(res, geom.size) + support_pt = geom.rot @ res + geom.pos + elif geomtype == int(GeomType.CYLINDER.value): + res = wp.vec3(0.0, 0.0, 0.0) + # set result in XY plane: support on circle + d = wp.sqrt(wp.dot(local_dir, local_dir)) + if d > MJ_MINVAL: + scl = geom.size[0] / d + res[0] = local_dir[0] * scl + res[1] = local_dir[1] * scl + # set result in Z direction + res[2] = wp.sign(local_dir[2]) * geom.size[1] + support_pt = geom.rot @ res + geom.pos + elif geomtype == int(GeomType.MESH.value): + max_dist = float(FLOAT_MIN) + if geom.graphadr == -1 or geom.vertnum < 10: + # exhaustive search over all vertices + for i in range(geom.vertnum): + vert = geom.vert[geom.vertadr + i] + dist = wp.dot(vert, local_dir) + if dist > max_dist: + max_dist = dist + support_pt = vert + else: + numvert = geom.graph[geom.graphadr] + vert_edgeadr = geom.graphadr + 2 + vert_globalid = geom.graphadr + 2 + numvert + edge_localid = geom.graphadr + 2 + 2 * numvert + # hillclimb until no change + prev = int(-1) + imax = int(0) + + while True: + prev = int(imax) + i = int(geom.graph[vert_edgeadr + imax]) + while geom.graph[edge_localid + i] >= 0: + subidx = geom.graph[edge_localid + i] + idx = geom.graph[vert_globalid + subidx] + dist = wp.dot(local_dir, geom.vert[geom.vertadr + idx]) + if dist > max_dist: + max_dist = dist + imax = int(subidx) + i += int(1) + if imax == prev: + break + imax = geom.graph[vert_globalid + imax] + support_pt = geom.vert[geom.vertadr + imax] + + support_pt = geom.rot @ support_pt + geom.pos + elif geomtype == int(GeomType.HFIELD.value): + max_dist = float(FLOAT_MIN) + for i in range(6): + vert = hfield_prism_vertex(geom.hfprism, i) + dist = wp.dot(vert, local_dir) + if dist > max_dist: + max_dist = dist + support_pt = vert + support_pt = geom.rot @ support_pt + geom.pos + + return wp.dot(support_pt, dir), support_pt + + +@wp.func +def _gjk_support( + # In: + geom1: Geom, + geom2: Geom, + geomtype1: int, + geomtype2: int, + dir: wp.vec3, +): + # Returns the distance between support points on two geoms, and the support point. + # Negative distance means objects are not intersecting along direction `dir`. + # Positive distance means objects are intersecting along the given direction `dir`. + + dist1, s1 = _gjk_support_geom(geom1, geomtype1, dir) + dist2, s2 = _gjk_support_geom(geom2, geomtype2, -dir) + + support_pt = s1 - s2 + return dist1 + dist2, support_pt + + +@wp.func +def _expand_polytope(count: int, prev_count: int, dists: vecc3, tris: mat2c3, p: matc3): + # expand polytope greedily + for j in range(count): + best = int(0) + dd = dists[0] + for i in range(1, 3 * prev_count): + if dists[i] < dd: + dd = dists[i] + best = i + + dists[best] = float(wp.static(2 * FLOAT_MAX)) + + parent_index = best // 3 + child_index = best % 3 + + # fill in the new triangle at the next index + tris[TRIS_DIM + j * 3 + 0] = tris[parent_index * 3 + child_index] + tris[TRIS_DIM + j * 3 + 1] = tris[parent_index * 3 + ((child_index + 1) % 3)] + tris[TRIS_DIM + j * 3 + 2] = p[parent_index] + + for r in range(wp.static(EPS_BEST_COUNT * 3)): + # swap triangles + swap = tris[TRIS_DIM + r] + tris[TRIS_DIM + r] = tris[r] + tris[r] = swap + + return dists, tris + + +@wp.func +def gjk_legacy( + # In: + gjk_iterations: int, + geom1: Geom, + geom2: Geom, + geomtype1: int, + geomtype2: int, +): + dir = wp.vec3(0.0, 0.0, 1.0) + dir_n = -dir + depth = float(FLOAT_MAX) + + dist_max, simplex0 = _gjk_support(geom1, geom2, geomtype1, geomtype2, dir) + dist_min, simplex1 = _gjk_support(geom1, geom2, geomtype1, geomtype2, dir_n) + + if dist_max < dist_min: + depth = dist_max + normal = dir + else: + depth = dist_min + normal = dir_n + + sd = simplex0 - simplex1 + dir = orthonormal(sd) + + dist_max, simplex3 = _gjk_support(geom1, geom2, geomtype1, geomtype2, dir) + + # Initialize a 2-simplex with simplex[2]==simplex[1]. This ensures the + # correct winding order for face normals defined below. Face 0 and face 3 + # are degenerate, and face 1 and 2 have opposing normals. + simplex = mat43() + simplex[0] = simplex0 + simplex[1] = simplex1 + simplex[2] = simplex[1] + simplex[3] = simplex3 + + if dist_max < depth: + depth = dist_max + normal = dir + if dist_min < depth: + depth = dist_min + normal = dir_n + + plane = mat43() + for _ in range(gjk_iterations): + # winding orders: plane[0] ccw, plane[1] cw, plane[2] ccw, plane[3] cw + plane[0] = wp.cross(simplex[3] - simplex[2], simplex[1] - simplex[2]) + plane[1] = wp.cross(simplex[3] - simplex[0], simplex[2] - simplex[0]) + plane[2] = wp.cross(simplex[3] - simplex[1], simplex[0] - simplex[1]) + plane[3] = wp.cross(simplex[2] - simplex[0], simplex[1] - simplex[0]) + + # Compute distance of each face halfspace to the origin. If dplane<0, then the + # origin is outside the halfspace. If dplane>0 then the origin is inside + # the halfspace defined by the face plane. + + dplane = wp.vec4(float(FLOAT_MAX)) + + plane0, p0 = gjk_normalize(plane[0]) + plane1, p1 = gjk_normalize(plane[1]) + plane2, p2 = gjk_normalize(plane[2]) + plane3, p3 = gjk_normalize(plane[3]) + + plane[0] = plane0 + plane[1] = plane1 + plane[2] = plane2 + plane[3] = plane3 + + if p0: + dplane[0] = wp.dot(plane[0], simplex[2]) + + if p1: + dplane[1] = wp.dot(plane[1], simplex[0]) + + if p2: + dplane[2] = wp.dot(plane[2], simplex[1]) + + if p3: + dplane[3] = wp.dot(plane[3], simplex[0]) + + # pick plane normal with minimum distance to the origin + i1 = wp.where(dplane[0] < dplane[1], 0, 1) + i2 = wp.where(dplane[2] < dplane[3], 2, 3) + index = wp.where(dplane[i1] < dplane[i2], i1, i2) + + if dplane[index] > 0.0: + # origin is inside the simplex, objects are intersecting + break + + # add new support point to the simplex + dist, simplex_i = _gjk_support(geom1, geom2, geomtype1, geomtype2, plane[index]) + simplex[index] = simplex_i + + if dist < depth: + depth = dist + normal = plane[index] + + # preserve winding order of the simplex faces + index1 = (index + 1) & 3 + index2 = (index + 2) & 3 + swap = simplex[index1] + simplex[index1] = simplex[index2] + simplex[index2] = swap + + if dist < 0.0: + break # objects are likely non-intersecting + + return simplex, normal + + +@wp.func +def epa_legacy( + # In: + epa_iterations: int, + geom1: Geom, + geom2: Geom, + geomtype1: int, + geomtype2: int, + depth_extension: float, + epa_exact_neg_distance: bool, + simplex: mat43, + normal: wp.vec3, +): + # get the support, if depth < 0: objects do not intersect + depth, _ = _gjk_support(geom1, geom2, geomtype1, geomtype2, normal) + + if depth < -depth_extension: + # Objects are not intersecting, and we do not obtain the closest points as + # specified by depth_extension. + return wp.nan, wp.vec3(wp.nan, wp.nan, wp.nan) + + if wp.static(epa_exact_neg_distance): + # Check closest points to all edges of the simplex, rather than just the + # face normals. This gives the exact depth/normal for the non-intersecting + # case. + for i in range(6): + i1 = VECI1[i] + i2 = VECI2[i] + + si1 = simplex[i1] + si2 = simplex[i2] + + if si1[0] != si2[0] or si1[1] != si2[1] or si1[2] != si2[2]: + v = si1 - si2 + alpha = wp.dot(si1, v) / wp.dot(v, v) + + # p0 is the closest segment point to the origin + p0 = wp.clamp(alpha, 0.0, 1.0) * v - si1 + p0, pf = gjk_normalize(p0) + + if pf: + depth2, _ = _gjk_support(geom1, geom2, geomtype1, geomtype2, p0) + + if depth2 < depth: + depth = depth2 + normal = p0 + + # supporting points for each triangle + p = matc3() + + # distance to the origin for candidate triangles + dists = vecc3() + + tris = mat2c3() + tris[0] = simplex[2] + tris[1] = simplex[1] + tris[2] = simplex[3] + + tris[3] = simplex[0] + tris[4] = simplex[2] + tris[5] = simplex[3] + + tris[6] = simplex[1] + tris[7] = simplex[0] + tris[8] = simplex[3] + + tris[9] = simplex[0] + tris[10] = simplex[1] + tris[11] = simplex[2] + + # Calculate the total number of iterations to avoid nested loop + # This is a hack to reduce compile time + count = int(4) + it = int(0) + for _ in range(wp.static(epa_iterations)): + it += count + count = wp.min(count * 3, EPS_BEST_COUNT) + + count = int(4) + i = int(0) + for _ in range(it): + # Loop through all triangles, and obtain distances to the origin for each + # new triangle candidate. + ti = 3 * i + n = wp.cross(tris[ti + 2] - tris[ti + 0], tris[ti + 1] - tris[ti + 0]) + + n, nf = gjk_normalize(n) + if not nf: + for j in range(3): + dists[i * 3 + j] = wp.static(float(2 * FLOAT_MAX)) + continue + + dist, pi = _gjk_support(geom1, geom2, geomtype1, geomtype2, n) + p[i] = pi + + if dist < depth: + depth = dist + normal = n + + # iterate over edges and get distance using support point + for j in range(3): + if wp.static(epa_exact_neg_distance): + # obtain closest point between new triangle edge and origin + tqj = tris[ti + j] + + if (p[i, 0] != tqj[0]) or (p[i, 1] != tqj[1]) or (p[i, 2] != tqj[2]): + v = p[i] - tris[ti + j] + alpha = wp.dot(p[i], v) / wp.dot(v, v) + p0 = wp.clamp(alpha, 0.0, 1.0) * v - p[i] + p0, pf = gjk_normalize(p0) + + if pf: + dist2, v = _gjk_support(geom1, geom2, geomtype1, geomtype2, p0) + + if dist2 < depth: + depth = dist2 + normal = p0 + + plane = wp.cross(p[i] - tris[ti + j], tris[ti + ((j + 1) % 3)] - tris[ti + j]) + plane, pf = gjk_normalize(plane) + + if pf: + dd = wp.dot(plane, tris[ti + j]) + else: + dd = float(FLOAT_MAX) + + if (dd < 0 and depth >= 0) or ( + tris[ti + ((j + 2) % 3)][0] == p[i][0] + and tris[ti + ((j + 2) % 3)][1] == p[i][1] + and tris[ti + ((j + 2) % 3)][2] == p[i][2] + ): + dists[i * 3 + j] = float(FLOAT_MAX) + else: + dists[i * 3 + j] = dd + + if i == count - 1: + prev_count = count + count = wp.min(count * 3, EPS_BEST_COUNT) + dists, tris = _expand_polytope(count, prev_count, dists, tris, p) + i = int(0) + else: + i += 1 + + return depth, normal + + +@wp.func +def multicontact_legacy( + # In: + geom1: Geom, + geom2: Geom, + geomtype1: int, + geomtype2: int, + depth_extension: float, + depth: float, + normal: wp.vec3, + ncontact: int, + npolygon: int, + perturbation_angle: float, +): + # Calculates multiple contact points given the normal from EPA. + # 1. Calculates the polygon on each shape by tiling the normal + # "perturbation_angle" (radians) in the orthogonal component of the normal. + # The "perturbation_angle" can be changed to depend on the depth of the + # contact, in a future version. + # 2. The normal is tilted "npolygon" times in the directions evenly + # spaced in the orthogonal component of the normal. + # (works well for >= 6, default is 8). + # 3. The intersection between these two polygons is calculated in 2D space + # (complement to the normal). If they intersect, extreme points in both + # directions are found. This can be modified to the extremes in the + # direction of eigenvectors of the variance of points of each polygon. If + # they do not intersect, the closest points of both polygons are found. + + assert ncontact <= MULTI_CONTACT_COUNT + assert npolygon <= MULTI_POLYGON_COUNT + + if depth < -depth_extension: + return 0, mat3c() + + dir = orthonormal(normal) + dir2 = wp.cross(normal, dir) + + angle = perturbation_angle + c = wp.cos(angle) + s = wp.sin(angle) + tc = 1.0 - c + + v1 = mat3p() + v2 = mat3p() + + contact_points = mat3c() + + # Obtain points on the polygon determined by the support and tilt angle, + # in the basis of the contact frame. + v1count = int(0) + v2count = int(0) + angle_ratio = wp.static(2.0 * wp.pi) / float(npolygon) + + for i in range(npolygon): + angle = angle_ratio * float(i) + axis = wp.cos(angle) * dir + wp.sin(angle) * dir2 + + # Axis-angle rotation matrix. See + # https://en.wikipedia.org/wiki/Rotation_matrix#Rotation_matrix_from_axis_and_angle + mat0 = c + axis[0] * axis[0] * tc + mat5 = c + axis[1] * axis[1] * tc + mat10 = c + axis[2] * axis[2] * tc + t1 = axis[0] * axis[1] * tc + t2 = axis[2] * s + mat4 = t1 + t2 + mat1 = t1 - t2 + t1 = axis[0] * axis[2] * tc + t2 = axis[1] * s + mat8 = t1 - t2 + mat2 = t1 + t2 + t1 = axis[1] * axis[2] * tc + t2 = axis[0] * s + mat9 = t1 + t2 + mat6 = t1 - t2 + + n = wp.vec3( + mat0 * normal[0] + mat1 * normal[1] + mat2 * normal[2], + mat4 * normal[0] + mat5 * normal[1] + mat6 * normal[2], + mat8 * normal[0] + mat9 * normal[1] + mat10 * normal[2], + ) + + _, p = _gjk_support_geom(geom1, geomtype1, n) + v1[v1count] = wp.vec3(wp.dot(p, dir), wp.dot(p, dir2), wp.dot(p, normal)) + + if i == 0: + v1count += 1 + elif any_different(v1[v1count], v1[v1count - 1]): + v1count += 1 + + n = -n + _, p = _gjk_support_geom(geom2, geomtype2, n) + v2[v2count] = wp.vec3(wp.dot(p, dir), wp.dot(p, dir2), wp.dot(p, normal)) + + if i == 0: + v2count += 1 + elif any_different(v2[v2count], v2[v2count - 1]): + v2count += 1 + + # remove duplicate vertices on the array boundary + if v1count > 1 and all_same(v1[v1count - 1], v1[0]): + v1count -= 1 + + if v2count > 1 and all_same(v2[v2count - 1], v2[0]): + v2count -= 1 + + # find an intersecting polygon between v1 and v2 in the 2D plane + out = mat43() + candCount = int(0) + + if v2count > 1: + for i in range(v1count): + m1a = v1[i] + is_in = bool(True) + + # check if point m1a is inside the v2 polygon on the 2D plane + for j in range(v2count): + j2 = (j + 1) % v2count + + # Checks that orientation of the triangle (v2[j], v2[j2], m1a) is + # counter-clockwise. If so, point m1a is inside the v2 polygon. + is_in = is_in and ((v2[j2][0] - v2[j][0]) * (m1a[1] - v2[j][1]) - (v2[j2][1] - v2[j][1]) * (m1a[0] - v2[j][0]) >= 0.0) + + if not is_in: + break + + if is_in: + if not candCount or m1a[0] < out[0, 0]: + out[0] = m1a + if not candCount or m1a[0] > out[1, 0]: + out[1] = m1a + if not candCount or m1a[1] < out[2, 1]: + out[2] = m1a + if not candCount or m1a[1] > out[3, 1]: + out[3] = m1a + candCount += 1 + + if v1count > 1: + for i in range(v2count): + m1a = v2[i] + is_in = bool(True) + + for j in range(v1count): + j2 = (j + 1) % v1count + is_in = is_in and (v1[j2][0] - v1[j][0]) * (m1a[1] - v1[j][1]) - (v1[j2][1] - v1[j][1]) * (m1a[0] - v1[j][0]) >= 0.0 + if not is_in: + break + + if is_in: + if not candCount or m1a[0] < out[0, 0]: + out[0] = m1a + if not candCount or m1a[0] > out[1, 0]: + out[1] = m1a + if not candCount or m1a[1] < out[2, 1]: + out[2] = m1a + if not candCount or m1a[1] > out[3, 1]: + out[3] = m1a + candCount += 1 + + if v1count > 1 and v2count > 1: + # Check all edge pairs, and store line segment intersections if they are + # on the edge of the boundary. + for i in range(v1count): + for j in range(v2count): + m1a = v1[i] + m1b = v1[(i + 1) % v1count] + m2a = v2[j] + m2b = v2[(j + 1) % v2count] + + det = (m2a[1] - m2b[1]) * (m1b[0] - m1a[0]) - (m1a[1] - m1b[1]) * (m2b[0] - m2a[0]) + + if wp.abs(det) > 1e-12: + a11 = (m2a[1] - m2b[1]) / det + a12 = (m2b[0] - m2a[0]) / det + a21 = (m1a[1] - m1b[1]) / det + a22 = (m1b[0] - m1a[0]) / det + b1 = m2a[0] - m1a[0] + b2 = m2a[1] - m1a[1] + + alpha = a11 * b1 + a12 * b2 + beta = a21 * b1 + a22 * b2 + if alpha >= 0.0 and alpha <= 1.0 and beta >= 0.0 and beta <= 1.0: + m0 = wp.vec3( + m1a[0] + alpha * (m1b[0] - m1a[0]), + m1a[1] + alpha * (m1b[1] - m1a[1]), + (m1a[2] + alpha * (m1b[2] - m1a[2]) + m2a[2] + beta * (m2b[2] - m2a[2])) * 0.5, + ) + if not candCount or m0[0] < out[0, 0]: + out[0] = m0 + if not candCount or m0[0] > out[1, 0]: + out[1] = m0 + if not candCount or m0[1] < out[2, 1]: + out[2] = m0 + if not candCount or m0[1] > out[3, 1]: + out[3] = m0 + candCount += 1 + + var_rx = wp.vec3(0.0) + contact_count = int(0) + if candCount > 0: + # Polygon intersection was found. + # TODO(btaba): replace the above routine with the manifold point routine + # from MJX. Deduplicate the points properly. + last_pt = wp.vec3(FLOAT_MAX, FLOAT_MAX, FLOAT_MAX) + + for k in range(ncontact): + pt = out[k, 0] * dir + out[k, 1] * dir2 + out[k, 2] * normal + + # skip contact points that are too close + if wp.length(pt - last_pt) <= 1e-6: + continue + + contact_points[contact_count] = pt + last_pt = pt + contact_count += 1 + + else: + # Polygon intersection was not found. Loop through all vertex pairs and + # calculate an approximate contact point. + minDist = float(0.0) + for i in range(v1count): + for j in range(v2count): + # Find the closest vertex pair. Calculate a contact point var_rx as the + # midpoint between the closest vertex pair. + m1 = v1[i] + m2 = v2[j] + dd = (m1[0] - m2[0]) * (m1[0] - m2[0]) + (m1[1] - m2[1]) * (m1[1] - m2[1]) + + if i != 0 and j != 0 or dd < minDist: + minDist = dd + var_rx = ((m1[0] + m2[0]) * dir + (m1[1] + m2[1]) * dir2 + (m1[2] + m2[2]) * normal) * 0.5 + + # Check for a closer point between a point on v2 and an edge on v1. + m1b = v1[(i + 1) % v1count] + m2b = v2[(j + 1) % v2count] + + if v1count > 1: + dd = (m1b[0] - m1[0]) * (m1b[0] - m1[0]) + (m1b[1] - m1[1]) * (m1b[1] - m1[1]) + t = ((m2[1] - m1[1]) * (m1b[0] - m1[0]) - (m2[0] - m1[0]) * (m1b[1] - m1[1])) / dd + dx = m2[0] + (m1b[1] - m1[1]) * t + dy = m2[1] - (m1b[0] - m1[0]) * t + dist = (dx - m2[0]) * (dx - m2[0]) + (dy - m2[1]) * (dy - m2[1]) + + if ( + (dist < minDist) + and ((dx - m1[0]) * (m1b[0] - m1[0]) + (dy - m1[1]) * (m1b[1] - m1[1]) >= 0) + and ((dx - m1b[0]) * (m1[0] - m1b[0]) + (dy - m1b[1]) * (m1[1] - m1b[1]) >= 0) + ): + alpha = wp.sqrt(((dx - m1[0]) * (dx - m1[0]) + (dy - m1[1]) * (dy - m1[1])) / dd) + minDist = dist + w = ((1.0 - alpha) * m1 + alpha * m1b + m2) * 0.5 + var_rx = w[0] * dir + w[1] * dir2 + w[2] * normal + + # check for a closer point between a point on v1 and an edge on v2 + if v2count > 1: + dd = (m2b[0] - m2[0]) * (m2b[0] - m2[0]) + (m2b[1] - m2[1]) * (m2b[1] - m2[1]) + t = ((m1[1] - m2[1]) * (m2b[0] - m2[0]) - (m1[0] - m2[0]) * (m2b[1] - m2[1])) / dd + dx = m1[0] + (m2b[1] - m2[1]) * t + dy = m1[1] - (m2b[0] - m2[0]) * t + dist = (dx - m1[0]) * (dx - m1[0]) + (dy - m1[1]) * (dy - m1[1]) + + if ( + dist < minDist + and (dx - m2[0]) * (m2b[0] - m2[0]) + (dy - m2[1]) * (m2b[1] - m2[1]) >= 0 + and (dx - m2b[0]) * (m2[0] - m2b[0]) + (dy - m2b[1]) * (m2[1] - m2b[1]) >= 0 + ): + alpha = wp.sqrt(((dx - m2[0]) * (dx - m2[0]) + (dy - m2[1]) * (dy - m2[1])) / dd) + minDist = dist + w = (m1 + (1.0 - alpha) * m2 + alpha * m2b) * 0.5 + var_rx = w[0] * dir + w[1] * dir2 + w[2] * normal + + for k in range(ncontact): + contact_points[k] = var_rx + + contact_count = 1 + + return contact_count, contact_points diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py new file mode 100644 index 00000000..8ccb259a --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py @@ -0,0 +1,320 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import warp as wp +from absl.testing import absltest + +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel + +from mujoco.mjx.third_party.mujoco_warp._src import test_util +from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import ccd +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import Geom +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model + +MAX_ITERATIONS = 10 + + +def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int): + @nested_kernel + def _gjk_kernel( + # Model: + geom_type: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + mesh_vertadr: wp.array(dtype=int), + mesh_vertnum: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + gid1: int, + gid2: int, + iterations: int, + vert: wp.array(dtype=wp.vec3), + vert1: wp.array(dtype=wp.vec3), + vert2: wp.array(dtype=wp.vec3), + vert_index1: wp.array(dtype=int), + vert_index2: wp.array(dtype=int), + face: wp.array(dtype=wp.vec3i), + face_pr: wp.array(dtype=wp.vec3), + face_norm2: wp.array(dtype=float), + face_index: wp.array(dtype=int), + face_map: wp.array(dtype=int), + horizon: wp.array(dtype=int), + # Out: + dist_out: wp.array(dtype=float), + pos_out: wp.array(dtype=wp.vec3), + ): + MESHGEOM = int(GeomType.MESH.value) + + geom1 = Geom() + geom1.index = -1 + geomtype1 = geom_type[gid1] + geom1.pos = geom_xpos_in[0, gid1] + geom1.rot = geom_xmat_in[0, gid1] + geom1.size = geom_size[0, gid1] + geom1.graphadr = -1 + + if geom_dataid[gid1] >= 0 and geom_type[gid1] == MESHGEOM: + dataid = geom_dataid[gid1] + geom1.vertadr = mesh_vertadr[dataid] + geom1.vertnum = mesh_vertnum[dataid] + geom1.vert = mesh_vert + + geom2 = Geom() + geom2.index = -1 + geomtype2 = geom_type[gid2] + geom2.pos = geom_xpos_in[0, gid2] + geom2.rot = geom_xmat_in[0, gid2] + geom2.size = geom_size[0, gid2] + geom2.graphadr = -1 + + if geom_dataid[gid2] >= 0 and geom_type[gid2] == MESHGEOM: + dataid = geom_dataid[gid2] + geom2.vertadr = mesh_vertadr[dataid] + geom2.vertnum = mesh_vertnum[dataid] + geom2.vert = mesh_vert + + x_1 = geom_xpos_in[0, gid1] + x_2 = geom_xpos_in[0, gid2] + + ( + dist, + x1, + x2, + ) = ccd( + 1e-6, + 1.0e30, + iterations, + iterations, + geom1, + geom2, + geomtype1, + geomtype2, + x_1, + x_2, + vert, + vert1, + vert2, + vert_index1, + vert_index2, + face, + face_pr, + face_norm2, + face_index, + face_map, + horizon, + ) + + dist_out[0] = dist + pos_out[0] = x1 + pos_out[1] = x2 + + vert = wp.array(shape=(iterations,), dtype=wp.vec3) + vert1 = wp.array(shape=(iterations,), dtype=wp.vec3) + vert2 = wp.array(shape=(iterations,), dtype=wp.vec3) + vert_index1 = wp.array(shape=(iterations,), dtype=int) + vert_index2 = wp.array(shape=(iterations,), dtype=int) + face = wp.array(shape=(2 * iterations,), dtype=wp.vec3i) + face_pr = wp.array(shape=(2 * iterations,), dtype=wp.vec3) + face_norm2 = wp.array(shape=(2 * iterations,), dtype=float) + face_index = wp.array(shape=(2 * iterations,), dtype=int) + face_map = wp.array(shape=(2 * iterations,), dtype=int) + horizon = wp.array(shape=(2 * iterations,), dtype=int) + dist_out = wp.array(shape=(1,), dtype=float) + pos_out = wp.array(shape=(2,), dtype=wp.vec3) + wp.launch( + _gjk_kernel, + dim=(1,), + inputs=[ + m.geom_type, + m.geom_dataid, + m.geom_size, + m.mesh_vertadr, + m.mesh_vertnum, + m.mesh_vert, + d.geom_xpos, + d.geom_xmat, + gid1, + gid2, + iterations, + vert, + vert1, + vert2, + vert_index1, + vert_index2, + face, + face_pr, + face_norm2, + face_index, + face_map, + horizon, + ], + outputs=[ + dist_out, + pos_out, + ], + ) + return dist_out.numpy()[0], pos_out.numpy()[0], pos_out.numpy()[1] + + +class GJKTest(absltest.TestCase): + """Tests for GJK/EPA.""" + + def test_spheres_distance(self): + """Test distance between two spheres.""" + + _, _, m, d = test_util.fixture( + xml=f""" + + + + + + + """ + ) + + dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS) + self.assertEqual(1.0, dist) + + def test_spheres_touching(self): + """Test two touching spheres have zero distance""" + + _, _, m, d = test_util.fixture( + xml=f""" + + + + + + + """ + ) + + dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS) + self.assertEqual(0.0, dist) + + def test_box_mesh_distance(self): + """Test distance between a mesh and box""" + + _, _, m, d = test_util.fixture( + xml=f""" + + + + + + + + + + """ + ) + + dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS) + self.assertAlmostEqual(0.1, dist) + + def test_sphere_sphere_contact(self): + """Test penetration depth between two spheres.""" + + _, _, m, d = test_util.fixture( + xml=f""" + + + + + + + """ + ) + + # TODO(kbayes): use margin trick instead of EPA for penetration recovery + dist, _, _ = _geom_dist(m, d, 0, 1, 500) + self.assertAlmostEqual(-2, dist) + + def test_box_box_contact(self): + """Test penetration between two boxes.""" + + _, _, m, d = test_util.fixture( + xml=f""" + + + + + + + """ + ) + dist, x1, x2 = _geom_dist(m, d, 0, 1, MAX_ITERATIONS) + self.assertAlmostEqual(-1, dist) + normal = wp.normalize(x1 - x2) + self.assertAlmostEqual(normal[0], 1) + self.assertAlmostEqual(normal[1], 0) + self.assertAlmostEqual(normal[2], 0) + + def test_mesh_mesh_contact(self): + """Test penetration between two meshes.""" + + _, _, m, d = test_util.fixture( + xml=f""" + + + + + + + + + + + + """ + ) + dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS) + self.assertAlmostEqual(-0.01, dist) + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py new file mode 100644 index 00000000..8397700c --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py @@ -0,0 +1,397 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +from typing import Tuple + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model + + +@wp.func +def _hfield_overlap_range( + # Model: + geom_dataid: wp.array(dtype=int), + geom_rbound: wp.array2d(dtype=float), + geom_margin: wp.array2d(dtype=float), + hfield_nrow: wp.array(dtype=int), + hfield_ncol: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + hfieldid: int, + geomid: int, + worldid: int, +) -> Tuple[int, int, int, int]: + """Returns min/max grid coordinates of height field cells overlapped by a geom's bounds. + + Args: + geom_dataid: Array of geom data IDs + geom_rbound: Array of geom bounding radii + geom_margin: Array of geom margins + hfield_nrow: Array of heightfield rows + hfield_ncol: Array of heightfield columns + hfield_size: Array of heightfield sizes + geom_xpos_in: Array of geom positions + geom_xmat_in: Array of geom orientation matrices + hfieldid: Index of the height field geom + geomid: Index of the other geom + worldid: Current world index + + Returns: + min_i, min_j, max_i, max_j: Grid coordinate bounds + """ + # get height field dimensions + dataid = geom_dataid[hfieldid] + nrow = hfield_nrow[dataid] + ncol = hfield_ncol[dataid] + size = hfield_size[dataid] # (x, y, z_top, z_bottom) + + # get positions and transforms + hf_pos = geom_xpos_in[worldid, hfieldid] + hf_mat = geom_xmat_in[worldid, hfieldid] + geom_pos = geom_xpos_in[worldid, geomid] + + # transform geom_pos to height field local space + local_pos = wp.transpose(hf_mat) @ (geom_pos - hf_pos) + + # get bounding radius of other geometry (including margin) + bound_radius = geom_rbound[worldid, geomid] + geom_margin[worldid, geomid] + + # calculate grid resolution + x_scale = 2.0 * size[0] / float(ncol - 1) + y_scale = 2.0 * size[1] / float(nrow - 1) + + # calculate min/max grid coordinates that could contain the object + min_i = wp.max(0, int((local_pos[0] - bound_radius + size[0]) / x_scale)) + max_i = wp.min(ncol - 2, int((local_pos[0] + bound_radius + size[0]) / x_scale) + 1) + min_j = wp.max(0, int((local_pos[1] - bound_radius + size[1]) / y_scale)) + max_j = wp.min(nrow - 2, int((local_pos[1] + bound_radius + size[1]) / y_scale) + 1) + + return min_i, min_j, max_i, max_j + + +@wp.func +def hfield_triangle_prism( + # Model: + geom_dataid: wp.array(dtype=int), + hfield_adr: wp.array(dtype=int), + hfield_nrow: wp.array(dtype=int), + hfield_ncol: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + hfield_data: wp.array(dtype=float), + # In: + hfieldid: int, + hftri_index: int, +) -> wp.mat33: + """Returns the vertices of a triangular prism for a heightfield triangle. + + Args: + geom_dataid: Array of geometry data IDs + hfield_adr: Array of heightfield addresses + hfield_nrow: Array of heightfield rows + hfield_ncol: Array of heightfield columns + hfield_size: Array of heightfield sizes + hfield_data: Array of heightfield data + hfieldid: Index of the height field geometry + hftri_index: Index of the triangle in the heightfield + + Returns: + 3x3 matrix containing the vertices of the triangular prism + """ + # https://mujoco.readthedocs.io/en/stable/XMLreference.html#asset-hfield + + # get heightfield dimensions + dataid = geom_dataid[hfieldid] + if dataid < 0 or hftri_index < 0: + return wp.mat33(0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0) + + nrow = hfield_nrow[dataid] + ncol = hfield_ncol[dataid] + size = hfield_size[dataid] # (x, y, z_top, z_bottom) + + # calculate which triangle in the grid + row = (hftri_index // 2) // (ncol - 1) + col = (hftri_index // 2) % (ncol - 1) + + # calculate vertices in 2D grid + x_scale = 2.0 * size[0] / float(ncol - 1) + y_scale = 2.0 * size[1] / float(nrow - 1) + + # grid coordinates (i, j) for triangle corners + i0 = col + j0 = row + i1 = i0 + 1 + j1 = j0 + 1 + + # convert grid coordinates to local space x, y coordinates + x0 = float(i0) * x_scale - size[0] + y0 = float(j0) * y_scale - size[1] + x1 = float(i1) * x_scale - size[0] + y1 = float(j1) * y_scale - size[1] + + # get height values at corners from hfield_data + base_addr = hfield_adr[dataid] + z00 = hfield_data[base_addr + j0 * ncol + i0] + z01 = hfield_data[base_addr + j1 * ncol + i0] + z10 = hfield_data[base_addr + j0 * ncol + i1] + z11 = hfield_data[base_addr + j1 * ncol + i1] + + # scale heights from range [0, 1] to [0, z_top] + z_top = size[2] + z00 = z00 * z_top + z01 = z01 * z_top + z10 = z10 * z_top + z11 = z11 * z_top + + # set bottom z-value + z_bottom = -size[3] + + # compress 6 prism vertices into 3x3 matrix, see hfield_prism_vertex for details + return wp.mat33( + x0, + y0, + z00, + x1, + y1, + z11, + wp.where(hftri_index % 2, 1.0, 0.0), + wp.where(hftri_index % 2, z10, z01), + z_bottom, + ) + + +@wp.func +def hfield_prism_vertex(prism: wp.mat33, vert_index: int) -> wp.vec3: + """Extracts vertices from a compressed triangular prism representation. + + The compression scheme stores a 6-vertex triangular prism using a 3x3 matrix: + - prism[0] = First vertex (x,y,z) - corner (i,j) + - prism[1] = Second vertex (x,y,z) - corner (i+1,j+1) + - prism[2,0] = Triangle type flag: 0 for even triangle (using corner (i,j+1)), + non-zero for odd triangle (using corner (i+1,j)) + - prism[2,1] = Z-coordinate of the third vertex + - prism[2,2] = Z-coordinate used for all bottom vertices (common z) + + In this way, we can reconstruct all 6 vertices of the prism by reusing + coordinates from the stored vertices. + + Args: + prism: 3x3 compressed representation of a triangular prism + vert_index: Index of vertex to extract (0-5) + + Returns: + The 3D coordinates of the requested vertex + """ + if vert_index == 0 or vert_index == 1: + return prism[vert_index] # first two vertices stored directly + + if vert_index == 2: # third vertex + if prism[2][0] == 0: # even triangle (i, j+1) + return wp.vec3(prism[0][0], prism[1][1], prism[2][1]) + else: # odd triangle (i+1, j) + return wp.vec3(prism[1][0], prism[0][1], prism[2][1]) + + if vert_index == 3 or vert_index == 4: # bottom vertices below 0 and 1 + return wp.vec3(prism[vert_index - 3][0], prism[vert_index - 3][1], prism[2][2]) + + if vert_index == 5: # bottom vertex below 2 + if prism[2][0] == 0: # even triangle + return wp.vec3(prism[0][0], prism[1][1], prism[2][2]) + else: # odd triangle + return wp.vec3(prism[1][0], prism[0][1], prism[2][2]) + + +@wp.kernel +def _hfield_midphase( + # Model: + geom_type: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_rbound: wp.array2d(dtype=float), + geom_margin: wp.array2d(dtype=float), + hfield_nrow: wp.array(dtype=int), + hfield_ncol: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + # Data in: + nconmax_in: int, + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + collision_pair_in: wp.array(dtype=wp.vec2i), + collision_hftri_index_in: wp.array(dtype=int), + collision_pairid_in: wp.array(dtype=int), + collision_worldid_in: wp.array(dtype=int), + # Data out: + collision_pair_out: wp.array(dtype=wp.vec2i), + collision_hftri_index_out: wp.array(dtype=int), + collision_pairid_out: wp.array(dtype=int), + collision_worldid_out: wp.array(dtype=int), + ncollision_out: wp.array(dtype=int), +): + """Midphase collision detection for heightfield triangles with other geoms. + + This kernel processes collision pairs where one geom is a heightfield (identified by + collision_hftri_index_in[pairid] == -1) and expands them into multiple collision pairs, + one for each potentially colliding triangle. + + Args: + geom_type: Array of geometry types + geom_dataid: Array of geometry data IDs + geom_rbound: Array of geometry bounding radii + geom_margin: Array of geometry margins + hfield_nrow: Array of heightfield rows + hfield_ncol: Array of heightfield columns + hfield_size: Array of heightfield sizes + nconmax_in: Max number of collisions + geom_xpos_in: Array of geometry positions + geom_xmat_in: Array of geometry orientation matrices + collision_pair_in: Array of collision pairs + collision_hftri_index_in: Array of heightfield triangle indices, -1 for heightfield + pairs + collision_pairid_in: Array of collision pair IDs + collision_worldid_in: Array of collision world IDs + + collision_pair_out: Output array of collision pairs + collision_hftri_index_out: Output array of heightfield triangle indices + collision_pairid_out: Output array of collision pair IDs + collision_worldid_out: Output array of collision world IDs + ncollision_out: Output counter for number of collisions + """ + pairid = wp.tid() + + # only process pairs that are marked for heightfield collision (-1) + # the buffer is cleared at the start of each frame in collision_driver.py + if collision_hftri_index_in[pairid] != -1: + return + + # get the collision pair info + pair = collision_pair_in[pairid] + worldid = collision_worldid_in[pairid] + pair_id = collision_pairid_in[pairid] + + # identify which geom is the heightfield + g1 = pair[0] + g2 = pair[1] + + hfieldid = g1 + geomid = g2 + + # if the first geom is not a heightfield, swap them + # in theory, shouldn't happen as _add_geom_pair already sorted the pair + if geom_type[g1] != int(GeomType.HFIELD.value): + hfieldid = g2 + geomid = g1 + + # get min/max grid coordinates for overlap region + min_i, min_j, max_i, max_j = _hfield_overlap_range( + geom_dataid, + geom_rbound, + geom_margin, + hfield_nrow, + hfield_ncol, + hfield_size, + geom_xpos_in, + geom_xmat_in, + hfieldid, + geomid, + worldid, + ) + + # get hfield dimensions for triangle index calculation + dataid = geom_dataid[hfieldid] + ncol = hfield_ncol[dataid] + + # loop through grid cells and add pairs for all triangles + for j in range(min_j, max_j + 1): + for i in range(min_i, max_i + 1): + # each grid cell contains two triangles + base_idx = ((j * (ncol - 1)) + i) * 2 + + # add both triangles from this cell + for t in range(2): + if i == 0 and j == 0 and t == 0: + # reuse the initial pair for the 1st triangle + new_pairid = pairid + else: + # for the rest create a new pair + new_pairid = wp.atomic_add(ncollision_out, 0, 1) + + if new_pairid >= nconmax_in: + return + + collision_pair_out[new_pairid] = pair + collision_hftri_index_out[new_pairid] = base_idx + t + collision_pairid_out[new_pairid] = pair_id + collision_worldid_out[new_pairid] = worldid + + +def hfield_midphase(m: Model, d: Data): + """Midphase collision detection for heightfield triangles with other geoms. + + Processes collision pairs from the broadphase where one geom is a heightfield and expands + them into multiple collision pairs, one for each potentially colliding triangle. The + function directly writes to the same collision buffers used by _add_geom_pair. + + Args: + m: Model containing geometry and heightfield data + - geom_type: Array of geometry types + - geom_dataid: Array of geometry data IDs + - hfield_nrow: Array of heightfield rows + - hfield_ncol: Array of heightfield columns + - hfield_size: Array of heightfield sizes + - geom_rbound: Array of geometry bounding radii + - geom_margin: Array of geometry margins + d: Data containing current state and collision information + - nconmax: Maximum number of contacts + - geom_xpos: Array of geometry positions + - geom_xmat: Array of geometry orientation matrices + - collision_pair: Array of collision pairs + - collision_hftri_index: Array of heightfield triangle indices + - collision_pairid: Array of collision pair IDs + - collision_worldid: Array of collision world IDs + - ncollision: Number of collisions + """ + # launch the midphase kernel to expand height field collision pairs + # write directly to the same buffers that _add_geom_pair writes to + wp.launch( + kernel=_hfield_midphase, + dim=d.nconmax, # launch threads to process all potential pairs + inputs=[ + m.geom_type, + m.geom_dataid, + m.geom_rbound, + m.geom_margin, + m.hfield_nrow, + m.hfield_ncol, + m.hfield_size, + d.nconmax, + d.geom_xpos, + d.geom_xmat, + d.collision_pair, + d.collision_hftri_index, + d.collision_pairid, + d.collision_worldid, + ], + outputs=[ + d.collision_pair, + d.collision_hftri_index, + d.collision_pairid, + d.collision_worldid, + d.ncollision, + ], + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py new file mode 100644 index 00000000..b0e6de66 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py @@ -0,0 +1,2949 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_triangle_prism +from mujoco.mjx.third_party.mujoco_warp._src.math import closest_segment_point +from mujoco.mjx.third_party.mujoco_warp._src.math import closest_segment_to_segment_points +from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame +from mujoco.mjx.third_party.mujoco_warp._src.math import normalize_with_norm +from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINMU +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope + +wp.set_module_options({"enable_backward": False}) + + +class vec8f(wp.types.vector(length=8, dtype=wp.float32)): + pass + + +class mat43f(wp.types.matrix(shape=(4, 3), dtype=wp.float32)): + pass + + +class mat83f(wp.types.matrix(shape=(8, 3), dtype=wp.float32)): + pass + + +@wp.struct +class Geom: + pos: wp.vec3 + rot: wp.mat33 + normal: wp.vec3 + size: wp.vec3 + hfprism: wp.mat33 + vertadr: int + vertnum: int + vert: wp.array(dtype=wp.vec3) + graphadr: int + graph: wp.array(dtype=int) + mesh_polynum: int + mesh_polyadr: int + mesh_polynormal: wp.array(dtype=wp.vec3) + mesh_polyvertadr: wp.array(dtype=int) + mesh_polyvertnum: wp.array(dtype=int) + mesh_polyvert: wp.array(dtype=int) + mesh_polymapadr: wp.array(dtype=int) + mesh_polymapnum: wp.array(dtype=int) + mesh_polymap: wp.array(dtype=int) + index: int + + +@wp.func +def _geom( + # Model: + geom_type: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + hfield_adr: wp.array(dtype=int), + hfield_nrow: wp.array(dtype=int), + hfield_ncol: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + hfield_data: wp.array(dtype=float), + mesh_vertadr: wp.array(dtype=int), + mesh_vertnum: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_graphadr: wp.array(dtype=int), + mesh_graph: wp.array(dtype=int), + mesh_polynum: wp.array(dtype=int), + mesh_polyadr: wp.array(dtype=int), + mesh_polynormal: wp.array(dtype=wp.vec3), + mesh_polyvertadr: wp.array(dtype=int), + mesh_polyvertnum: wp.array(dtype=int), + mesh_polyvert: wp.array(dtype=int), + mesh_polymapadr: wp.array(dtype=int), + mesh_polymapnum: wp.array(dtype=int), + mesh_polymap: wp.array(dtype=int), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + worldid: int, + gid: int, + hftri_index: int, +) -> Geom: + geom = Geom() + geom.pos = geom_xpos_in[worldid, gid] + rot = geom_xmat_in[worldid, gid] + geom.rot = rot + geom.size = geom_size[worldid, gid] + geom.normal = wp.vec3(rot[0, 2], rot[1, 2], rot[2, 2]) # plane + dataid = geom_dataid[gid] + + # If geom is MESH, get mesh verts + if dataid >= 0 and geom_type[gid] == int(GeomType.MESH.value): + geom.vertadr = mesh_vertadr[dataid] + geom.vertnum = mesh_vertnum[dataid] + geom.graphadr = mesh_graphadr[dataid] + geom.mesh_polynum = mesh_polynum[dataid] + geom.mesh_polyadr = mesh_polyadr[dataid] + else: + geom.vertadr = -1 + geom.vertnum = -1 + geom.graphadr = -1 + geom.mesh_polynum = -1 + geom.mesh_polyadr = -1 + + if geom_type[gid] == int(GeomType.MESH.value): + geom.vert = mesh_vert + geom.graph = mesh_graph + geom.mesh_polynormal = mesh_polynormal + geom.mesh_polyvertadr = mesh_polyvertadr + geom.mesh_polyvertnum = mesh_polyvertnum + geom.mesh_polyvert = mesh_polyvert + geom.mesh_polymapadr = mesh_polymapadr + geom.mesh_polymapnum = mesh_polymapnum + geom.mesh_polymap = mesh_polymap + + # If geom is HFIELD triangle, compute triangle prism verts + if geom_type[gid] == int(GeomType.HFIELD.value): + geom.hfprism = hfield_triangle_prism( + geom_dataid, hfield_adr, hfield_nrow, hfield_ncol, hfield_size, hfield_data, gid, hftri_index + ) + + geom.index = -1 + return geom + + +@wp.func +def write_contact( + # Data in: + nconmax_in: int, + # In: + dist_in: float, + pos_in: wp.vec3, + frame_in: wp.mat33, + margin_in: float, + gap_in: float, + condim_in: int, + friction_in: vec5, + solref_in: wp.vec2f, + solreffriction_in: wp.vec2f, + solimp_in: vec5, + geoms_in: wp.vec2i, + worldid_in: int, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + active = (dist_in - margin_in) < 0 + if active: + cid = wp.atomic_add(ncon_out, 0, 1) + if cid < nconmax_in: + contact_dist_out[cid] = dist_in + contact_pos_out[cid] = pos_in + contact_frame_out[cid] = frame_in + contact_geom_out[cid] = geoms_in + contact_worldid_out[cid] = worldid_in + includemargin = margin_in - gap_in + contact_includemargin_out[cid] = includemargin + contact_dim_out[cid] = condim_in + contact_friction_out[cid] = friction_in + contact_solref_out[cid] = solref_in + contact_solreffriction_out[cid] = solreffriction_in + contact_solimp_out[cid] = solimp_in + + +@wp.func +def _plane_sphere(plane_normal: wp.vec3, plane_pos: wp.vec3, sphere_pos: wp.vec3, sphere_radius: float): + dist = wp.dot(sphere_pos - plane_pos, plane_normal) - sphere_radius + pos = sphere_pos - plane_normal * (sphere_radius + 0.5 * dist) + return dist, pos + + +@wp.func +def plane_sphere( + # Data in: + nconmax_in: int, + # In: + plane: Geom, + sphere: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + dist, pos = _plane_sphere(plane.normal, plane.pos, sphere.pos, sphere.size[0]) + + write_contact( + nconmax_in, + dist, + pos, + make_frame(plane.normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def _sphere_sphere( + # Data in: + nconmax_in: int, + # In: + pos1: wp.vec3, + radius1: float, + pos2: wp.vec3, + radius2: float, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + dir = pos2 - pos1 + dist = wp.length(dir) + if dist == 0.0: + n = wp.vec3(1.0, 0.0, 0.0) + else: + n = dir / dist + dist = dist - (radius1 + radius2) + pos = pos1 + n * (radius1 + 0.5 * dist) + + write_contact( + nconmax_in, + dist, + pos, + make_frame(n), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def _sphere_sphere_ext( + # Data in: + nconmax_in: int, + # In: + pos1: wp.vec3, + radius1: float, + pos2: wp.vec3, + radius2: float, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + mat1: wp.mat33, + mat2: wp.mat33, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + dir = pos2 - pos1 + dist = wp.length(dir) + if dist == 0.0: + # Use cross product of z axes like MuJoCo + axis1 = wp.vec3(mat1[0, 2], mat1[1, 2], mat1[2, 2]) + axis2 = wp.vec3(mat2[0, 2], mat2[1, 2], mat2[2, 2]) + n = wp.cross(axis1, axis2) + n = wp.normalize(n) + else: + n = dir / dist + dist = dist - (radius1 + radius2) + pos = pos1 + n * (radius1 + 0.5 * dist) + + write_contact( + nconmax_in, + dist, + pos, + make_frame(n), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def sphere_sphere( + # Data in: + nconmax_in: int, + # In: + sphere1: Geom, + sphere2: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + _sphere_sphere( + nconmax_in, + sphere1.pos, + sphere1.size[0], + sphere2.pos, + sphere2.size[0], + worldid, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def sphere_capsule( + # Data in: + nconmax_in: int, + # In: + sphere: Geom, + cap: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + """Calculates one contact between a sphere and a capsule.""" + axis = wp.vec3(cap.rot[0, 2], cap.rot[1, 2], cap.rot[2, 2]) + length = cap.size[1] + segment = axis * length + + # Find closest point on capsule centerline to sphere center + pt = closest_segment_point(cap.pos - segment, cap.pos + segment, sphere.pos) + + # Treat as sphere-sphere collision between sphere and closest point + _sphere_sphere( + nconmax_in, + sphere.pos, + sphere.size[0], + pt, + cap.size[0], + worldid, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def capsule_capsule( + # Data in: + nconmax_in: int, + # In: + cap1: Geom, + cap2: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + axis1 = wp.vec3(cap1.rot[0, 2], cap1.rot[1, 2], cap1.rot[2, 2]) + axis2 = wp.vec3(cap2.rot[0, 2], cap2.rot[1, 2], cap2.rot[2, 2]) + length1 = cap1.size[1] + length2 = cap2.size[1] + seg1 = axis1 * length1 + seg2 = axis2 * length2 + + pt1, pt2 = closest_segment_to_segment_points( + cap1.pos - seg1, + cap1.pos + seg1, + cap2.pos - seg2, + cap2.pos + seg2, + ) + + _sphere_sphere( + nconmax_in, + pt1, + cap1.size[0], + pt2, + cap2.size[0], + worldid, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def plane_capsule( + # Data in: + nconmax_in: int, + # In: + plane: Geom, + cap: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + """Calculates two contacts between a capsule and a plane.""" + n = plane.normal + axis = wp.vec3(cap.rot[0, 2], cap.rot[1, 2], cap.rot[2, 2]) + # align contact frames with capsule axis + b, b_norm = normalize_with_norm(axis - n * wp.dot(n, axis)) + + if b_norm < 0.5: + if -0.5 < n[1] and n[1] < 0.5: + b = wp.vec3(0.0, 1.0, 0.0) + else: + b = wp.vec3(0.0, 0.0, 1.0) + + c = wp.cross(n, b) + frame = wp.mat33(n[0], n[1], n[2], b[0], b[1], b[2], c[0], c[1], c[2]) + segment = axis * cap.size[1] + + dist1, pos1 = _plane_sphere(n, plane.pos, cap.pos + segment, cap.size[0]) + write_contact( + nconmax_in, + dist1, + pos1, + frame, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + dist2, pos2 = _plane_sphere(n, plane.pos, cap.pos - segment, cap.size[0]) + write_contact( + nconmax_in, + dist2, + pos2, + frame, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def plane_ellipsoid( + # Data in: + nconmax_in: int, + # In: + plane: Geom, + ellipsoid: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + sphere_support = -wp.normalize(wp.cw_mul(wp.transpose(ellipsoid.rot) @ plane.normal, ellipsoid.size)) + pos = ellipsoid.pos + ellipsoid.rot @ wp.cw_mul(sphere_support, ellipsoid.size) + dist = wp.dot(plane.normal, pos - plane.pos) + pos = pos - plane.normal * dist * 0.5 + + write_contact( + nconmax_in, + dist, + pos, + make_frame(plane.normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def plane_box( + # Data in: + nconmax_in: int, + # In: + plane: Geom, + box: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + count = int(0) + corner = wp.vec3() + dist = wp.dot(box.pos - plane.pos, plane.normal) + + # test all corners, pick bottom 4 + for i in range(8): + # get corner in local coordinates + corner.x = wp.where(i & 1, box.size.x, -box.size.x) + corner.y = wp.where(i & 2, box.size.y, -box.size.y) + corner.z = wp.where(i & 4, box.size.z, -box.size.z) + + # get corner in global coordinates relative to box center + corner = box.rot * corner + + # compute distance to plane, skip if too far or pointing up + ldist = wp.dot(plane.normal, corner) + if dist + ldist > margin or ldist > 0: + continue + + cdist = dist + ldist + frame = make_frame(plane.normal) + pos = corner + box.pos + (plane.normal * cdist / -2.0) + write_contact( + nconmax_in, + cdist, + pos, + frame, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + count += 1 + if count >= 4: + break + + +_HUGE_VAL = 1e6 + + +@wp.func +def plane_convex( + # Data in: + nconmax_in: int, + # In: + plane: Geom, + convex: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + """Calculates contacts between a plane and a convex object.""" + + # get points in the convex frame + plane_pos = wp.transpose(convex.rot) @ (plane.pos - convex.pos) + n = wp.transpose(convex.rot) @ plane.normal + + # Store indices in vec4 + indices = wp.vec4i(-1, -1, -1, -1) + + # exhaustive search over all vertices + if convex.graphadr == -1 or convex.vertnum < 10: + # Find support points + max_support = wp.float32(-_HUGE_VAL) + for i in range(convex.vertnum): + support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n) + max_support = wp.max(support, max_support) + + threshold = wp.max(0.0, max_support - 1e-3) + # Find point a (first support point) + a_dist = wp.float32(-_HUGE_VAL) + for i in range(convex.vertnum): + support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n) + dist = wp.where(support > threshold, 0.0, -_HUGE_VAL) + if dist > a_dist: + indices[0] = i + a_dist = dist + a = convex.vert[convex.vertadr + indices[0]] + + # Find point b (furthest from a) + b_dist = wp.float32(-_HUGE_VAL) + for i in range(convex.vertnum): + support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n) + dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL) + dist = wp.length_sq(a - convex.vert[convex.vertadr + i]) + dist_mask + if dist > b_dist: + indices[1] = i + b_dist = dist + b = convex.vert[convex.vertadr + indices[1]] + + # Find point c (furthest along axis orthogonal to a-b) + ab = wp.cross(n, a - b) + c_dist = wp.float32(-_HUGE_VAL) + for i in range(convex.vertnum): + support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n) + dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL) + dist = wp.length_sq(ab - convex.vert[convex.vertadr + i]) + dist_mask + if dist > c_dist: + indices[2] = i + c_dist = dist + c = convex.vert[convex.vertadr + indices[2]] + + # Find point d (furthest from other triangle edges) + ac = wp.cross(n, a - c) + bc = wp.cross(n, b - c) + d_dist = wp.float32(-_HUGE_VAL) + for i in range(convex.vertnum): + support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n) + dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL) + ap = ac - convex.vert[convex.vertadr + i] + bp = bc - convex.vert[convex.vertadr + i] + dist_ap = wp.abs(wp.dot(ap, ac)) + dist_mask + dist_bp = wp.abs(wp.dot(bp, bc)) + dist_mask + if dist_ap + dist_bp > d_dist: + indices[3] = i + d_dist = dist_ap + dist_bp + + else: + numvert = convex.graph[convex.graphadr] + vert_edgeadr = convex.graphadr + 2 + vert_globalid = convex.graphadr + 2 + numvert + edge_localid = convex.graphadr + 2 + 2 * numvert + + # Find support points + max_support = wp.float32(-_HUGE_VAL) + + # hillclimb until no change + prev = int(-1) + imax = int(0) + + while True: + prev = int(imax) + i = int(convex.graph[vert_edgeadr + imax]) + while convex.graph[edge_localid + i] >= 0: + subidx = convex.graph[edge_localid + i] + idx = convex.graph[vert_globalid + subidx] + support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n) + if support > max_support: + max_support = support + imax = int(subidx) + i += int(1) + if imax == prev: + break + + threshold = wp.max(0.0, max_support - 1e-3) + + a_dist = wp.float32(-_HUGE_VAL) + # hillclimb until no change + prev = int(-1) + imax = int(0) + + while True: + prev = int(imax) + i = int(convex.graph[vert_edgeadr + imax]) + while convex.graph[edge_localid + i] >= 0: + subidx = convex.graph[edge_localid + i] + idx = convex.graph[vert_globalid + subidx] + support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n) + dist = wp.where(support > threshold, 0.0, -_HUGE_VAL) + if dist > a_dist: + a_dist = dist + imax = int(subidx) + i += int(1) + if imax == prev: + break + imax = convex.graph[vert_globalid + imax] + a = convex.vert[convex.vertadr + imax] + indices[0] = imax + + # Find point b (furthest from a) + b_dist = wp.float32(-_HUGE_VAL) + # hillclimb until no change + prev = int(-1) + imax = int(0) + + while True: + prev = int(imax) + i = int(convex.graph[vert_edgeadr + imax]) + while convex.graph[edge_localid + i] >= 0: + subidx = convex.graph[edge_localid + i] + idx = convex.graph[vert_globalid + subidx] + support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n) + dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL) + dist = wp.length_sq(a - convex.vert[convex.vertadr + idx]) + dist_mask + if dist > b_dist: + b_dist = dist + imax = int(subidx) + i += int(1) + if imax == prev: + break + imax = convex.graph[vert_globalid + imax] + b = convex.vert[convex.vertadr + imax] + indices[1] = imax + + # Find point c (furthest along axis orthogonal to a-b) + ab = wp.cross(n, a - b) + c_dist = wp.float32(-_HUGE_VAL) + # hillclimb until no change + prev = int(-1) + imax = int(0) + + while True: + prev = int(imax) + i = int(convex.graph[vert_edgeadr + imax]) + while convex.graph[edge_localid + i] >= 0: + subidx = convex.graph[edge_localid + i] + idx = convex.graph[vert_globalid + subidx] + support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n) + dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL) + dist = wp.length_sq(ab - convex.vert[convex.vertadr + idx]) + dist_mask + if dist > c_dist: + c_dist = dist + imax = int(subidx) + i += int(1) + if imax == prev: + break + imax = convex.graph[vert_globalid + imax] + c = convex.vert[convex.vertadr + imax] + indices[2] = imax + + # Find point d (furthest from other triangle edges) + ac = wp.cross(n, a - c) + bc = wp.cross(n, b - c) + d_dist = wp.float32(-_HUGE_VAL) + # hillclimb until no change + prev = int(-1) + imax = int(0) + + while True: + prev = int(imax) + i = int(convex.graph[vert_edgeadr + imax]) + while convex.graph[edge_localid + i] >= 0: + subidx = convex.graph[edge_localid + i] + idx = convex.graph[vert_globalid + subidx] + support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n) + dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL) + ap = ac - convex.vert[convex.vertadr + idx] + bp = bc - convex.vert[convex.vertadr + idx] + dist_ap = wp.abs(wp.dot(ap, ac)) + dist_mask + dist_bp = wp.abs(wp.dot(bp, bc)) + dist_mask + if dist_ap + dist_bp > d_dist: + d_dist = dist_ap + dist_bp + imax = int(subidx) + i += int(1) + if imax == prev: + break + imax = convex.graph[vert_globalid + imax] + indices[3] = imax + + # Write contacts + frame = make_frame(plane.normal) + for i in range(3, -1, -1): + idx = indices[i] + count = int(0) + for j in range(i + 1): + if indices[j] == idx: + count = count + 1 + + # Check if the index is unique (appears exactly once) + if count == 1: + pos = convex.vert[convex.vertadr + idx] + pos = convex.pos + convex.rot @ pos + support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n) + dist = -support + pos = pos - 0.5 * dist * plane.normal + write_contact( + nconmax_in, + dist, + pos, + frame, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def sphere_cylinder( + # Data in: + nconmax_in: int, + # In: + sphere: Geom, + cylinder: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + axis = wp.vec3( + cylinder.rot[0, 2], + cylinder.rot[1, 2], + cylinder.rot[2, 2], + ) + + vec = sphere.pos - cylinder.pos + x = wp.dot(vec, axis) + + a_proj = axis * x + p_proj = vec - a_proj + p_proj_sqr = wp.dot(p_proj, p_proj) + + collide_side = wp.abs(x) < cylinder.size[1] + collide_cap = p_proj_sqr < (cylinder.size[0] * cylinder.size[0]) + + if collide_side and collide_cap: + dist_cap = cylinder.size[1] - wp.abs(x) + dist_radius = cylinder.size[0] - wp.sqrt(p_proj_sqr) + + if dist_cap < dist_radius: + collide_side = False + else: + collide_cap = False + + # Side collision + if collide_side: + pos_target = cylinder.pos + a_proj + + _sphere_sphere_ext( + nconmax_in, + sphere.pos, + sphere.size[0], + pos_target, + cylinder.size[0], + worldid, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + sphere.rot, + cylinder.rot, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + return + + # Cap collision + if collide_cap: + if x > 0.0: + # top cap + pos_cap = cylinder.pos + axis * cylinder.size[1] + plane_normal = axis + else: + # bottom cap + pos_cap = cylinder.pos - axis * cylinder.size[1] + plane_normal = -axis + + dist, pos_contact = _plane_sphere(plane_normal, pos_cap, sphere.pos, sphere.size[0]) + plane_normal = -plane_normal # Flip normal after position calculation + + write_contact( + nconmax_in, + dist, + pos_contact, + make_frame(plane_normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + return + + # Corner collision + inv_len = 1.0 / wp.sqrt(p_proj_sqr) + p_proj = p_proj * (cylinder.size[0] * inv_len) + + cap_offset = axis * (wp.sign(x) * cylinder.size[1]) + pos_corner = cylinder.pos + cap_offset + p_proj + + _sphere_sphere_ext( + nconmax_in, + sphere.pos, + sphere.size[0], + pos_corner, + 0.0, + worldid, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + sphere.rot, + cylinder.rot, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def plane_cylinder( + # Data in: + nconmax_in: int, + # In: + plane: Geom, + cylinder: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + """Calculates contacts between a cylinder and a plane.""" + # Extract plane normal and cylinder axis + n = plane.normal + axis = wp.vec3(cylinder.rot[0, 2], cylinder.rot[1, 2], cylinder.rot[2, 2]) + + # Project, make sure axis points toward plane + prjaxis = wp.dot(n, axis) + if prjaxis > 0: + axis = -axis + prjaxis = -prjaxis + + # Compute normal distance from plane to cylinder center + dist0 = wp.dot(cylinder.pos - plane.pos, n) + + # Remove component of -normal along cylinder axis + vec = axis * prjaxis - n + len_sqr = wp.dot(vec, vec) + + # If vector is nondegenerate, normalize and scale by radius + # Otherwise use cylinder's x-axis scaled by radius + vec = wp.where( + len_sqr >= 1e-12, + vec * (cylinder.size[0] / wp.sqrt(len_sqr)), + wp.vec3(cylinder.rot[0, 0], cylinder.rot[1, 0], cylinder.rot[2, 0]) * cylinder.size[0], + ) + + # Project scaled vector on normal + prjvec = wp.dot(vec, n) + + # Scale cylinder axis by half-length + axis = axis * cylinder.size[1] + prjaxis = prjaxis * cylinder.size[1] + + frame = make_frame(n) + + # First contact point (end cap closer to plane) + dist1 = dist0 + prjaxis + prjvec + if dist1 <= margin: + pos1 = cylinder.pos + vec + axis - n * (dist1 * 0.5) + write_contact( + nconmax_in, + dist1, + pos1, + frame, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + else: + # If nearest point is above margin, no contacts + return + + # Second contact point (end cap farther from plane) + dist2 = dist0 - prjaxis + prjvec + if dist2 <= margin: + pos2 = cylinder.pos + vec - axis - n * (dist2 * 0.5) + write_contact( + nconmax_in, + dist2, + pos2, + make_frame(plane.normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + # Try triangle contact points on side closer to plane + prjvec1 = -prjvec * 0.5 + dist3 = dist0 + prjaxis + prjvec1 + if dist3 <= margin: + # Compute sideways vector scaled by radius*sqrt(3)/2 + vec1 = wp.cross(vec, axis) + vec1 = wp.normalize(vec1) * (cylinder.size[0] * wp.sqrt(3.0) * 0.5) + + # Add contact point A - adjust to closest side + pos3 = cylinder.pos + vec1 + axis - vec * 0.5 - n * (dist3 * 0.5) + write_contact( + nconmax_in, + dist3, + pos3, + frame, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + # Add contact point B - adjust to closest side + pos4 = cylinder.pos - vec1 + axis - vec * 0.5 - n * (dist3 * 0.5) + write_contact( + nconmax_in, + dist3, + pos4, + frame, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def contact_params( + # Model: + geom_condim: wp.array(dtype=int), + geom_priority: wp.array(dtype=int), + geom_solmix: wp.array2d(dtype=float), + geom_solref: wp.array2d(dtype=wp.vec2), + geom_solimp: wp.array2d(dtype=vec5), + geom_friction: wp.array2d(dtype=wp.vec3), + geom_margin: wp.array2d(dtype=float), + geom_gap: wp.array2d(dtype=float), + pair_dim: wp.array(dtype=int), + pair_solref: wp.array2d(dtype=wp.vec2), + pair_solreffriction: wp.array2d(dtype=wp.vec2), + pair_solimp: wp.array2d(dtype=vec5), + pair_margin: wp.array2d(dtype=float), + pair_gap: wp.array2d(dtype=float), + pair_friction: wp.array2d(dtype=vec5), + # Data in: + collision_pair_in: wp.array(dtype=wp.vec2i), + collision_pairid_in: wp.array(dtype=int), + # In: + cid: int, + worldid: int, +): + geoms = collision_pair_in[cid] + pairid = collision_pairid_in[cid] + + if pairid > -1: + margin = pair_margin[worldid, pairid] + gap = pair_gap[worldid, pairid] + condim = pair_dim[pairid] + friction = pair_friction[worldid, pairid] + solref = pair_solref[worldid, pairid] + solreffriction = pair_solreffriction[worldid, pairid] + solimp = pair_solimp[worldid, pairid] + else: + g1 = geoms[0] + g2 = geoms[1] + + p1 = geom_priority[g1] + p2 = geom_priority[g2] + + solmix1 = geom_solmix[worldid, g1] + solmix2 = geom_solmix[worldid, g2] + + mix = solmix1 / (solmix1 + solmix2) + mix = wp.where((solmix1 < MJ_MINVAL) and (solmix2 < MJ_MINVAL), 0.5, mix) + mix = wp.where((solmix1 < MJ_MINVAL) and (solmix2 >= MJ_MINVAL), 0.0, mix) + mix = wp.where((solmix1 >= MJ_MINVAL) and (solmix2 < MJ_MINVAL), 1.0, mix) + mix = wp.where(p1 == p2, mix, wp.where(p1 > p2, 1.0, 0.0)) + + margin = wp.max(geom_margin[worldid, g1], geom_margin[worldid, g2]) + gap = wp.max(geom_gap[worldid, g1], geom_gap[worldid, g2]) + + condim1 = geom_condim[g1] + condim2 = geom_condim[g2] + condim = wp.where(p1 == p2, wp.max(condim1, condim2), wp.where(p1 > p2, condim1, condim2)) + + max_geom_friction = wp.max(geom_friction[worldid, g1], geom_friction[worldid, g2]) + friction = vec5( + wp.max(MJ_MINMU, max_geom_friction[0]), + wp.max(MJ_MINMU, max_geom_friction[0]), + wp.max(MJ_MINMU, max_geom_friction[1]), + wp.max(MJ_MINMU, max_geom_friction[2]), + wp.max(MJ_MINMU, max_geom_friction[2]), + ) + + if geom_solref[worldid, g1].x > 0.0 and geom_solref[worldid, g2].x > 0.0: + solref = mix * geom_solref[worldid, g1] + (1.0 - mix) * geom_solref[worldid, g2] + else: + solref = wp.min(geom_solref[worldid, g1], geom_solref[worldid, g2]) + + solreffriction = wp.vec2(0.0, 0.0) + + solimp = mix * geom_solimp[worldid, g1] + (1.0 - mix) * geom_solimp[worldid, g2] + + return geoms, margin, gap, condim, friction, solref, solreffriction, solimp + + +@wp.func +def _sphere_box( + # Data in: + nconmax_in: int, + # In: + sphere_pos: wp.vec3, + sphere_size: float, + box_pos: wp.vec3, + box_rot: wp.mat33, + box_size: wp.vec3, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + center = wp.transpose(box_rot) @ (sphere_pos - box_pos) + + clamped = wp.max(-box_size, wp.min(box_size, center)) + clamped_dir, dist = normalize_with_norm(clamped - center) + + if dist - sphere_size > margin: + return + + # sphere center inside box + if dist <= MJ_MINVAL: + closest = 2.0 * (box_size[0] + box_size[1] + box_size[2]) + k = wp.int32(0) + for i in range(6): + face_dist = wp.abs(wp.where(i % 2, 1.0, -1.0) * box_size[i / 2] - center[i / 2]) + if closest > face_dist: + closest = face_dist + k = i + + nearest = wp.vec3(0.0) + nearest[k / 2] = wp.where(k % 2, -1.0, 1.0) + pos = center + nearest * (sphere_size - closest) / 2.0 + contact_normal = box_rot @ nearest + contact_dist = -closest - sphere_size + + else: + deepest = center + clamped_dir * sphere_size + pos = 0.5 * (clamped + deepest) + contact_normal = box_rot @ clamped_dir + contact_dist = dist - sphere_size + + contact_pos = box_pos + box_rot @ pos + write_contact( + nconmax_in, + contact_dist, + contact_pos, + make_frame(contact_normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def sphere_box( + # Data in: + nconmax_in: int, + # In: + sphere: Geom, + box: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + _sphere_box( + nconmax_in, + sphere.pos, + sphere.size[0], + box.pos, + box.rot, + box.size, + worldid, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def capsule_box( + # Data in: + nconmax_in: int, + # In: + cap: Geom, + box: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + """Calculates contacts between a capsule and a box.""" + # Based on the mjc implementation + boxmatT = wp.transpose(box.rot) + pos = boxmatT @ (cap.pos - box.pos) + axis = boxmatT @ wp.vec3(cap.rot[0, 2], cap.rot[1, 2], cap.rot[2, 2]) + halfaxis = axis * cap.size[1] # halfaxis is the capsule direction + axisdir = wp.int32(halfaxis[0] > 0.0) + 2 * wp.int32(halfaxis[1] > 0.0) + 4 * wp.int32(halfaxis[2] > 0.0) + + bestdistmax = margin + 2.0 * (cap.size[0] + cap.size[1] + box.size[0] + box.size[1] + box.size[2]) + + # keep track of closest point + bestdist = wp.float32(bestdistmax) + bestsegmentpos = wp.float32(-12) + + # cltype: encoded collision configuration + # cltype / 3 == 0 : lower corner is closest to the capsule + # == 2 : upper corner is closest to the capsule + # == 1 : middle of the edge is closest to the capsule + # cltype % 3 == 0 : lower corner is closest to the box + # == 2 : upper corner is closest to the box + # == 1 : middle of the capsule is closest to the box + cltype = wp.int32(-4) + + # clface: index of the closest face of the box to the capsule + # -1: no face is closest (edge or corner is closest) + # 0, 1, 2: index of the axis perpendicular to the closest face + clface = wp.int32(-12) + + # first: consider cases where a face of the box is closest + for i in range(-1, 2, 2): + axisTip = pos + wp.float32(i) * halfaxis + boxPoint = wp.vec3(axisTip) + + n_out = wp.int32(0) + ax_out = wp.int32(-1) + + for j in range(3): + if boxPoint[j] < -box.size[j]: + n_out += 1 + ax_out = j + boxPoint[j] = -box.size[j] + elif boxPoint[j] > box.size[j]: + n_out += 1 + ax_out = j + boxPoint[j] = box.size[j] + + if n_out > 1: + continue + + dist = wp.length_sq(boxPoint - axisTip) + + if dist < bestdist: + bestdist = dist + bestsegmentpos = wp.float32(i) + cltype = -2 + i + clface = ax_out + + # second: consider cases where an edge of the box is closest + clcorner = wp.int32(-123) # which corner is the closest + cledge = wp.int32(-123) # which axis + bestboxpos = wp.float32(0.0) + + for i in range(8): + for j in range(3): + if i & (1 << j) != 0: + continue + + c2 = wp.int32(-123) + + # box_pt is the starting point (corner) on the box + box_pt = wp.cw_mul( + wp.vec3( + wp.where(i & 1, 1.0, -1.0), + wp.where(i & 2, 1.0, -1.0), + wp.where(i & 4, 1.0, -1.0), + ), + box.size, + ) + box_pt[j] = 0.0 + + # find closest point between capsule and the edge + dif = box_pt - pos + + u = -box.size[j] * dif[j] + v = wp.dot(halfaxis, dif) + ma = box.size[j] * box.size[j] + mb = -box.size[j] * halfaxis[j] + mc = cap.size[1] * cap.size[1] + det = ma * mc - mb * mb + if wp.abs(det) < MJ_MINVAL: + continue + + idet = 1.0 / det + # sX : X=1 means middle of segment. X=0 or 2 one or the other end + + x1 = wp.float32((mc * u - mb * v) * idet) + x2 = wp.float32((ma * v - mb * u) * idet) + + s1 = wp.int32(1) + s2 = wp.int32(1) + + if x1 > 1: + x1 = 1.0 + s1 = 2 + x2 = (v - mb) / mc + elif x1 < -1: + x1 = -1.0 + s1 = 0 + x2 = (v + mb) / mc + + x2_over = x2 > 1.0 + if x2_over or x2 < -1.0: + if x2_over: + x2 = 1.0 + s2 = 2 + x1 = (u - mb) / ma + else: + x2 = -1.0 + s2 = 0 + x1 = (u + mb) / ma + + if x1 > 1: + x1 = 1.0 + s1 = 2 + elif x1 < -1: + x1 = -1.0 + s1 = 0 + + dif -= halfaxis * x2 + dif[j] += box.size[j] * x1 + + # encode relative positions of the closest points + ct = s1 * 3 + s2 + + dif_sq = wp.length_sq(dif) + if dif_sq < bestdist - MJ_MINVAL: + bestdist = dif_sq + bestsegmentpos = x2 + bestboxpos = x1 + # ct<6 means closest point on box is at lower end or middle of edge + c2 = ct / 6 + + clcorner = i + (1 << j) * c2 # index of closest box corner + cledge = j # axis index of closest box edge + cltype = ct # encoded collision configuration + + best = wp.float32(0.0) + + p = wp.vec2(pos.x, pos.y) + dd = wp.vec2(halfaxis.x, halfaxis.y) + s = wp.vec2(box.size.x, box.size.y) + secondpos = wp.float32(-4.0) + + uu = dd.x * s.y + vv = dd.y * s.x + w_neg = dd.x * p.y - dd.y * p.x < 0 + + best = wp.float32(-1.0) + + ee1 = uu - vv + ee2 = uu + vv + + if wp.abs(ee1) > best: + best = wp.abs(ee1) + c1 = wp.where((ee1 < 0) == w_neg, 0, 3) + + if wp.abs(ee2) > best: + best = wp.abs(ee2) + c1 = wp.where((ee2 > 0) == w_neg, 1, 2) + + if cltype == -4: # invalid type + return + + if cltype >= 0 and cltype / 3 != 1: # closest to a corner of the box + c1 = axisdir ^ clcorner + # Calculate relative orientation between capsule and corner + # There are two possible configurations: + # 1. Capsule axis points toward/away from corner + # 2. Capsule axis aligns with a face or edge + if c1 != 0 and c1 != 7: # create second contact point + if c1 == 1 or c1 == 2 or c1 == 4: + mul = 1 + else: + mul = -1 + c1 = 7 - c1 + + # "de" and "dp" distance from first closest point on the capsule to both ends of it + # mul is a direction along the capsule's axis + + if c1 == 1: + ax = 0 + ax1 = 1 + ax2 = 2 + elif c1 == 2: + ax = 1 + ax1 = 2 + ax2 = 0 + elif c1 == 4: + ax = 2 + ax1 = 0 + ax2 = 1 + + if axis[ax] * axis[ax] > 0.5: # second point along the edge of the box + m = 2.0 * box.size[ax] / wp.abs(halfaxis[ax]) + secondpos = min(1.0 - wp.float32(mul) * bestsegmentpos, m) + else: # second point along a face of the box + # check for overshoot again + m = 2.0 * min( + box.size[ax1] / wp.abs(halfaxis[ax1]), + box.size[ax2] / wp.abs(halfaxis[ax2]), + ) + secondpos = -min(1.0 + wp.float32(mul) * bestsegmentpos, m) + secondpos *= wp.float32(mul) + + elif cltype >= 0 and cltype / 3 == 1: # we are on box's edge + # Calculate relative orientation between capsule and edge + # Two possible configurations: + # - T configuration: c1 = 2^n (no additional contacts) + # - X configuration: c1 != 2^n (potential additional contacts) + c1 = axisdir ^ clcorner + c1 &= 7 - (1 << cledge) # mask out edge axis to determine configuration + + if c1 == 1 or c1 == 2 or c1 == 4: # create second contact point + if cledge == 0: + ax1 = 1 + ax2 = 2 + if cledge == 1: + ax1 = 2 + ax2 = 0 + if cledge == 2: + ax1 = 0 + ax2 = 1 + ax = cledge + + # find which face the capsule has a lower angle, and switch the axis + if wp.abs(axis[ax1]) > wp.abs(axis[ax2]): + ax1 = ax2 + ax2 = 3 - ax - ax1 + + # mul determines direction along capsule axis for second contact point + if c1 & (1 << ax2): + mul = 1 + secondpos = 1.0 - bestsegmentpos + else: + mul = -1 + secondpos = 1.0 + bestsegmentpos + + # now find out whether we point towards the opposite side or towards one of the sides + # and also find the farthest point along the capsule that is above the box + + e1 = 2.0 * box.size[ax2] / wp.abs(halfaxis[ax2]) + secondpos = min(e1, secondpos) + + if ((axisdir & (1 << ax)) != 0) == ((c1 & (1 << ax2)) != 0): + e2 = 1.0 - bestboxpos + else: + e2 = 1.0 + bestboxpos + + e1 = box.size[ax] * e2 / wp.abs(halfaxis[ax]) + + secondpos = min(e1, secondpos) + secondpos *= wp.float32(mul) + + elif cltype < 0: + # similarly we handle the case when one capsule's end is closest to a face of the box + # and find where is the other end pointing to and clamping to the farthest point + # of the capsule that's above the box + # if the closest point is inside the box there's no need for a second point + + if clface != -1: # create second contact point + mul = wp.where(cltype == -3, 1, -1) + secondpos = 2.0 + + tmp1 = pos - halfaxis * wp.float32(mul) + + for i in range(3): + if i != clface: + ha_r = wp.float32(mul) / halfaxis[i] + e1 = (box.size[i] - tmp1[i]) * ha_r + if 0 < e1 and e1 < secondpos: + secondpos = e1 + + e1 = (-box.size[i] - tmp1[i]) * ha_r + if 0 < e1 and e1 < secondpos: + secondpos = e1 + + secondpos *= wp.float32(mul) + + # create sphere in original orientation at first contact point + s1_pos_l = pos + halfaxis * bestsegmentpos + s1_pos_g = box.rot @ s1_pos_l + box.pos + + # collide with sphere + _sphere_box( + nconmax_in, + s1_pos_g, + cap.size[0], + box.pos, + box.rot, + box.size, + worldid, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + if secondpos > -3: # secondpos was modified + s2_pos_l = pos + halfaxis * (secondpos + bestsegmentpos) + s2_pos_g = box.rot @ s2_pos_l + box.pos + _sphere_box( + nconmax_in, + s2_pos_g, + cap.size[0], + box.pos, + box.rot, + box.size, + worldid, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@wp.func +def _compute_rotmore(face_idx: int) -> wp.mat33: + rotmore = wp.mat33(0.0) + + if face_idx == 0: + rotmore[0, 2] = -1.0 + rotmore[1, 1] = +1.0 + rotmore[2, 0] = +1.0 + elif face_idx == 1: + rotmore[0, 0] = +1.0 + rotmore[1, 2] = -1.0 + rotmore[2, 1] = +1.0 + elif face_idx == 2: + rotmore[0, 0] = +1.0 + rotmore[1, 1] = +1.0 + rotmore[2, 2] = +1.0 + elif face_idx == 3: + rotmore[0, 2] = +1.0 + rotmore[1, 1] = +1.0 + rotmore[2, 0] = -1.0 + elif face_idx == 4: + rotmore[0, 0] = +1.0 + rotmore[1, 2] = +1.0 + rotmore[2, 1] = -1.0 + elif face_idx == 5: + rotmore[0, 0] = -1.0 + rotmore[1, 1] = +1.0 + rotmore[2, 2] = -1.0 + + return rotmore + + +@wp.func +def box_box( + # Data in: + nconmax_in: int, + # In: + box1: Geom, + box2: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2f, + solreffriction: wp.vec2f, + solimp: vec5, + geoms: wp.vec2i, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + # Compute transforms between box's frames + + pos21 = wp.transpose(box1.rot) @ (box2.pos - box1.pos) + pos12 = wp.transpose(box2.rot) @ (box1.pos - box2.pos) + + rot21 = wp.transpose(box1.rot) @ box2.rot + rot12 = wp.transpose(rot21) + + rot21abs = wp.matrix_from_rows(wp.abs(rot21[0]), wp.abs(rot21[1]), wp.abs(rot21[2])) + rot12abs = wp.transpose(rot21abs) + + plen2 = rot21abs @ box2.size + plen1 = rot12abs @ box1.size + + # Compute axis of maximum separation + s_sum_3 = 3.0 * (box1.size + box2.size) + separation = wp.float32(margin + s_sum_3[0] + s_sum_3[1] + s_sum_3[2]) + axis_code = wp.int32(-1) + + # First test: consider boxes' face normals + for i in range(3): + c1 = -wp.abs(pos21[i]) + box1.size[i] + plen2[i] + + c2 = -wp.abs(pos12[i]) + box2.size[i] + plen1[i] + + if c1 < -margin or c2 < -margin: + return + + if c1 < separation: + separation = c1 + axis_code = i + 3 * wp.int32(pos21[i] < 0) + 0 # Face of box1 + if c2 < separation: + separation = c2 + axis_code = i + 3 * wp.int32(pos12[i] < 0) + 6 # Face of box2 + + clnorm = wp.vec3(0.0) + inv = wp.bool(False) + cle1 = wp.int32(0) + cle2 = wp.int32(0) + + # Second test: consider cross products of boxes' edges + for i in range(3): + for j in range(3): + # Compute cross product of box edges (potential separating axis) + if i == 0: + cross_axis = wp.vec3(0.0, -rot12[j, 2], rot12[j, 1]) + elif i == 1: + cross_axis = wp.vec3(rot12[j, 2], 0.0, -rot12[j, 0]) + else: + cross_axis = wp.vec3(-rot12[j, 1], rot12[j, 0], 0.0) + + cross_length = wp.length(cross_axis) + if cross_length < MJ_MINVAL: + continue + + cross_axis /= cross_length + + box_dist = wp.dot(pos21, cross_axis) + c3 = wp.float32(0.0) + + # Project box half-sizes onto the potential separating axis + for k in range(3): + if k != i: + c3 += box1.size[k] * wp.abs(cross_axis[k]) + if k != j: + c3 += box2.size[k] * rot21abs[i, 3 - k - j] / cross_length + + c3 -= wp.abs(box_dist) + + # Early exit: no collision if separated along this axis + if c3 < -margin: + return + + # Track minimum separation and which edge-edge pair it occurs on + if c3 < separation * (1.0 - 1e-12): + separation = c3 + # Determine which corners/edges are closest + cle1 = 0 + cle2 = 0 + + for k in range(3): + if k != i and (int(cross_axis[k] > 0) ^ int(box_dist < 0)): + cle1 += 1 << k + if k != j: + if int(rot21[i, 3 - k - j] > 0) ^ int(box_dist < 0) ^ int((k - j + 3) % 3 == 1): + cle2 += 1 << k + + axis_code = 12 + i * 3 + j + clnorm = cross_axis + inv = box_dist < 0 + + # No axis with separation < margin found + if axis_code == -1: + return + + points = mat83f() + depth = vec8f() + max_con_pair = 8 + # 8 contacts should suffice for most configurations + + if axis_code < 12: + # Handle face-vertex collision + face_idx = axis_code % 6 + box_idx = axis_code / 6 + rotmore = _compute_rotmore(face_idx) + + r = rotmore @ wp.where(box_idx, rot12, rot21) + p = rotmore @ wp.where(box_idx, pos12, pos21) + ss = wp.abs(rotmore @ wp.where(box_idx, box2.size, box1.size)) + s = wp.where(box_idx, box1.size, box2.size) + rt = wp.transpose(r) + + lx, ly, hz = ss[0], ss[1], ss[2] + p[2] -= hz + + clcorner = wp.int32(0) # corner of non-face box with least axis separation + + for i in range(3): + if r[2, i] < 0: + clcorner += 1 << i + + lp = p + for i in range(wp.static(3)): + lp += rt[i] * s[i] * wp.where(clcorner & 1 << i, 1.0, -1.0) + + m = wp.int32(1) + dirs = wp.int32(0) + + cn1 = wp.vec3(0.0) + cn2 = wp.vec3(0.0) + + for i in range(3): + if wp.abs(r[2, i]) < 0.5: + if not dirs: + cn1 = rt[i] * s[i] * wp.where(clcorner & (1 << i), -2.0, 2.0) + else: + cn2 = rt[i] * s[i] * wp.where(clcorner & (1 << i), -2.0, 2.0) + + dirs += 1 + + k = dirs * dirs + + # Find potential contact points + + n = wp.int32(0) + + for i in range(k): + for q in range(2): + # lines_a and lines_b (lines between corners) computed on the fly + lav = lp + wp.where(i < 2, wp.vec3(0.0), wp.where(i == 2, cn1, cn2)) + lbv = wp.where(i == 0 or i == 3, cn1, cn2) + + if wp.abs(lbv[q]) > MJ_MINVAL: + br = 1.0 / lbv[q] + for j in range(-1, 2, 2): + l = ss[q] * wp.float32(j) + c1 = (l - lav[q]) * br + if c1 < 0 or c1 > 1: + continue + c2 = lav[1 - q] + lbv[1 - q] * c1 + if wp.abs(c2) > ss[1 - q]: + continue + + points[n] = lav + c1 * lbv + n += 1 + + if dirs == 2: + ax = cn1[0] + bx = cn2[0] + ay = cn1[1] + by = cn2[1] + C = 1.0 / (ax * by - bx * ay) + + for i in range(4): + llx = wp.where(i / 2, lx, -lx) + lly = wp.where(i % 2, ly, -ly) + + x = llx - lp[0] + y = lly - lp[1] + + u = (x * by - y * bx) * C + v = (y * ax - x * ay) * C + + if u > 0 and v > 0 and u < 1 and v < 1: + points[n] = wp.vec3(llx, lly, lp[2] + u * cn1[2] + v * cn2[2]) + n += 1 + + for i in range(1 << dirs): + tmpv = lp + wp.float32(i & 1) * cn1 + wp.float32((i & 2) != 0) * cn2 + if tmpv[0] > -lx and tmpv[0] < lx and tmpv[1] > -ly and tmpv[1] < ly: + points[n] = tmpv + n += 1 + + m = n + n = wp.int32(0) + + for i in range(m): + if points[i][2] > margin: + continue + if i != n: + points[n] = points[i] + + points[n, 2] *= 0.5 + depth[n] = points[n, 2] + n += 1 + + # Set up contact frame + rw = wp.where(box_idx, box2.rot, box1.rot) @ wp.transpose(rotmore) + pw = wp.where(box_idx, box2.pos, box1.pos) + normal = wp.where(box_idx, -1.0, 1.0) * wp.transpose(rw)[2] + + else: + # Handle edge-edge collision + edge1 = (axis_code - 12) / 3 + edge2 = (axis_code - 12) % 3 + + # Set up non-contacting edges ax1, ax2 for box2 and pax1, pax2 for box 1 + ax1 = wp.int(1 - (edge2 & 1)) + ax2 = wp.int(2 - (edge2 & 2)) + + pax1 = wp.int(1 - (edge1 & 1)) + pax2 = wp.int(2 - (edge1 & 2)) + + if rot21abs[edge1, ax1] < rot21abs[edge1, ax2]: + ax1, ax2 = ax2, ax1 + + if rot12abs[edge2, pax1] < rot12abs[edge2, pax2]: + pax1, pax2 = pax2, pax1 + + rotmore = _compute_rotmore(wp.where(cle1 & (1 << pax2), pax2, pax2 + 3)) + + # Transform coordinates for edge-edge contact calculation + p = rotmore @ pos21 + rnorm = rotmore @ clnorm + r = rotmore @ rot21 + rt = wp.transpose(r) + s = wp.abs(wp.transpose(rotmore) @ box1.size) + + lx, ly, hz = s[0], s[1], s[2] + p[2] -= hz + + # Calculate closest box2 face + + points[0] = ( + p + + rt[ax1] * box2.size[ax1] * wp.where(cle2 & (1 << ax1), 1.0, -1.0) + + rt[ax2] * box2.size[ax2] * wp.where(cle2 & (1 << ax2), 1.0, -1.0) + ) + points[1] = points[0] - rt[edge2] * box2.size[edge2] + points[0] += rt[edge2] * box2.size[edge2] + + points[2] = ( + p + + rt[ax1] * box2.size[ax1] * wp.where(cle2 & (1 << ax1), -1.0, 1.0) + + rt[ax2] * box2.size[ax2] * wp.where(cle2 & (1 << ax2), 1.0, -1.0) + ) + + points[3] = points[2] - rt[edge2] * box2.size[edge2] + points[2] += rt[edge2] * box2.size[edge2] + + n = 4 + + # Set up coordinate axes for contact face of box2 + axi_lp = points[0] + axi_cn1 = points[1] - points[0] + axi_cn2 = points[2] - points[0] + + # Check if contact normal is valid + if wp.abs(rnorm[2]) < MJ_MINVAL: + return # Shouldn't happen + + # Calculate inverse normal for projection + innorm = wp.where(inv, -1.0, 1.0) / rnorm[2] + + pu = mat43f() + + # Project points onto contact plane + for i in range(4): + pu[i] = points[i] + c_scl = points[i, 2] * wp.where(inv, -1.0, 1.0) * innorm + points[i] -= rnorm * c_scl + + pts_lp = points[0] + pts_cn1 = points[1] - points[0] + pts_cn2 = points[2] - points[0] + + n = wp.int32(0) + + for i in range(4): + for q in range(2): + la = pts_lp[q] + wp.where(i < 2, 0.0, wp.where(i == 2, pts_cn1[q], pts_cn2[q])) + lb = wp.where(i == 0 or i == 3, pts_cn1[q], pts_cn2[q]) + lc = pts_lp[1 - q] + wp.where(i < 2, 0.0, wp.where(i == 2, pts_cn1[1 - q], pts_cn2[1 - q])) + ld = wp.where(i == 0 or i == 3, pts_cn1[1 - q], pts_cn2[1 - q]) + + # linesu_a and linesu_b (lines between corners) computed on the fly + lua = axi_lp + wp.where(i < 2, wp.vec3(0.0), wp.where(i == 2, axi_cn1, axi_cn2)) + lub = wp.where(i == 0 or i == 3, axi_cn1, axi_cn2) + + if wp.abs(lb) > MJ_MINVAL: + br = 1.0 / lb + for j in range(-1, 2, 2): + if n == max_con_pair: + break + l = s[q] * wp.float32(j) + c1 = (l - la) * br + if c1 < 0 or c1 > 1: + continue + c2 = lc + ld * c1 + if wp.abs(c2) > s[1 - q]: + continue + if (lua[2] + lub[2] * c1) * innorm > margin: + continue + + points[n] = lua * 0.5 + c1 * lub * 0.5 + points[n, q] += 0.5 * l + points[n, 1 - q] += 0.5 * c2 + depth[n] = points[n, 2] * innorm * 2.0 + n += 1 + + nl = n + + ax = pts_cn1[0] + bx = pts_cn2[0] + ay = pts_cn1[1] + by = pts_cn2[1] + C = 1.0 / (ax * by - bx * ay) + + for i in range(4): + if n == max_con_pair: + break + llx = wp.where(i / 2, lx, -lx) + lly = wp.where(i % 2, ly, -ly) + + x = llx - pts_lp[0] + y = lly - pts_lp[1] + + u = (x * by - y * bx) * C + v = (y * ax - x * ay) * C + + if nl == 0: + if (u < 0 or u > 0) and (v < 0 or v > 1): + continue + elif u < 0 or v < 0 or u > 1 or v > 1: + continue + + u = wp.clamp(u, 0.0, 1.0) + v = wp.clamp(v, 0.0, 1.0) + w = 1.0 - u - v + vtmp = pu[0] * w + pu[1] * u + pu[2] * v + + points[n] = wp.vec3(llx, lly, 0.0) + + vtmp2 = points[n] - vtmp + tc1 = wp.length_sq(vtmp2) + if vtmp[2] > 0 and tc1 > margin * margin: + continue + + points[n] = 0.5 * (points[n] + vtmp) + + depth[n] = wp.sqrt(tc1) * wp.where(vtmp[2] < 0, -1.0, 1.0) + n += 1 + + nf = n + + for i in range(4): + if n >= max_con_pair: + break + x = pu[i, 0] + y = pu[i, 1] + if nl == 0 and nf != 0: + if (x < -lx or x > lx) and (y < -ly or y > ly): + continue + elif x < -lx or x > lx or y < -ly or y > ly: + continue + + c1 = wp.float32(0) + + for j in range(2): + if pu[i, j] < -s[j]: + c1 += (pu[i, j] + s[j]) * (pu[i, j] + s[j]) + elif pu[i, j] > s[j]: + c1 += (pu[i, j] - s[j]) * (pu[i, j] - s[j]) + + c1 += pu[i, 2] * innorm * pu[i, 2] * innorm + + if pu[i, 2] > 0 and c1 > margin * margin: + continue + + tmp_p = wp.vec3(pu[i, 0], pu[i, 1], 0.0) + + for j in range(2): + if pu[i, j] < -s[j]: + tmp_p[j] = -s[j] * 0.5 + elif pu[i, j] > s[j]: + tmp_p[j] = +s[j] * 0.5 + + tmp_p += pu[i] + points[n] = tmp_p * 0.5 + + depth[n] = wp.sqrt(c1) * wp.where(pu[i, 2] < 0, -1.0, 1.0) + n += 1 + + # Set up contact data for all points + rw = box1.rot @ wp.transpose(rotmore) + pw = box1.pos + normal = wp.where(inv, -1.0, 1.0) * rw @ rnorm + + frame = make_frame(normal) + coff = wp.atomic_add(ncon_out, 0, n) + + for i in range(min(nconmax_in - coff, n)): + points[i, 2] += hz + pos = rw @ points[i] + pw + + cid = coff + i + + contact_dist_out[cid] = depth[i] + contact_pos_out[cid] = pos + contact_frame_out[cid] = frame + contact_geom_out[cid] = geoms + contact_worldid_out[cid] = worldid + contact_includemargin_out[cid] = margin - gap + contact_dim_out[cid] = condim + contact_friction_out[cid] = friction + contact_solref_out[cid] = solref + contact_solreffriction_out[cid] = solreffriction + contact_solimp_out[cid] = solimp + + +_PRIMITIVE_COLLISIONS = { + (GeomType.PLANE.value, GeomType.SPHERE.value): plane_sphere, + (GeomType.PLANE.value, GeomType.CAPSULE.value): plane_capsule, + (GeomType.PLANE.value, GeomType.ELLIPSOID.value): plane_ellipsoid, + (GeomType.PLANE.value, GeomType.CYLINDER.value): plane_cylinder, + (GeomType.PLANE.value, GeomType.BOX.value): plane_box, + (GeomType.PLANE.value, GeomType.MESH.value): plane_convex, + (GeomType.SPHERE.value, GeomType.SPHERE.value): sphere_sphere, + (GeomType.SPHERE.value, GeomType.CAPSULE.value): sphere_capsule, + (GeomType.SPHERE.value, GeomType.CYLINDER.value): sphere_cylinder, + (GeomType.SPHERE.value, GeomType.BOX.value): sphere_box, + (GeomType.CAPSULE.value, GeomType.CAPSULE.value): capsule_capsule, + (GeomType.CAPSULE.value, GeomType.BOX.value): capsule_box, + (GeomType.BOX.value, GeomType.BOX.value): box_box, +} + + +# TODO(team): _check_collisions shared utility +def _check_primitive_collisions(): + prev_idx = -1 + for types in _PRIMITIVE_COLLISIONS.keys(): + idx = upper_trid_index(len(GeomType), types[0], types[1]) + if types[1] < types[0] or idx <= prev_idx: + return False + prev_idx = idx + return True + + +assert _check_primitive_collisions(), "_PRIMITIVE_COLLISIONS is in invalid order" + +_primitive_collisions_types = [] +_primitive_collisions_func = [] + + +def _primitive_narrowphase_builder(m: Model): + for types, func in _PRIMITIVE_COLLISIONS.items(): + idx = upper_trid_index(len(GeomType), types[0], types[1]) + if m.geom_pair_type_count[idx] and types not in _primitive_collisions_types: + _primitive_collisions_types.append(types) + _primitive_collisions_func.append(func) + + @wp.kernel + def _primitive_narrowphase( + # Model: + geom_type: wp.array(dtype=int), + geom_condim: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_priority: wp.array(dtype=int), + geom_solmix: wp.array2d(dtype=float), + geom_solref: wp.array2d(dtype=wp.vec2), + geom_solimp: wp.array2d(dtype=vec5), + geom_size: wp.array2d(dtype=wp.vec3), + geom_friction: wp.array2d(dtype=wp.vec3), + geom_margin: wp.array2d(dtype=float), + geom_gap: wp.array2d(dtype=float), + hfield_adr: wp.array(dtype=int), + hfield_nrow: wp.array(dtype=int), + hfield_ncol: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + hfield_data: wp.array(dtype=float), + mesh_vertadr: wp.array(dtype=int), + mesh_vertnum: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_graphadr: wp.array(dtype=int), + mesh_graph: wp.array(dtype=int), + mesh_polynum: wp.array(dtype=int), + mesh_polyadr: wp.array(dtype=int), + mesh_polynormal: wp.array(dtype=wp.vec3), + mesh_polyvertadr: wp.array(dtype=int), + mesh_polyvertnum: wp.array(dtype=int), + mesh_polyvert: wp.array(dtype=int), + mesh_polymapadr: wp.array(dtype=int), + mesh_polymapnum: wp.array(dtype=int), + mesh_polymap: wp.array(dtype=int), + pair_dim: wp.array(dtype=int), + pair_solref: wp.array2d(dtype=wp.vec2), + pair_solreffriction: wp.array2d(dtype=wp.vec2), + pair_solimp: wp.array2d(dtype=vec5), + pair_margin: wp.array2d(dtype=float), + pair_gap: wp.array2d(dtype=float), + pair_friction: wp.array2d(dtype=vec5), + # Data in: + nconmax_in: int, + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + collision_pair_in: wp.array(dtype=wp.vec2i), + collision_hftri_index_in: wp.array(dtype=int), + collision_pairid_in: wp.array(dtype=int), + collision_worldid_in: wp.array(dtype=int), + ncollision_in: wp.array(dtype=int), + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), + ): + tid = wp.tid() + + if tid >= ncollision_in[0]: + return + + geoms = collision_pair_in[tid] + g1 = geoms[0] + g2 = geoms[1] + + type1 = geom_type[g1] + type2 = geom_type[g2] + + worldid = collision_worldid_in[tid] + + _, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params( + geom_condim, + geom_priority, + geom_solmix, + geom_solref, + geom_solimp, + geom_friction, + geom_margin, + geom_gap, + pair_dim, + pair_solref, + pair_solreffriction, + pair_solimp, + pair_margin, + pair_gap, + pair_friction, + collision_pair_in, + collision_pairid_in, + tid, + worldid, + ) + + hftri_index = collision_hftri_index_in[tid] + + geom1 = _geom( + geom_type, + geom_dataid, + geom_size, + hfield_adr, + hfield_nrow, + hfield_ncol, + hfield_size, + hfield_data, + mesh_vertadr, + mesh_vertnum, + mesh_vert, + mesh_graphadr, + mesh_graph, + mesh_polynum, + mesh_polyadr, + mesh_polynormal, + mesh_polyvertadr, + mesh_polyvertnum, + mesh_polyvert, + mesh_polymapadr, + mesh_polymapnum, + mesh_polymap, + geom_xpos_in, + geom_xmat_in, + worldid, + g1, + hftri_index, + ) + + geom2 = _geom( + geom_type, + geom_dataid, + geom_size, + hfield_adr, + hfield_nrow, + hfield_ncol, + hfield_size, + hfield_data, + mesh_vertadr, + mesh_vertnum, + mesh_vert, + mesh_graphadr, + mesh_graph, + mesh_polynum, + mesh_polyadr, + mesh_polynormal, + mesh_polyvertadr, + mesh_polyvertnum, + mesh_polyvert, + mesh_polymapadr, + mesh_polymapnum, + mesh_polymap, + geom_xpos_in, + geom_xmat_in, + worldid, + g2, + hftri_index, + ) + + for i in range(wp.static(len(_primitive_collisions_func))): + collision_type1 = wp.static(_primitive_collisions_types[i][0]) + collision_type2 = wp.static(_primitive_collisions_types[i][1]) + + if collision_type1 == type1 and collision_type2 == type2: + wp.static(_primitive_collisions_func[i])( + nconmax_in, + geom1, + geom2, + worldid, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + return _primitive_narrowphase + + +@event_scope +def primitive_narrowphase(m: Model, d: Data): + """Runs collision detection on primitive geom pairs discovered during broadphase. + + This function processes collision pairs involving primitive shapes that were + identified during the broadphase stage. It computes detailed contact information + such as distance, position, and frame, and populates the `d.contact` array. + + The primitive geom types handled are PLANE, SPHERE, CAPSULE, CYLINDER, BOX. + + It also handles collisions between planes and convex hulls. + + To improve performance, it dynamically builds and launches a kernel tailored to + the specific primitive collision types present in the model, avoiding + unnecessary checks for non-existent collision pairs. + """ + # we need to figure out how to keep the overhead of this small - not launching anything + # for pair types without collisions, as well as updating the launch dimensions. + wp.launch( + _primitive_narrowphase_builder(m), + dim=d.nconmax, + inputs=[ + m.geom_type, + m.geom_condim, + m.geom_dataid, + m.geom_priority, + m.geom_solmix, + m.geom_solref, + m.geom_solimp, + m.geom_size, + m.geom_friction, + m.geom_margin, + m.geom_gap, + m.hfield_adr, + m.hfield_nrow, + m.hfield_ncol, + m.hfield_size, + m.hfield_data, + m.mesh_vertadr, + m.mesh_vertnum, + m.mesh_vert, + m.mesh_graphadr, + m.mesh_graph, + m.mesh_polynum, + m.mesh_polyadr, + m.mesh_polynormal, + m.mesh_polyvertadr, + m.mesh_polyvertnum, + m.mesh_polyvert, + m.mesh_polymapadr, + m.mesh_polymapnum, + m.mesh_polymap, + m.pair_dim, + m.pair_solref, + m.pair_solreffriction, + m.pair_solimp, + m.pair_margin, + m.pair_gap, + m.pair_friction, + d.nconmax, + d.geom_xpos, + d.geom_xmat, + d.collision_pair, + d.collision_hftri_index, + d.collision_pairid, + d.collision_worldid, + d.ncollision, + ], + outputs=[ + d.ncon, + d.contact.dist, + d.contact.pos, + d.contact.frame, + d.contact.includemargin, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + d.contact.dim, + d.contact.geom, + d.contact.worldid, + ], + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py new file mode 100644 index 00000000..88457e91 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py @@ -0,0 +1,599 @@ +# Copyright 2025 The Physics-Next Project Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +from typing import Tuple + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src import math +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import _geom +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import contact_params +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import write_contact +from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 +from mujoco.mjx.third_party.mujoco_warp._src.util_misc import halton +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope + + +@wp.struct +class OptimizationParams: + rel_mat: wp.mat33 + rel_pos: wp.vec3 + attr1: wp.vec3 + attr2: wp.vec3 + + +@wp.struct +class AABB: + min: wp.vec3 + max: wp.vec3 + + +@wp.func +def transform_aabb(aabb_pos: wp.vec3, aabb_size: wp.vec3, pos: wp.vec3, ori: wp.mat33) -> AABB: + aabb = AABB() + aabb.max = wp.vec3(-1000000000.0, -1000000000.0, -1000000000.0) + aabb.min = wp.vec3(1000000000.0, 1000000000.0, 1000000000.0) + + for i in range(8): + vec = wp.vec3( + aabb_size.x * (1.0 if (i & 1) else -1.0), + aabb_size.y * (1.0 if (i & 2) else -1.0), + aabb_size.z * (1.0 if (i & 4) else -1.0), + ) + + frame_vec = ori * (vec + aabb_pos) + pos + + aabb.min = wp.min(aabb.min, frame_vec) + aabb.max = wp.max(aabb.max, frame_vec) + + return aabb + + +@wp.func +def sphere(p: wp.vec3, size: wp.vec3) -> float: + return wp.length(p) - size[0] + + +@wp.func +def ellipsoid(p: wp.vec3, size: wp.vec3) -> float: + scaled_p = wp.vec3(p[0] / size[0], p[1] / size[1], p[2] / size[2]) + k0 = wp.length(scaled_p) + k1 = wp.length(wp.vec3(p[0] / (size[0] ** 2.0), p[1] / (size[1] ** 2.0), p[2] / (size[2] ** 2.0))) + if k1 != 0.0: + denom = k1 + else: + denom = 1e-12 + return k0 * (k0 - 1.0) / denom + + +@wp.func +def grad_sphere(p: wp.vec3) -> wp.vec3: + c = wp.length(p) + if c > 1e-9: + return p / c + else: + wp.vec3(0.0) + + +@wp.func +def grad_ellipsoid(p: wp.vec3, size: wp.vec3) -> wp.vec3: + a = wp.vec3(p[0] / size[0], p[1] / size[1], p[2] / size[2]) + + b = wp.vec3(a[0] / size[0], a[1] / size[1], a[2] / size[2]) + k0 = wp.length(a) + k1 = wp.length(b) + invK0 = 1.0 / k0 + invK1 = 1.0 / k1 + + gk0 = b * invK0 + gk1 = wp.vec3( + b[0] * invK1 / (size[0] * size[0]), + b[1] * invK1 / (size[1] * size[1]), + b[2] * invK1 / (size[2] * size[2]), + ) + df_dk0 = (2.0 * k0 - 1.0) * invK1 + df_dk1 = k0 * (k0 - 1.0) * invK1 * invK1 + + raw_grad = gk0 * df_dk0 - gk1 * df_dk1 + return raw_grad / wp.length(raw_grad) + + +@wp.func +def user_sdf(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> float: + wp.printf("ERROR: user_sdf function must be implemented by user code\n") + return 0.0 + + +@wp.func +def user_sdf_grad(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> wp.vec3: + wp.printf("ERROR: user_sdf_grad function must be implemented by user code\n") + return wp.vec3(0.0) + + +@wp.func +def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int) -> float: + if type == int(GeomType.SPHERE.value): + return sphere(p, attr) + elif type == int(GeomType.ELLIPSOID.value): + return ellipsoid(p, attr) + elif type == int(GeomType.SDF.value): + return user_sdf(p, attr, sdf_type) + wp.printf("ERROR: SDF type not implemented\n") + return 0.0 + + +@wp.func +def sdf_grad(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int) -> wp.vec3: + if type == int(GeomType.SPHERE.value): + return grad_sphere(p) + elif type == int(GeomType.ELLIPSOID.value): + return grad_ellipsoid(p, attr) + elif type == int(GeomType.SDF.value): + return user_sdf_grad(p, attr, sdf_type) + wp.printf("ERROR: SDF grad type not implemented\n") + return wp.vec3(0.0) + + +@wp.func +def clearance( + type1: int, p1: wp.vec3, p2: wp.vec3, s1: wp.vec3, s2: wp.vec3, sdf_type1: int, sdf_type2: int, sfd_intersection: bool +) -> float: + sdf1 = sdf(type1, p1, s1, sdf_type1) + sdf2 = sdf(int(GeomType.SDF.value), p2, s2, sdf_type2) + if sfd_intersection: + return wp.max(sdf1, sdf2) + else: + return sdf1 + sdf2 + wp.abs(wp.max(sdf1, sdf2)) + + +@wp.func +def compute_grad( + type1: int, p1: wp.vec3, p2: wp.vec3, params: OptimizationParams, sdf_type1: int, sdf_type2: int, sfd_intersection: bool +) -> wp.vec3: + A = sdf(type1, p1, params.attr1, sdf_type1) + B = sdf(int(GeomType.SDF.value), p2, params.attr2, sdf_type2) + grad1 = sdf_grad(type1, p1, params.attr1, sdf_type1) + grad2 = sdf_grad(int(GeomType.SDF.value), p2, params.attr2, sdf_type2) + grad1_transformed = params.rel_mat * grad1 + if sfd_intersection: + if A > B: + return grad1_transformed + else: + return grad2 + else: + gradient = grad2 + grad1_transformed + max_val = wp.max(A, B) + if A > B: + max_grad = grad1_transformed + else: + max_grad = grad2 + sign = wp.sign(max_val) + gradient += max_grad * sign + return gradient + + +@wp.func +def gradient_step( + type1: int, x: wp.vec3, params: OptimizationParams, sdf_type1: int, sdf_type2: int, niter: int, sfd_intersection: bool +) -> Tuple[float, wp.vec3]: + amin = 1e-4 + rho = 0.5 + c = 0.1 + dist = float(1e10) + + for _ in range(niter): + alpha = float(2.0) + x2 = wp.vec3(x[0], x[1], x[2]) + x1 = params.rel_mat * x2 + params.rel_pos + grad = compute_grad(type1, x1, x2, params, sdf_type1, sdf_type2, sfd_intersection) + dist0 = clearance(type1, x1, x, params.attr1, params.attr2, sdf_type1, sdf_type2, sfd_intersection) + grad_dot = wp.dot(grad, grad) + + if grad_dot < 1e-12: + return dist0, x + + wolfe = -c * alpha * grad_dot + while True: + alpha *= rho + wolfe *= rho + + x = x2 - grad * alpha + x1 = params.rel_mat * x + params.rel_pos + dist = clearance(type1, x1, x, params.attr1, params.attr2, sdf_type1, sdf_type2, sfd_intersection) + + if alpha <= amin or (dist - dist0) <= wolfe: + break + if dist > dist0: + return dist, x + return dist, x + + +@wp.func +def gradient_descent( + # In: + type1: int, + x0_initial: wp.vec3, + attr1: wp.vec3, + attr2: wp.vec3, + pos1: wp.vec3, + rot1: wp.mat33, + pos2: wp.vec3, + rot2: wp.mat33, + sdf_type1: int, + sdf_type2: int, + sdf_iterations: int, +) -> Tuple[float, wp.vec3, wp.vec3]: + params = OptimizationParams() + params.rel_mat = wp.transpose(rot1) * rot2 + params.rel_pos = wp.transpose(rot1) * (pos2 - pos1) + params.attr1 = attr1 + params.attr2 = attr2 + + # Collision phase (10 iterations, sfd_intersection=False) + dist, x = gradient_step(type1, x0_initial, params, sdf_type1, sdf_type2, sdf_iterations, False) + + # Intersection phase (1 iteration, sfd_intersection=True) + dist, x = gradient_step(type1, x, params, sdf_type1, sdf_type2, 1, True) + + # Midsurface calculation + x_1 = params.rel_mat * x + params.rel_pos + + grad1 = sdf_grad(type1, x_1, params.attr1, sdf_type1) + grad1 = wp.transpose(params.rel_mat) * grad1 + grad1 = wp.normalize(grad1) + + grad2 = sdf_grad(int(GeomType.SDF.value), x, params.attr2, sdf_type2) + grad2 = wp.normalize(grad2) + + n = grad1 - grad2 + n = wp.normalize(n) + pos = rot2 * x + pos2 + n = rot2 * n + pos3 = pos - n * dist / 2.0 + return dist, pos3, n + + +@wp.kernel +def _sdf_narrowphase( + # Model: + geom_type: wp.array(dtype=int), + geom_condim: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_priority: wp.array(dtype=int), + geom_solmix: wp.array2d(dtype=float), + geom_solref: wp.array2d(dtype=wp.vec2), + geom_solimp: wp.array2d(dtype=vec5), + geom_size: wp.array2d(dtype=wp.vec3), + geom_aabb: wp.array2d(dtype=wp.vec3), + geom_pos: wp.array2d(dtype=wp.vec3), + geom_quat: wp.array2d(dtype=wp.quat), + geom_friction: wp.array2d(dtype=wp.vec3), + geom_margin: wp.array2d(dtype=float), + geom_gap: wp.array2d(dtype=float), + hfield_adr: wp.array(dtype=int), + hfield_nrow: wp.array(dtype=int), + hfield_ncol: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + hfield_data: wp.array(dtype=float), + mesh_vertadr: wp.array(dtype=int), + mesh_vertnum: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_graphadr: wp.array(dtype=int), + mesh_graph: wp.array(dtype=int), + mesh_polynum: wp.array(dtype=int), + mesh_polyadr: wp.array(dtype=int), + mesh_polynormal: wp.array(dtype=wp.vec3), + mesh_polyvertadr: wp.array(dtype=int), + mesh_polyvertnum: wp.array(dtype=int), + mesh_polyvert: wp.array(dtype=int), + mesh_polymapadr: wp.array(dtype=int), + mesh_polymapnum: wp.array(dtype=int), + mesh_polymap: wp.array(dtype=int), + pair_dim: wp.array(dtype=int), + pair_solref: wp.array2d(dtype=wp.vec2), + pair_solreffriction: wp.array2d(dtype=wp.vec2), + pair_solimp: wp.array2d(dtype=vec5), + pair_margin: wp.array2d(dtype=float), + pair_gap: wp.array2d(dtype=float), + pair_friction: wp.array2d(dtype=vec5), + # In: + plugin: wp.array(dtype=int), + plugin_attr: wp.array(dtype=wp.vec3f), + geom_plugin_index: wp.array(dtype=int), + # Data in: + nconmax_in: int, + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + collision_pair_in: wp.array(dtype=wp.vec2i), + collision_hftri_index_in: wp.array(dtype=int), + collision_pairid_in: wp.array(dtype=int), + collision_worldid_in: wp.array(dtype=int), + ncollision_in: wp.array(dtype=int), + # In: + sdf_initpoints: int, + sdf_iterations: int, + # Data out: + ncon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), +): + tid = wp.tid() + + if tid >= ncollision_in[0]: + return + + geoms = collision_pair_in[tid] + + g2 = geoms[1] + type2 = geom_type[g2] + if type2 != int(GeomType.SDF.value): + return + + worldid = collision_worldid_in[tid] + + _, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params( + geom_condim, + geom_priority, + geom_solmix, + geom_solref, + geom_solimp, + geom_friction, + geom_margin, + geom_gap, + pair_dim, + pair_solref, + pair_solreffriction, + pair_solimp, + pair_margin, + pair_gap, + pair_friction, + collision_pair_in, + collision_pairid_in, + tid, + worldid, + ) + g1 = geoms[0] + + hftri_index = collision_hftri_index_in[tid] + + geom1 = _geom( + geom_type, + geom_dataid, + geom_size, + hfield_adr, + hfield_nrow, + hfield_ncol, + hfield_size, + hfield_data, + mesh_vertadr, + mesh_vertnum, + mesh_vert, + mesh_graphadr, + mesh_graph, + mesh_polynum, + mesh_polyadr, + mesh_polynormal, + mesh_polyvertadr, + mesh_polyvertnum, + mesh_polyvert, + mesh_polymapadr, + mesh_polymapnum, + mesh_polymap, + geom_xpos_in, + geom_xmat_in, + worldid, + g1, + hftri_index, + ) + geom2 = _geom( + geom_type, + geom_dataid, + geom_size, + hfield_adr, + hfield_nrow, + hfield_ncol, + hfield_size, + hfield_data, + mesh_vertadr, + mesh_vertnum, + mesh_vert, + mesh_graphadr, + mesh_graph, + mesh_polynum, + mesh_polyadr, + mesh_polynormal, + mesh_polyvertadr, + mesh_polyvertnum, + mesh_polyvert, + mesh_polymapadr, + mesh_polymapnum, + mesh_polymap, + geom_xpos_in, + geom_xmat_in, + worldid, + g2, + hftri_index, + ) + + type1 = geom_type[g1] + g1_plugin = geom_plugin_index[g1] + g2_plugin = geom_plugin_index[g2] + + g2_to_g1_rot = wp.transpose(geom2.rot) * geom1.rot + g2_to_g1_pos = wp.transpose(geom2.rot) * (geom1.pos - geom2.pos) + + aabb_pos = geom_aabb[g1, 0] + aabb_size = geom_aabb[g1, 1] + aabb1 = transform_aabb(aabb_pos, aabb_size, g2_to_g1_pos, g2_to_g1_rot) + + aabb_pos = geom_aabb[g2, 0] + aabb_size = geom_aabb[g2, 1] + aabb2 = transform_aabb(aabb_pos, aabb_size, wp.vec3(0.0), wp.mat33(1.0)) + + aabb_intersection = AABB() + aabb_intersection.min = wp.max(aabb1.min, aabb2.min) + aabb_intersection.max = wp.min(aabb1.max, aabb2.max) + + geom_pos2 = geom_pos[worldid, g2] + quat2 = geom_quat[worldid, g2] + geom_mat2 = math.quat_to_mat(quat2) + rot2 = math.mul(geom2.rot, math.transpose(geom_mat2)) + pos2 = wp.sub(geom2.pos, math.mul(rot2, geom_pos2)) + + if type1 == int(GeomType.SDF.value): + geom_pos1 = geom_pos[worldid, g1] + quat1 = geom_quat[worldid, g1] + geom_mat1 = math.quat_to_mat(quat1) + rot1 = math.mul(geom1.rot, math.transpose(geom_mat1)) + pos1 = wp.sub(geom1.pos, math.mul(rot1, geom_pos1)) + attr1 = plugin_attr[g1_plugin] + g1_plugin_id = plugin[g1_plugin] + else: + pos1 = geom1.pos + rot1 = geom1.rot + attr1 = geom1.size + g1_plugin_id = -1 + + for i in range(sdf_initpoints): + x_g2 = wp.vec3( + aabb_intersection.min[0] + (aabb_intersection.max[0] - aabb_intersection.min[0]) * halton(i, 2), + aabb_intersection.min[1] + (aabb_intersection.max[1] - aabb_intersection.min[1]) * halton(i, 3), + aabb_intersection.min[2] + (aabb_intersection.max[2] - aabb_intersection.min[2]) * halton(i, 5), + ) + + x = geom2.rot * x_g2 + geom2.pos + x0_initial = wp.transpose(rot2) * (x - pos2) + + dist, pos, n = gradient_descent( + type1, x0_initial, attr1, plugin_attr[g2_plugin], pos1, rot1, pos2, rot2, g1_plugin_id, plugin[g2_plugin], sdf_iterations + ) + + write_contact( + nconmax_in, + dist, + pos, + make_frame(n), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + worldid, + ncon_out, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + ) + + +@event_scope +def sdf_narrowphase(m: Model, d: Data): + wp.launch( + _sdf_narrowphase, + dim=d.nconmax, + inputs=[ + m.geom_type, + m.geom_condim, + m.geom_dataid, + m.geom_priority, + m.geom_solmix, + m.geom_solref, + m.geom_solimp, + m.geom_size, + m.geom_aabb, + m.geom_pos, + m.geom_quat, + m.geom_friction, + m.geom_margin, + m.geom_gap, + m.hfield_adr, + m.hfield_nrow, + m.hfield_ncol, + m.hfield_size, + m.hfield_data, + m.mesh_vertadr, + m.mesh_vertnum, + m.mesh_vert, + m.mesh_graphadr, + m.mesh_graph, + m.mesh_polynum, + m.mesh_polyadr, + m.mesh_polynormal, + m.mesh_polyvertadr, + m.mesh_polyvertnum, + m.mesh_polyvert, + m.mesh_polymapadr, + m.mesh_polymapnum, + m.mesh_polymap, + m.pair_dim, + m.pair_solref, + m.pair_solreffriction, + m.pair_solimp, + m.pair_margin, + m.pair_gap, + m.pair_friction, + m.plugin, + m.plugin_attr, + m.geom_plugin_index, + d.nconmax, + d.geom_xpos, + d.geom_xmat, + d.collision_pair, + d.collision_hftri_index, + d.collision_pairid, + d.collision_worldid, + d.ncollision, + m.opt.sdf_initpoints, + m.opt.sdf_iterations, + ], + outputs=[ + d.ncon, + d.contact.dist, + d.contact.pos, + d.contact.frame, + d.contact.includemargin, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + d.contact.dim, + d.contact.geom, + d.contact.worldid, + ], + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py new file mode 100644 index 00000000..9ab71173 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py @@ -0,0 +1,1908 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src import math +from mujoco.mjx.third_party.mujoco_warp._src import support +from mujoco.mjx.third_party.mujoco_warp._src import types +from mujoco.mjx.third_party.mujoco_warp._src.types import ConstraintType +from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 +from mujoco.mjx.third_party.mujoco_warp._src.types import vec11 +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope + +wp.config.enable_backward = False + + +@wp.kernel +def zero_constraint_counts( + # Data out: + ne_out: wp.array(dtype=int), + ne_connect_out: wp.array(dtype=int), + ne_weld_out: wp.array(dtype=int), + ne_jnt_out: wp.array(dtype=int), + ne_ten_out: wp.array(dtype=int), + nf_out: wp.array(dtype=int), + nl_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), +): + worldid = wp.tid() + + # Zero all constraint counters + ne_out[worldid] = 0 + ne_connect_out[worldid] = 0 + ne_weld_out[worldid] = 0 + ne_jnt_out[worldid] = 0 + ne_ten_out[worldid] = 0 + nf_out[worldid] = 0 + nl_out[worldid] = 0 + nefc_out[worldid] = 0 + + +@wp.func +def _update_efc_row( + # In: + worldid: int, + timestep: float, + refsafe: int, + efcid: int, + pos_aref: float, + pos_imp: float, + invweight: float, + solref: wp.vec2, + solimp: vec5, + margin: float, + vel: float, + frictionloss: float, + type: int, + id: int, + # Data out: + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), +): + # Calculate kbi + timeconst = solref[0] + dampratio = solref[1] + dmin = solimp[0] + dmax = solimp[1] + width = solimp[2] + mid = solimp[3] + power = solimp[4] + + # TODO(team): wp.static? + if not refsafe: + timeconst = wp.max(timeconst, 2.0 * timestep) + + dmin = wp.clamp(dmin, types.MJ_MINIMP, types.MJ_MAXIMP) + dmax = wp.clamp(dmax, types.MJ_MINIMP, types.MJ_MAXIMP) + width = wp.max(types.MJ_MINVAL, width) + mid = wp.clamp(mid, types.MJ_MINIMP, types.MJ_MAXIMP) + power = wp.max(1.0, power) + + # See https://mujoco.readthedocs.io/en/latest/modeling.html#solver-parameters + k = 1.0 / (dmax * dmax * timeconst * timeconst * dampratio * dampratio) + b = 2.0 / (dmax * timeconst) + k = wp.where(solref[0] <= 0, -solref[0] / (dmax * dmax), k) + b = wp.where(solref[1] <= 0, -solref[1] / dmax, b) + + imp_x = wp.abs(pos_imp) / width + imp_a = (1.0 / wp.pow(mid, power - 1.0)) * wp.pow(imp_x, power) + imp_b = 1.0 - (1.0 / wp.pow(1.0 - mid, power - 1.0)) * wp.pow(1.0 - imp_x, power) + imp_y = wp.where(imp_x < mid, imp_a, imp_b) + imp = dmin + imp_y * (dmax - dmin) + imp = wp.clamp(imp, dmin, dmax) + imp = wp.where(imp_x > 1.0, dmax, imp) + + # Update constraints + efc_D_out[worldid, efcid] = 1.0 / wp.max(invweight * (1.0 - imp) / imp, types.MJ_MINVAL) + efc_vel_out[worldid, efcid] = vel + efc_aref_out[worldid, efcid] = -k * imp * pos_aref - b * vel + efc_pos_out[worldid, efcid] = pos_aref + margin + efc_margin_out[worldid, efcid] = margin + efc_frictionloss_out[worldid, efcid] = frictionloss + efc_type_out[worldid, efcid] = type + efc_id_out[worldid, efcid] = id + + +@wp.kernel +def _efc_equality_connect( + # Model: + nv: int, + nsite: int, + opt_timestep: wp.array(dtype=float), + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + body_invweight0: wp.array2d(dtype=wp.vec2), + dof_bodyid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + eq_obj1id: wp.array(dtype=int), + eq_obj2id: wp.array(dtype=int), + eq_objtype: wp.array(dtype=int), + eq_solref: wp.array2d(dtype=wp.vec2), + eq_solimp: wp.array2d(dtype=vec5), + eq_data: wp.array2d(dtype=vec11), + eq_connect_adr: wp.array(dtype=int), + # Data in: + njmax_in: int, + qvel_in: wp.array2d(dtype=float), + eq_active_in: wp.array2d(dtype=bool), + xpos_in: wp.array2d(dtype=wp.vec3), + xmat_in: wp.array2d(dtype=wp.mat33), + site_xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + # In: + refsafe_in: int, + # Data out: + ne_connect_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), +): + """Calculates constraint rows for connect equality constraints.""" + + worldid, i_eq_connect_adr = wp.tid() + timestep = opt_timestep[worldid] + i_eq = eq_connect_adr[i_eq_connect_adr] + + if not eq_active_in[worldid, i_eq]: + return + + wp.atomic_add(ne_connect_out, worldid, 3) + efcid = wp.atomic_add(nefc_out, worldid, 3) + + if efcid + 3 >= njmax_in: + return + + data = eq_data[worldid, i_eq] + anchor1 = wp.vec3f(data[0], data[1], data[2]) + anchor2 = wp.vec3f(data[3], data[4], data[5]) + + obj1id = eq_obj1id[i_eq] + obj2id = eq_obj2id[i_eq] + + if nsite and eq_objtype[i_eq] == wp.static(types.ObjType.SITE.value): + # body1id stores the index of site_bodyid. + body1id = site_bodyid[obj1id] + body2id = site_bodyid[obj2id] + pos1 = site_xpos_in[worldid, obj1id] + pos2 = site_xpos_in[worldid, obj2id] + else: + body1id = obj1id + body2id = obj2id + pos1 = xpos_in[worldid, body1id] + xmat_in[worldid, body1id] @ anchor1 + pos2 = xpos_in[worldid, body2id] + xmat_in[worldid, body2id] @ anchor2 + + # error is difference in global positions + pos = pos1 - pos2 + + # compute Jacobian difference (opposite of contact: 0 - 1) + Jqvel = wp.vec3f(0.0, 0.0, 0.0) + for dofid in range(nv): # TODO: parallelize + jacp1, _ = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + pos1, + body1id, + dofid, + worldid, + ) + jacp2, _ = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + pos2, + body2id, + dofid, + worldid, + ) + j1mj2 = jacp1 - jacp2 + efc_J_out[worldid, efcid + 0, dofid] = j1mj2[0] + efc_J_out[worldid, efcid + 1, dofid] = j1mj2[1] + efc_J_out[worldid, efcid + 2, dofid] = j1mj2[2] + Jqvel += j1mj2 * qvel_in[worldid, dofid] + + invweight = body_invweight0[worldid, body1id][0] + body_invweight0[worldid, body2id][0] + pos_imp = wp.length(pos) + + solref = eq_solref[worldid, i_eq] + solimp = eq_solimp[worldid, i_eq] + + for i in range(3): + efcidi = efcid + i + + _update_efc_row( + worldid, + timestep, + refsafe_in, + efcidi, + pos[i], + pos_imp, + invweight, + solref, + solimp, + 0.0, + Jqvel[i], + 0.0, + ConstraintType.EQUALITY.value, + i_eq, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + +@wp.kernel +def _efc_equality_joint( + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + qpos0: wp.array2d(dtype=float), + jnt_qposadr: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + dof_invweight0: wp.array2d(dtype=float), + eq_obj1id: wp.array(dtype=int), + eq_obj2id: wp.array(dtype=int), + eq_solref: wp.array2d(dtype=wp.vec2), + eq_solimp: wp.array2d(dtype=vec5), + eq_data: wp.array2d(dtype=vec11), + eq_jnt_adr: wp.array(dtype=int), + # Data in: + njmax_in: int, + qpos_in: wp.array2d(dtype=float), + qvel_in: wp.array2d(dtype=float), + eq_active_in: wp.array2d(dtype=bool), + # In: + refsafe_in: int, + # Data out: + ne_jnt_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), +): + worldid, i_eq_joint_adr = wp.tid() + timestep = opt_timestep[worldid] + i_eq = eq_jnt_adr[i_eq_joint_adr] + if not eq_active_in[worldid, i_eq]: + return + + wp.atomic_add(ne_jnt_out, worldid, 1) + efcid = wp.atomic_add(nefc_out, worldid, 1) + + if efcid >= njmax_in: + return + + for i in range(nv): + efc_J_out[worldid, efcid, i] = 0.0 + + jntid_1 = eq_obj1id[i_eq] + jntid_2 = eq_obj2id[i_eq] + data = eq_data[worldid, i_eq] + dofadr1 = jnt_dofadr[jntid_1] + qposadr1 = jnt_qposadr[jntid_1] + efc_J_out[worldid, efcid, dofadr1] = 1.0 + + if jntid_2 > -1: + # Two joint constraint + qposadr2 = jnt_qposadr[jntid_2] + dofadr2 = jnt_dofadr[jntid_2] + dif = qpos_in[worldid, qposadr2] - qpos0[worldid, qposadr2] + + # Horner's method for polynomials + rhs = data[0] + dif * (data[1] + dif * (data[2] + dif * (data[3] + dif * data[4]))) + deriv_2 = data[1] + dif * (2.0 * data[2] + dif * (3.0 * data[3] + dif * 4.0 * data[4])) + + pos = qpos_in[worldid, qposadr1] - qpos0[worldid, qposadr1] - rhs + Jqvel = qvel_in[worldid, dofadr1] - qvel_in[worldid, dofadr2] * deriv_2 + invweight = dof_invweight0[worldid, dofadr1] + dof_invweight0[worldid, dofadr2] + + efc_J_out[worldid, efcid, dofadr2] = -deriv_2 + else: + # Single joint constraint + pos = qpos_in[worldid, qposadr1] - qpos0[worldid, qposadr1] - data[0] + Jqvel = qvel_in[worldid, dofadr1] + invweight = dof_invweight0[worldid, dofadr1] + + # Update constraint parameters + _update_efc_row( + worldid, + timestep, + refsafe_in, + efcid, + pos, + pos, + invweight, + eq_solref[worldid, i_eq], + eq_solimp[worldid, i_eq], + 0.0, + Jqvel, + 0.0, + ConstraintType.EQUALITY.value, + i_eq, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + +@wp.kernel +def _efc_equality_tendon( + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + eq_obj1id: wp.array(dtype=int), + eq_obj2id: wp.array(dtype=int), + eq_solref: wp.array2d(dtype=wp.vec2), + eq_solimp: wp.array2d(dtype=vec5), + eq_data: wp.array2d(dtype=vec11), + eq_ten_adr: wp.array(dtype=int), + tendon_length0: wp.array2d(dtype=float), + tendon_invweight0: wp.array2d(dtype=float), + # Data in: + njmax_in: int, + qvel_in: wp.array2d(dtype=float), + eq_active_in: wp.array2d(dtype=bool), + ten_length_in: wp.array2d(dtype=float), + ten_J_in: wp.array3d(dtype=float), + # In: + refsafe_in: int, + # Data out: + ne_ten_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), +): + worldid, tenid = wp.tid() + timestep = opt_timestep[worldid] + eqid = eq_ten_adr[tenid] + + if not eq_active_in[worldid, eqid]: + return + + wp.atomic_add(ne_ten_out, worldid, 1) + efcid = wp.atomic_add(nefc_out, worldid, 1) + + if efcid >= njmax_in: + return + + obj1id = eq_obj1id[eqid] + obj2id = eq_obj2id[eqid] + + data = eq_data[worldid, eqid] + solref = eq_solref[worldid, eqid] + solimp = eq_solimp[worldid, eqid] + pos1 = ten_length_in[worldid, obj1id] - tendon_length0[worldid, obj1id] + jac1 = ten_J_in[worldid, obj1id] + + if obj2id > -1: + invweight = tendon_invweight0[worldid, obj1id] + tendon_invweight0[worldid, obj2id] + + pos2 = ten_length_in[worldid, obj2id] - tendon_length0[worldid, obj2id] + jac2 = ten_J_in[worldid, obj2id] + + dif = pos2 + dif2 = dif * dif + dif3 = dif2 * dif + dif4 = dif3 * dif + + pos = pos1 - (data[0] + data[1] * dif + data[2] * dif2 + data[3] * dif3 + data[4] * dif4) + deriv = data[1] + 2.0 * data[2] * dif + 3.0 * data[3] * dif2 + 4.0 * data[4] * dif3 + else: + invweight = tendon_invweight0[worldid, obj1id] + pos = pos1 - data[0] + deriv = 0.0 + + Jqvel = float(0.0) + for i in range(nv): + if deriv != 0.0: + J = jac1[i] + jac2[i] * -deriv + else: + J = jac1[i] + efc_J_out[worldid, efcid, i] = J + Jqvel += J * qvel_in[worldid, i] + + _update_efc_row( + worldid, + timestep, + refsafe_in, + efcid, + pos, + pos, + invweight, + solref, + solimp, + 0.0, + Jqvel, + 0.0, + ConstraintType.EQUALITY.value, + eqid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + +@wp.kernel +def _efc_friction_dof( + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + dof_invweight0: wp.array2d(dtype=float), + dof_frictionloss: wp.array2d(dtype=float), + dof_solimp: wp.array2d(dtype=vec5), + dof_solref: wp.array2d(dtype=wp.vec2), + # Data in: + njmax_in: int, + qvel_in: wp.array2d(dtype=float), + # In: + refsafe_in: int, + # Data out: + nf_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + timestep = opt_timestep[worldid] + + if dof_frictionloss[worldid, dofid] <= 0.0: + return + + efcid = wp.atomic_add(nefc_out, worldid, 1) + + if efcid >= njmax_in: + return + + wp.atomic_add(nf_out, worldid, 1) + + for i in range(nv): + efc_J_out[worldid, efcid, i] = 0.0 + + efc_J_out[worldid, efcid, dofid] = 1.0 + Jqvel = qvel_in[worldid, dofid] + + _update_efc_row( + worldid, + timestep, + refsafe_in, + efcid, + 0.0, + 0.0, + dof_invweight0[worldid, dofid], + dof_solref[worldid, dofid], + dof_solimp[worldid, dofid], + 0.0, + Jqvel, + dof_frictionloss[worldid, dofid], + ConstraintType.FRICTION_DOF.value, + dofid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + +@wp.kernel +def _efc_friction_tendon( + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + tendon_solref_fri: wp.array2d(dtype=wp.vec2), + tendon_solimp_fri: wp.array2d(dtype=vec5), + tendon_frictionloss: wp.array2d(dtype=float), + tendon_invweight0: wp.array2d(dtype=float), + # Data in: + qvel_in: wp.array2d(dtype=float), + ten_J_in: wp.array3d(dtype=float), + # In: + refsafe_in: int, + # Data out: + nf_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), +): + worldid, tenid = wp.tid() + timestep = opt_timestep[worldid] + + frictionloss = tendon_frictionloss[worldid, tenid] + if frictionloss <= 0.0: + return + + efcid = wp.atomic_add(nefc_out, worldid, 1) + wp.atomic_add(nf_out, worldid, 1) + + Jqvel = float(0.0) + + # TODO(team): parallelize + for i in range(nv): + J = ten_J_in[worldid, tenid, i] + efc_J_out[worldid, efcid, i] = J + Jqvel += J * qvel_in[worldid, i] + + _update_efc_row( + worldid, + timestep, + refsafe_in, + efcid, + 0.0, + 0.0, + tendon_invweight0[worldid, tenid], + tendon_solref_fri[worldid, tenid], + tendon_solimp_fri[worldid, tenid], + 0.0, + Jqvel, + frictionloss, + ConstraintType.FRICTION_TENDON.value, + tenid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + +@wp.kernel +def _efc_equality_weld( + # Model: + nv: int, + nsite: int, + opt_timestep: wp.array(dtype=float), + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + body_invweight0: wp.array2d(dtype=wp.vec2), + dof_bodyid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + site_quat: wp.array2d(dtype=wp.quat), + eq_obj1id: wp.array(dtype=int), + eq_obj2id: wp.array(dtype=int), + eq_objtype: wp.array(dtype=int), + eq_solref: wp.array2d(dtype=wp.vec2), + eq_solimp: wp.array2d(dtype=vec5), + eq_data: wp.array2d(dtype=vec11), + eq_wld_adr: wp.array(dtype=int), + # Data in: + njmax_in: int, + qvel_in: wp.array2d(dtype=float), + eq_active_in: wp.array2d(dtype=bool), + xpos_in: wp.array2d(dtype=wp.vec3), + xquat_in: wp.array2d(dtype=wp.quat), + xmat_in: wp.array2d(dtype=wp.mat33), + site_xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + # In: + refsafe_in: int, + # Data out: + ne_weld_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), +): + worldid, i_eq_weld_adr = wp.tid() + timestep = opt_timestep[worldid] + i_eq = eq_wld_adr[i_eq_weld_adr] + if not eq_active_in[worldid, i_eq]: + return + + wp.atomic_add(ne_weld_out, worldid, 6) + efcid = wp.atomic_add(nefc_out, worldid, 6) + + if efcid + 6 >= njmax_in: + return + + is_site = eq_objtype[i_eq] == wp.static(types.ObjType.SITE.value) and nsite > 0 + + obj1id = eq_obj1id[i_eq] + obj2id = eq_obj2id[i_eq] + + data = eq_data[worldid, i_eq] + anchor1 = wp.vec3(data[0], data[1], data[2]) + anchor2 = wp.vec3(data[3], data[4], data[5]) + relpose = wp.quat(data[6], data[7], data[8], data[9]) + torquescale = data[10] + + if is_site: + # body1id stores the index of site_bodyid. + body1id = site_bodyid[obj1id] + body2id = site_bodyid[obj2id] + pos1 = site_xpos_in[worldid, obj1id] + pos2 = site_xpos_in[worldid, obj2id] + + quat = math.mul_quat(xquat_in[worldid, body1id], site_quat[worldid, obj1id]) + quat1 = math.quat_inv(math.mul_quat(xquat_in[worldid, body2id], site_quat[worldid, obj2id])) + + else: + body1id = obj1id + body2id = obj2id + pos1 = xpos_in[worldid, body1id] + xmat_in[worldid, body1id] @ anchor2 + pos2 = xpos_in[worldid, body2id] + xmat_in[worldid, body2id] @ anchor1 + + quat = math.mul_quat(xquat_in[worldid, body1id], relpose) + quat1 = math.quat_inv(xquat_in[worldid, body2id]) + + # compute Jacobian difference (opposite of contact: 0 - 1) + Jqvelp = wp.vec3f(0.0, 0.0, 0.0) + Jqvelr = wp.vec3f(0.0, 0.0, 0.0) + + for dofid in range(nv): # TODO: parallelize + jacp1, jacr1 = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + pos1, + body1id, + dofid, + worldid, + ) + jacp2, jacr2 = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + pos2, + body2id, + dofid, + worldid, + ) + + jacdifp = jacp1 - jacp2 + for i in range(wp.static(3)): + efc_J_out[worldid, efcid + i, dofid] = jacdifp[i] + + jacdifr = (jacr1 - jacr2) * torquescale + jacdifrq = math.mul_quat(math.quat_mul_axis(quat1, jacdifr), quat) + jacdifr = 0.5 * wp.vec3(jacdifrq[1], jacdifrq[2], jacdifrq[3]) + + for i in range(wp.static(3)): + efc_J_out[worldid, efcid + 3 + i, dofid] = jacdifr[i] + + Jqvelp += jacdifp * qvel_in[worldid, dofid] + Jqvelr += jacdifr * qvel_in[worldid, dofid] + + # error is difference in global position and orientation + cpos = pos1 - pos2 + + crotq = math.mul_quat(quat1, quat) # copy axis components + crot = wp.vec3(crotq[1], crotq[2], crotq[3]) * torquescale + + invweight_t = body_invweight0[worldid, body1id][0] + body_invweight0[worldid, body2id][0] + + pos_imp = wp.sqrt(wp.length_sq(cpos) + wp.length_sq(crot)) + + solref = eq_solref[worldid, i_eq] + solimp = eq_solimp[worldid, i_eq] + + for i in range(3): + _update_efc_row( + worldid, + timestep, + refsafe_in, + efcid + i, + cpos[i], + pos_imp, + invweight_t, + solref, + solimp, + 0.0, + Jqvelp[i], + 0.0, + ConstraintType.EQUALITY.value, + i_eq, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + invweight_r = body_invweight0[worldid, body1id][1] + body_invweight0[worldid, body2id][1] + + for i in range(3): + _update_efc_row( + worldid, + timestep, + refsafe_in, + efcid + 3 + i, + crot[i], + pos_imp, + invweight_r, + solref, + solimp, + 0.0, + Jqvelr[i], + 0.0, + ConstraintType.EQUALITY.value, + i_eq, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + +@wp.kernel +def _efc_limit_slide_hinge( + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + jnt_qposadr: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + jnt_solref: wp.array2d(dtype=wp.vec2), + jnt_solimp: wp.array2d(dtype=vec5), + jnt_range: wp.array2d(dtype=wp.vec2), + jnt_margin: wp.array2d(dtype=float), + jnt_limited_slide_hinge_adr: wp.array(dtype=int), + dof_invweight0: wp.array2d(dtype=float), + # Data in: + njmax_in: int, + qpos_in: wp.array2d(dtype=float), + qvel_in: wp.array2d(dtype=float), + # In: + refsafe_in: int, + # Data out: + nl_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), +): + worldid, jntlimitedid = wp.tid() + timestep = opt_timestep[worldid] + jntid = jnt_limited_slide_hinge_adr[jntlimitedid] + jntrange = jnt_range[worldid, jntid] + + qpos = qpos_in[worldid, jnt_qposadr[jntid]] + jntmargin = jnt_margin[worldid, jntid] + dist_min, dist_max = qpos - jntrange[0], jntrange[1] - qpos + pos = wp.min(dist_min, dist_max) - jntmargin + active = pos < 0 + + if active: + wp.atomic_add(nl_out, worldid, 1) + efcid = wp.atomic_add(nefc_out, worldid, 1) + + if efcid >= njmax_in: + return + + for i in range(nv): + efc_J_out[worldid, efcid, i] = 0.0 + + dofadr = jnt_dofadr[jntid] + + J = float(dist_min < dist_max) * 2.0 - 1.0 + efc_J_out[worldid, efcid, dofadr] = J + Jqvel = J * qvel_in[worldid, dofadr] + + _update_efc_row( + worldid, + timestep, + refsafe_in, + efcid, + pos, + pos, + dof_invweight0[worldid, dofadr], + jnt_solref[worldid, jntid], + jnt_solimp[worldid, jntid], + jntmargin, + Jqvel, + 0.0, + ConstraintType.LIMIT_JOINT.value, + jntid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + +@wp.kernel +def _efc_limit_ball( + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + jnt_qposadr: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + jnt_solref: wp.array2d(dtype=wp.vec2), + jnt_solimp: wp.array2d(dtype=vec5), + jnt_range: wp.array2d(dtype=wp.vec2), + jnt_margin: wp.array2d(dtype=float), + jnt_limited_ball_adr: wp.array(dtype=int), + dof_invweight0: wp.array2d(dtype=float), + # Data in: + njmax_in: int, + qpos_in: wp.array2d(dtype=float), + qvel_in: wp.array2d(dtype=float), + # In: + refsafe_in: int, + # Data out: + nl_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), +): + worldid, jntlimitedid = wp.tid() + timestep = opt_timestep[worldid] + jntid = jnt_limited_ball_adr[jntlimitedid] + qposadr = jnt_qposadr[jntid] + + qpos = qpos_in[worldid] + jnt_quat = wp.quat(qpos[qposadr + 0], qpos[qposadr + 1], qpos[qposadr + 2], qpos[qposadr + 3]) + jnt_quat = wp.normalize(jnt_quat) + axis_angle = math.quat_to_vel(jnt_quat) + jntrange = jnt_range[worldid, jntid] + axis, angle = math.normalize_with_norm(axis_angle) + jntmargin = jnt_margin[worldid, jntid] + + pos = wp.max(jntrange[0], jntrange[1]) - angle - jntmargin + active = pos < 0 + + if active: + wp.atomic_add(nl_out, worldid, 1) + efcid = wp.atomic_add(nefc_out, worldid, 1) + + if efcid >= njmax_in: + return + + for i in range(nv): + efc_J_out[worldid, efcid, i] = 0.0 + + dofadr = jnt_dofadr[jntid] + + efc_J_out[worldid, efcid, dofadr + 0] = -axis[0] + efc_J_out[worldid, efcid, dofadr + 1] = -axis[1] + efc_J_out[worldid, efcid, dofadr + 2] = -axis[2] + + Jqvel = -axis[0] * qvel_in[worldid, dofadr + 0] + Jqvel -= axis[1] * qvel_in[worldid, dofadr + 1] + Jqvel -= axis[2] * qvel_in[worldid, dofadr + 2] + + _update_efc_row( + worldid, + timestep, + refsafe_in, + efcid, + pos, + pos, + dof_invweight0[worldid, dofadr], + jnt_solref[worldid, jntid], + jnt_solimp[worldid, jntid], + jntmargin, + Jqvel, + 0.0, + ConstraintType.LIMIT_JOINT.value, + jntid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + +@wp.kernel +def _efc_limit_tendon( + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + jnt_dofadr: wp.array(dtype=int), + tendon_adr: wp.array(dtype=int), + tendon_num: wp.array(dtype=int), + tendon_limited_adr: wp.array(dtype=int), + tendon_solref_lim: wp.array2d(dtype=wp.vec2), + tendon_solimp_lim: wp.array2d(dtype=vec5), + tendon_range: wp.array2d(dtype=wp.vec2), + tendon_margin: wp.array2d(dtype=float), + tendon_invweight0: wp.array2d(dtype=float), + wrap_objid: wp.array(dtype=int), + wrap_type: wp.array(dtype=int), + # Data in: + njmax_in: int, + qvel_in: wp.array2d(dtype=float), + ten_length_in: wp.array2d(dtype=float), + ten_J_in: wp.array3d(dtype=float), + # In: + refsafe_in: int, + # Data out: + nl_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), +): + worldid, tenlimitedid = wp.tid() + timestep = opt_timestep[worldid] + tenid = tendon_limited_adr[tenlimitedid] + + tenrange = tendon_range[worldid, tenid] + length = ten_length_in[worldid, tenid] + dist_min, dist_max = length - tenrange[0], tenrange[1] - length + tenmargin = tendon_margin[worldid, tenid] + pos = wp.min(dist_min, dist_max) - tenmargin + active = pos < 0 + + if active: + wp.atomic_add(nl_out, worldid, 1) + efcid = wp.atomic_add(nefc_out, worldid, 1) + + if efcid >= njmax_in: + return + + for i in range(nv): + efc_J_out[worldid, efcid, i] = 0.0 + + Jqvel = float(0.0) + scl = float(dist_min < dist_max) * 2.0 - 1.0 + + adr = tendon_adr[tenid] + if wrap_type[adr] == wp.static(types.WrapType.JOINT.value): + ten_num = tendon_num[tenid] + for i in range(ten_num): + dofadr = jnt_dofadr[wrap_objid[adr + i]] + J = scl * ten_J_in[worldid, tenid, dofadr] + efc_J_out[worldid, efcid, dofadr] = J + Jqvel += J * qvel_in[worldid, dofadr] + else: + for i in range(nv): + J = scl * ten_J_in[worldid, tenid, i] + efc_J_out[worldid, efcid, i] = J + Jqvel += J * qvel_in[worldid, i] + + _update_efc_row( + worldid, + timestep, + refsafe_in, + efcid, + pos, + pos, + tendon_invweight0[worldid, tenid], + tendon_solref_lim[worldid, tenid], + tendon_solimp_lim[worldid, tenid], + tenmargin, + Jqvel, + 0.0, + ConstraintType.LIMIT_TENDON.value, + tenid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + +@wp.kernel +def _efc_contact_pyramidal( + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + opt_impratio: wp.array(dtype=float), + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + body_invweight0: wp.array2d(dtype=wp.vec2), + dof_bodyid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + # Data in: + njmax_in: int, + ncon_in: wp.array(dtype=int), + qvel_in: wp.array2d(dtype=float), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + # In: + refsafe_in: int, + dist_in: wp.array(dtype=float), + condim_in: wp.array(dtype=int), + includemargin_in: wp.array(dtype=float), + worldid_in: wp.array(dtype=int), + geom_in: wp.array(dtype=wp.vec2i), + pos_in: wp.array(dtype=wp.vec3), + frame_in: wp.array(dtype=wp.mat33), + friction_in: wp.array(dtype=vec5), + solref_in: wp.array(dtype=wp.vec2), + solimp_in: wp.array(dtype=vec5), + # Data out: + nefc_out: wp.array(dtype=int), + contact_efc_address_out: wp.array2d(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), +): + conid, dimid = wp.tid() + + if conid >= ncon_in[0]: + return + + condim = condim_in[conid] + + if condim == 1 and dimid > 0: + return + elif condim > 1 and dimid >= 2 * (condim - 1): + return + + includemargin = includemargin_in[conid] + pos = dist_in[conid] - includemargin + active = pos < 0 + + if active: + worldid = worldid_in[conid] + efcid = wp.atomic_add(nefc_out, worldid, 1) + + if efcid >= njmax_in: + return + + timestep = opt_timestep[worldid] + impratio = opt_impratio[worldid] + contact_efc_address_out[conid, dimid] = efcid + + geom = geom_in[conid] + body1 = geom_bodyid[geom[0]] + body2 = geom_bodyid[geom[1]] + + con_pos = pos_in[conid] + frame = frame_in[conid] + + # pyramidal has common invweight across all edges + invweight = body_invweight0[worldid, body1][0] + body_invweight0[worldid, body2][0] + + if condim > 1: + dimid2 = dimid / 2 + 1 + + friction = friction_in[conid] + fri0 = friction[0] + frii = friction[dimid2 - 1] + invweight = invweight + fri0 * fri0 * invweight + invweight = invweight * 2.0 * fri0 * fri0 / impratio + + Jqvel = float(0.0) + for i in range(nv): + J = float(0.0) + Ji = float(0.0) + jac1p, jac1r = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + con_pos, + body1, + i, + worldid, + ) + jac2p, jac2r = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + con_pos, + body2, + i, + worldid, + ) + jacp_dif = jac2p - jac1p + for xyz in range(3): + J += frame[0, xyz] * jacp_dif[xyz] + + if condim > 1: + if dimid2 < 3: + Ji += frame[dimid2, xyz] * jacp_dif[xyz] + else: + Ji += frame[dimid2 - 3, xyz] * (jac2r[xyz] - jac1r[xyz]) + + if condim > 1: + if dimid % 2 == 0: + J += Ji * frii + else: + J -= Ji * frii + + efc_J_out[worldid, efcid, i] = J + Jqvel += J * qvel_in[worldid, i] + + if condim == 1: + efc_type = int(ConstraintType.CONTACT_FRICTIONLESS.value) + else: + efc_type = int(ConstraintType.CONTACT_PYRAMIDAL.value) + + _update_efc_row( + worldid, + timestep, + refsafe_in, + efcid, + pos, + pos, + invweight, + solref_in[conid], + solimp_in[conid], + includemargin, + Jqvel, + 0.0, + efc_type, + conid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + +@wp.kernel +def _efc_contact_elliptic( + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + opt_impratio: wp.array(dtype=float), + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + body_invweight0: wp.array2d(dtype=wp.vec2), + dof_bodyid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + # Data in: + njmax_in: int, + ncon_in: wp.array(dtype=int), + qvel_in: wp.array2d(dtype=float), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + # In: + refsafe_in: int, + dist_in: wp.array(dtype=float), + condim_in: wp.array(dtype=int), + includemargin_in: wp.array(dtype=float), + worldid_in: wp.array(dtype=int), + geom_in: wp.array(dtype=wp.vec2i), + pos_in: wp.array(dtype=wp.vec3), + frame_in: wp.array(dtype=wp.mat33), + friction_in: wp.array(dtype=vec5), + solref_in: wp.array(dtype=wp.vec2), + solreffriction_in: wp.array(dtype=wp.vec2), + solimp_in: wp.array(dtype=vec5), + # Data out: + nefc_out: wp.array(dtype=int), + contact_efc_address_out: wp.array2d(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), +): + conid, dimid = wp.tid() + + if conid >= ncon_in[0]: + return + + condim = condim_in[conid] + + if dimid > condim - 1: + return + + includemargin = includemargin_in[conid] + pos = dist_in[conid] - includemargin + active = pos < 0.0 + + if active: + worldid = worldid_in[conid] + efcid = wp.atomic_add(nefc_out, worldid, 1) + + if efcid >= njmax_in: + return + + timestep = opt_timestep[worldid] + impratio = opt_impratio[worldid] + contact_efc_address_out[conid, dimid] = efcid + + geom = geom_in[conid] + body1 = geom_bodyid[geom[0]] + body2 = geom_bodyid[geom[1]] + + cpos = pos_in[conid] + frame = frame_in[conid] + + # TODO(team): parallelize J and Jqvel computation? + Jqvel = float(0.0) + for i in range(nv): + J = float(0.0) + jac1p, jac1r = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + cpos, + body1, + i, + worldid, + ) + jac2p, jac2r = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + cpos, + body2, + i, + worldid, + ) + for xyz in range(3): + if dimid < 3: + jac_dif = jac2p[xyz] - jac1p[xyz] + J += frame[dimid, xyz] * jac_dif + else: + jac_dif = jac2r[xyz] - jac1r[xyz] + J += frame[dimid - 3, xyz] * jac_dif + + efc_J_out[worldid, efcid, i] = J + Jqvel += J * qvel_in[worldid, i] + + invweight = body_invweight0[worldid, body1][0] + body_invweight0[worldid, body2][0] + + ref = solref_in[conid] + pos_aref = pos + + if dimid > 0: + solreffriction = solreffriction_in[conid] + + # non-normal directions use solreffriction (if non-zero) + if solreffriction[0] or solreffriction[1]: + ref = solreffriction + + # TODO(team): precompute 1 / impratio + invweight = invweight / impratio + friction = friction_in[conid] + + if dimid > 1: + fri0 = friction[0] + frii = friction[dimid - 1] + fri = fri0 * fri0 / (frii * frii) + invweight *= fri + + pos_aref = 0.0 + + if condim == 1: + efc_type = int(ConstraintType.CONTACT_FRICTIONLESS.value) + else: + efc_type = int(ConstraintType.CONTACT_ELLIPTIC.value) + + _update_efc_row( + worldid, + timestep, + refsafe_in, + efcid, + pos_aref, + pos, + invweight, + ref, + solimp_in[conid], + includemargin, + Jqvel, + 0.0, + efc_type, + conid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + +@wp.kernel +def _num_equality( + # Data in: + ne_connect_in: wp.array(dtype=int), + ne_weld_in: wp.array(dtype=int), + ne_jnt_in: wp.array(dtype=int), + ne_ten_in: wp.array(dtype=int), + # Data out: + ne_out: wp.array(dtype=int), +): + worldid = wp.tid() + ne = ne_connect_in[worldid] + ne_weld_in[worldid] + ne_jnt_in[worldid] + ne_ten_in[worldid] + ne_out[worldid] = ne + + +@event_scope +def make_constraint(m: types.Model, d: types.Data): + """Creates constraint jacobians and other supporting data.""" + + wp.launch( + zero_constraint_counts, + dim=d.nworld, + inputs=[ + d.ne, + d.ne_connect, + d.ne_weld, + d.ne_jnt, + d.ne_ten, + d.nf, + d.nl, + d.nefc, + ], + ) + + if not (m.opt.disableflags & types.DisableBit.CONSTRAINT.value): + refsafe = m.opt.disableflags & types.DisableBit.REFSAFE + + if not (m.opt.disableflags & types.DisableBit.EQUALITY.value): + wp.launch( + _efc_equality_connect, + dim=(d.nworld, m.eq_connect_adr.size), + inputs=[ + m.nv, + m.nsite, + m.opt.timestep, + m.body_parentid, + m.body_rootid, + m.body_invweight0, + m.dof_bodyid, + m.site_bodyid, + m.eq_obj1id, + m.eq_obj2id, + m.eq_objtype, + m.eq_solref, + m.eq_solimp, + m.eq_data, + m.eq_connect_adr, + d.njmax, + d.qvel, + d.eq_active, + d.xpos, + d.xmat, + d.site_xpos, + d.subtree_com, + d.cdof, + refsafe, + ], + outputs=[ + d.ne_connect, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + ], + ) + wp.launch( + _efc_equality_weld, + dim=(d.nworld, m.eq_wld_adr.size), + inputs=[ + m.nv, + m.nsite, + m.opt.timestep, + m.body_parentid, + m.body_rootid, + m.body_invweight0, + m.dof_bodyid, + m.site_bodyid, + m.site_quat, + m.eq_obj1id, + m.eq_obj2id, + m.eq_objtype, + m.eq_solref, + m.eq_solimp, + m.eq_data, + m.eq_wld_adr, + d.njmax, + d.qvel, + d.eq_active, + d.xpos, + d.xquat, + d.xmat, + d.site_xpos, + d.subtree_com, + d.cdof, + refsafe, + ], + outputs=[ + d.ne_weld, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + ], + ) + wp.launch( + _efc_equality_joint, + dim=(d.nworld, m.eq_jnt_adr.size), + inputs=[ + m.nv, + m.opt.timestep, + m.qpos0, + m.jnt_qposadr, + m.jnt_dofadr, + m.dof_invweight0, + m.eq_obj1id, + m.eq_obj2id, + m.eq_solref, + m.eq_solimp, + m.eq_data, + m.eq_jnt_adr, + d.njmax, + d.qpos, + d.qvel, + d.eq_active, + refsafe, + ], + outputs=[ + d.ne_jnt, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + ], + ) + wp.launch( + _efc_equality_tendon, + dim=(d.nworld, m.eq_ten_adr.size), + inputs=[ + m.nv, + m.opt.timestep, + m.eq_obj1id, + m.eq_obj2id, + m.eq_solref, + m.eq_solimp, + m.eq_data, + m.eq_ten_adr, + m.tendon_length0, + m.tendon_invweight0, + d.njmax, + d.qvel, + d.eq_active, + d.ten_length, + d.ten_J, + refsafe, + ], + outputs=[ + d.ne_ten, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + ], + ) + + wp.launch( + _num_equality, + dim=(d.nworld,), + inputs=[ + d.ne_connect, + d.ne_weld, + d.ne_jnt, + d.ne_ten, + ], + outputs=[ + d.ne, + ], + ) + + if not (m.opt.disableflags & types.DisableBit.FRICTIONLOSS.value): + wp.launch( + _efc_friction_dof, + dim=(d.nworld, m.nv), + inputs=[ + m.nv, + m.opt.timestep, + m.dof_invweight0, + m.dof_frictionloss, + m.dof_solimp, + m.dof_solref, + d.njmax, + d.qvel, + refsafe, + ], + outputs=[ + d.nf, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + ], + ) + + wp.launch( + _efc_friction_tendon, + dim=(d.nworld, m.ntendon), + inputs=[ + m.nv, + m.opt.timestep, + m.tendon_solref_fri, + m.tendon_solimp_fri, + m.tendon_frictionloss, + m.tendon_invweight0, + d.qvel, + d.ten_J, + refsafe, + ], + outputs=[ + d.nf, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + ], + ) + + # limit + if not (m.opt.disableflags & types.DisableBit.LIMIT.value): + limit_ball = m.jnt_limited_ball_adr.size > 0 + if limit_ball: + wp.launch( + _efc_limit_ball, + dim=(d.nworld, m.jnt_limited_ball_adr.size), + inputs=[ + m.nv, + m.opt.timestep, + m.jnt_qposadr, + m.jnt_dofadr, + m.jnt_solref, + m.jnt_solimp, + m.jnt_range, + m.jnt_margin, + m.jnt_limited_ball_adr, + m.dof_invweight0, + d.njmax, + d.qpos, + d.qvel, + refsafe, + ], + outputs=[ + d.nl, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + ], + ) + + limit_slide_hinge = m.jnt_limited_slide_hinge_adr.size > 0 + if limit_slide_hinge: + wp.launch( + _efc_limit_slide_hinge, + dim=(d.nworld, m.jnt_limited_slide_hinge_adr.size), + inputs=[ + m.nv, + m.opt.timestep, + m.jnt_qposadr, + m.jnt_dofadr, + m.jnt_solref, + m.jnt_solimp, + m.jnt_range, + m.jnt_margin, + m.jnt_limited_slide_hinge_adr, + m.dof_invweight0, + d.njmax, + d.qpos, + d.qvel, + refsafe, + ], + outputs=[ + d.nl, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + ], + ) + + limit_tendon = m.tendon_limited_adr.size > 0 + if limit_tendon: + wp.launch( + _efc_limit_tendon, + dim=(d.nworld, m.tendon_limited_adr.size), + inputs=[ + m.nv, + m.opt.timestep, + m.jnt_dofadr, + m.tendon_adr, + m.tendon_num, + m.tendon_limited_adr, + m.tendon_solref_lim, + m.tendon_solimp_lim, + m.tendon_range, + m.tendon_margin, + m.tendon_invweight0, + m.wrap_objid, + m.wrap_type, + d.njmax, + d.qvel, + d.ten_length, + d.ten_J, + refsafe, + ], + outputs=[ + d.nl, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + ], + ) + + # contact + if not (m.opt.disableflags & types.DisableBit.CONTACT.value): + if m.opt.cone == types.ConeType.PYRAMIDAL.value: + wp.launch( + _efc_contact_pyramidal, + dim=(d.nconmax, 2 * (m.condim_max - 1) if m.condim_max > 1 else 1), + inputs=[ + m.nv, + m.opt.timestep, + m.opt.impratio, + m.body_parentid, + m.body_rootid, + m.body_invweight0, + m.dof_bodyid, + m.geom_bodyid, + d.njmax, + d.ncon, + d.qvel, + d.subtree_com, + d.cdof, + refsafe, + d.contact.dist, + d.contact.dim, + d.contact.includemargin, + d.contact.worldid, + d.contact.geom, + d.contact.pos, + d.contact.frame, + d.contact.friction, + d.contact.solref, + d.contact.solimp, + ], + outputs=[ + d.nefc, + d.contact.efc_address, + d.efc.type, + d.efc.id, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + ], + ) + elif m.opt.cone == types.ConeType.ELLIPTIC.value: + wp.launch( + _efc_contact_elliptic, + dim=(d.nconmax, m.condim_max), + inputs=[ + m.nv, + m.opt.timestep, + m.opt.impratio, + m.body_parentid, + m.body_rootid, + m.body_invweight0, + m.dof_bodyid, + m.geom_bodyid, + d.njmax, + d.ncon, + d.qvel, + d.subtree_com, + d.cdof, + refsafe, + d.contact.dist, + d.contact.dim, + d.contact.includemargin, + d.contact.worldid, + d.contact.geom, + d.contact.pos, + d.contact.frame, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + ], + outputs=[ + d.nefc, + d.contact.efc_address, + d.efc.type, + d.efc.id, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + ], + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint_test.py new file mode 100644 index 00000000..28cffc8c --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint_test.py @@ -0,0 +1,217 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Tests for constraint functions.""" + +import mujoco +import numpy as np +import warp as wp +from absl.testing import absltest +from absl.testing import parameterized + +import mujoco_warp as mjwarp + +from mujoco.mjx.third_party.mujoco_warp._src import test_util +from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType + +# tolerance for difference between MuJoCo and MJWarp constraint calculations, +# mostly due to float precision +_TOLERANCE = 5e-5 + + +def _assert_eq(a, b, name): + tol = _TOLERANCE * 10 # avoid test noise + err_msg = f"mismatch: {name}" + np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol) + + +class ConstraintTest(parameterized.TestCase): + @parameterized.parameters( + (ConeType.PYRAMIDAL, 1, 1), + (ConeType.PYRAMIDAL, 1, 3), + (ConeType.PYRAMIDAL, 1, 4), + (ConeType.PYRAMIDAL, 1, 6), + (ConeType.PYRAMIDAL, 3, 3), + (ConeType.PYRAMIDAL, 3, 4), + (ConeType.PYRAMIDAL, 3, 6), + (ConeType.PYRAMIDAL, 4, 4), + (ConeType.PYRAMIDAL, 4, 6), + (ConeType.PYRAMIDAL, 6, 6), + (ConeType.ELLIPTIC, 1, 1), + (ConeType.ELLIPTIC, 1, 3), + (ConeType.ELLIPTIC, 1, 4), + (ConeType.ELLIPTIC, 1, 6), + (ConeType.ELLIPTIC, 3, 3), + (ConeType.ELLIPTIC, 3, 4), + (ConeType.ELLIPTIC, 3, 6), + (ConeType.ELLIPTIC, 4, 4), + (ConeType.ELLIPTIC, 4, 6), + (ConeType.ELLIPTIC, 6, 6), + ) + def test_condim(self, cone, condim1, condim2): + """Test condim.""" + xml = f""" + + + + + + + + + + + + + """ + + _, mjd, m, d = test_util.fixture(xml=xml, cone=cone) + + for arr in ( + d.efc.D, + d.efc.aref, + d.efc.pos, + d.efc.margin, + ): + arr.zero_() + + # fill with nan to check whether we are not reading uninitialized values + d.efc.J.fill_(wp.nan) + + mjwarp.make_constraint(m, d) + + _assert_eq(d.efc.J.numpy()[0, : mjd.nefc, :].reshape(-1), mjd.efc_J, "efc_J") + _assert_eq(d.efc.D.numpy()[0, : mjd.nefc], mjd.efc_D, "efc_D") + _assert_eq(d.efc.aref.numpy()[0, : mjd.nefc], mjd.efc_aref, "efc_aref") + _assert_eq(d.efc.pos.numpy()[0, : mjd.nefc], mjd.efc_pos, "efc_pos") + _assert_eq(d.efc.margin.numpy()[0, : mjd.nefc], mjd.efc_margin, "efc_margin") + + @parameterized.parameters( + mujoco.mjtCone.mjCONE_PYRAMIDAL, + mujoco.mjtCone.mjCONE_ELLIPTIC, + ) + def test_constraints(self, cone): + """Test constraints.""" + for key in range(3): + mjm, mjd, m, d = test_util.fixture("constraints.xml", sparse=False, cone=cone, keyframe=key) + + for arr in ( + d.efc.D, + d.efc.aref, + d.efc.pos, + d.efc.margin, + d.ne, + d.nefc, + d.nf, + d.nl, + ): + arr.zero_() + + d.efc.J.fill_(wp.nan) + + mjwarp.make_constraint(m, d) + + _assert_eq(d.ne.numpy()[0], mjd.ne, "ne") + _assert_eq(d.nefc.numpy()[0], mjd.nefc, "nefc") + _assert_eq(d.nf.numpy()[0], mjd.nf, "nf") + _assert_eq(d.nl.numpy()[0], mjd.nl, "nl") + _assert_eq(d.efc.J.numpy()[0, : mjd.nefc, :].reshape(-1), mjd.efc_J, "efc_J") + _assert_eq(d.efc.D.numpy()[0, : mjd.nefc], mjd.efc_D, "efc_D") + _assert_eq(d.efc.vel.numpy()[0, : mjd.nefc], mjd.efc_vel, "efc_vel") + _assert_eq(d.efc.aref.numpy()[0, : mjd.nefc], mjd.efc_aref, "efc_aref") + _assert_eq(d.efc.pos.numpy()[0, : mjd.nefc], mjd.efc_pos, "efc_pos") + _assert_eq(d.efc.margin.numpy()[0, : mjd.nefc], mjd.efc_margin, "efc_margin") + _assert_eq(d.efc.type.numpy()[0, : mjd.nefc], mjd.efc_type, "efc_type") + + def test_limit_tendon(self): + """Test limit tendon constraints.""" + for keyframe in range(-1, 1): + _, mjd, m, d = test_util.fixture("tendon/tendon_limit.xml", sparse=False, keyframe=keyframe) + + for arr in (d.nefc, d.nl, d.efc.J, d.efc.D, d.efc.aref, d.efc.pos, d.efc.margin): + arr.zero_() + + mjwarp.make_constraint(m, d) + + _assert_eq(d.nefc.numpy()[0], mjd.nefc, "nefc") + _assert_eq(d.nl.numpy()[0], mjd.nl, "nl") + _assert_eq(d.efc.J.numpy()[0, : mjd.nefc, :].reshape(-1), mjd.efc_J, "efc_J") + _assert_eq(d.efc.D.numpy()[0, : mjd.nefc], mjd.efc_D, "efc_D") + _assert_eq(d.efc.vel.numpy()[0, : mjd.nefc], mjd.efc_vel, "efc_vel") + _assert_eq(d.efc.aref.numpy()[0, : mjd.nefc], mjd.efc_aref, "efc_aref") + _assert_eq(d.efc.pos.numpy()[0, : mjd.nefc], mjd.efc_pos, "efc_pos") + _assert_eq(d.efc.margin.numpy()[0, : mjd.nefc], mjd.efc_margin, "efc_margin") + _assert_eq(d.efc.type.numpy()[0, : mjd.nefc], mjd.efc_type, "efc_type") + + def test_equality_tendon(self): + """Test equality tendon constraints.""" + + _, mjd, m, d = test_util.fixture( + xml=""" + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + """ + ) + + mjwarp.make_constraint(m, d) + + _assert_eq(d.nefc.numpy()[0], mjd.nefc, "nefc") + _assert_eq(d.ne.numpy()[0], mjd.ne, "ne") + _assert_eq(d.efc.J.numpy()[0, : mjd.nefc, :].reshape(-1), mjd.efc_J, "efc_J") + _assert_eq(d.efc.D.numpy()[0, : mjd.nefc], mjd.efc_D, "efc_D") + _assert_eq(d.efc.vel.numpy()[0, : mjd.nefc], mjd.efc_vel, "efc_vel") + _assert_eq(d.efc.aref.numpy()[0, : mjd.nefc], mjd.efc_aref, "efc_aref") + _assert_eq(d.efc.pos.numpy()[0, : mjd.nefc], mjd.efc_pos, "efc_pos") + _assert_eq(d.efc.margin.numpy()[0, : mjd.nefc], mjd.efc_margin, "efc_margin") + _assert_eq(d.efc.type.numpy()[0, : mjd.nefc], mjd.efc_type, "efc_type") + + +if __name__ == "__main__": + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py new file mode 100644 index 00000000..dcd70f4b --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py @@ -0,0 +1,206 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src.support import mul_m +from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import DynType +from mujoco.mjx.third_party.mujoco_warp._src.types import GainType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import vec10f +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope + +wp.set_module_options({"enable_backward": False}) + + +# TODO(team): improve performance with tile operations? +@wp.kernel +def _qderiv_actuator_passive( + # Model: + nu: int, + opt_timestep: wp.array(dtype=float), + opt_is_sparse: bool, + dof_damping: wp.array2d(dtype=float), + actuator_dyntype: wp.array(dtype=int), + actuator_gaintype: wp.array(dtype=int), + actuator_biastype: wp.array(dtype=int), + actuator_actadr: wp.array(dtype=int), + actuator_actnum: wp.array(dtype=int), + actuator_gainprm: wp.array2d(dtype=vec10f), + actuator_biasprm: wp.array2d(dtype=vec10f), + # Data in: + act_in: wp.array2d(dtype=float), + ctrl_in: wp.array2d(dtype=float), + actuator_moment_in: wp.array3d(dtype=float), + qM_in: wp.array3d(dtype=float), + # In: + qMi: wp.array(dtype=int), + qMj: wp.array(dtype=int), + actuation_enabled: bool, + passive_enabled: bool, + # Data out: + qM_integration_out: wp.array3d(dtype=float), +): + worldid, elemid = wp.tid() + dofiid = qMi[elemid] + dofjid = qMj[elemid] + + qderiv = float(0.0) + for actid in range(nu): + if actuation_enabled: + if actuator_gaintype[actid] == int(GainType.AFFINE.value): + gain = actuator_gainprm[worldid, actid][2] + else: + gain = 0.0 + + if actuator_biastype[actid] == int(BiasType.AFFINE.value): + bias = actuator_biasprm[worldid, actid][2] + else: + bias = 0.0 + + if actuator_dyntype[actid] != int(DynType.NONE.value): + act_first = actuator_actadr[actid] + act_last = act_first + actuator_actnum[actid] - 1 + vel = bias + gain * act_in[worldid, act_last] + else: + vel = bias + gain * ctrl_in[worldid, actid] + + qderiv += actuator_moment_in[worldid, actid, dofiid] * actuator_moment_in[worldid, actid, dofjid] * vel + + if passive_enabled and dofiid == dofjid: + qderiv -= dof_damping[worldid, dofiid] / float(nu) + + qderiv *= opt_timestep[worldid] + + if opt_is_sparse: + qM_integration_out[worldid, 0, elemid] = qM_in[worldid, 0, elemid] - qderiv + else: + qM = qM_in[worldid, dofiid, dofjid] - qderiv + qM_integration_out[worldid, dofiid, dofjid] = qM + if dofiid != dofjid: + qM_integration_out[worldid, dofjid, dofiid] = qM + + +# TODO(team): improve performance with tile operations? +@wp.kernel +def _qderiv_tendon_damping( + # Model: + ntendon: int, + opt_timestep: wp.array(dtype=float), + opt_is_sparse: bool, + tendon_damping: wp.array2d(dtype=float), + # Data in: + ten_J_in: wp.array3d(dtype=float), + # In: + qMi: wp.array(dtype=int), + qMj: wp.array(dtype=int), + # Data out: + qM_integration_out: wp.array3d(dtype=float), +): + worldid, elemid = wp.tid() + dofiid = qMi[elemid] + dofjid = qMj[elemid] + + qderiv = float(0.0) + for tenid in range(ntendon): + qderiv -= ten_J_in[worldid, tenid, dofiid] * ten_J_in[worldid, tenid, dofjid] * tendon_damping[worldid, tenid] + qderiv *= opt_timestep[worldid] + + if opt_is_sparse: + qM_integration_out[worldid, 0, elemid] -= qderiv + else: + qM_integration_out[worldid, dofiid, dofjid] -= qderiv + if dofiid != dofjid: + qM_integration_out[worldid, dofjid, dofiid] -= qderiv + + +@wp.kernel +def _qfrc_forward( + # Data in: + qfrc_smooth_in: wp.array2d(dtype=float), + qfrc_constraint_in: wp.array2d(dtype=float), + # Data out: + qfrc_integration_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + qfrc_integration_out[worldid, dofid] = qfrc_smooth_in[worldid, dofid] + qfrc_constraint_in[worldid, dofid] + + +@event_scope +def deriv_smooth_vel(m: Model, d: Data, flg_forward: bool = True): + """Analytical derivative of smooth forces w.r.t. velocities. + + Args: + m (Model): The model containing kinematic and dynamic information (device). + d (Data): The data object containing the current state and output arrays (device). + flg_forward (bool, optional): If True forward dynamics else inverse dynamics routine. + Default is True. + """ + actuation_enabled = not (m.opt.disableflags & DisableBit.ACTUATION) + passive_enabled = not (m.opt.disableflags & DisableBit.PASSIVE) + + qMi = m.qM_fullm_i if m.opt.is_sparse else m.dof_tri_row + qMj = m.qM_fullm_j if m.opt.is_sparse else m.dof_tri_col + if actuation_enabled or passive_enabled: + wp.launch( + _qderiv_actuator_passive, + dim=(d.nworld, qMi.size), + inputs=[ + m.nu, + m.opt.timestep, + m.opt.is_sparse, + m.dof_damping, + m.actuator_dyntype, + m.actuator_gaintype, + m.actuator_biastype, + m.actuator_actadr, + m.actuator_actnum, + m.actuator_gainprm, + m.actuator_biasprm, + d.act, + d.ctrl, + d.actuator_moment, + d.qM, + qMi, + qMj, + actuation_enabled, + passive_enabled, + ], + outputs=[d.qM_integration], + ) + + if passive_enabled: + wp.launch( + _qderiv_tendon_damping, + dim=(d.nworld, qMi.size), + inputs=[m.ntendon, m.opt.timestep, m.opt.is_sparse, m.tendon_damping, d.ten_J, qMi, qMj], + outputs=[d.qM_integration], + ) + + if flg_forward: + wp.launch( + _qfrc_forward, + dim=(d.nworld, m.nv), + inputs=[d.qfrc_smooth, d.qfrc_constraint], + outputs=[d.qfrc_integration], + ) + else: + # qfrc = qM @ qacc + mul_m(m, d, d.qfrc_integration, d.qacc, d.inverse_mul_m_skip, d.qM_integration) + + # TODO(team): rne derivative diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py new file mode 100644 index 00000000..a27c38a1 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py @@ -0,0 +1,1037 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +from typing import Optional + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src import collision_driver +from mujoco.mjx.third_party.mujoco_warp._src import constraint +from mujoco.mjx.third_party.mujoco_warp._src import derivative +from mujoco.mjx.third_party.mujoco_warp._src import math +from mujoco.mjx.third_party.mujoco_warp._src import passive +from mujoco.mjx.third_party.mujoco_warp._src import sensor +from mujoco.mjx.third_party.mujoco_warp._src import smooth +from mujoco.mjx.third_party.mujoco_warp._src import solver +from mujoco.mjx.third_party.mujoco_warp._src import util_misc +from mujoco.mjx.third_party.mujoco_warp._src.support import xfrc_accumulate +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import DynType +from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import GainType +from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType +from mujoco.mjx.third_party.mujoco_warp._src.types import JointType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import TileSet +from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType +from mujoco.mjx.third_party.mujoco_warp._src.types import vec10f +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel + +wp.set_module_options({"enable_backward": False}) + +# RK4 tableau +_RK4_A = [ + [0.5, 0.0, 0.0], + [0.0, 0.5, 0.0], + [0.0, 0.0, 1.0], +] +_RK4_B = [1.0 / 6.0, 1.0 / 3.0, 1.0 / 3.0, 1.0 / 6.0] + + +@wp.kernel +def _next_position( + # Model: + opt_timestep: wp.array(dtype=float), + jnt_type: wp.array(dtype=int), + jnt_qposadr: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + # Data in: + qpos_in: wp.array2d(dtype=float), + qvel_in: wp.array2d(dtype=float), + # In: + qvel_scale_in: float, + # Data out: + qpos_out: wp.array2d(dtype=float), +): + worldid, jntid = wp.tid() + timestep = opt_timestep[worldid] + + jnttype = jnt_type[jntid] + qpos_adr = jnt_qposadr[jntid] + dof_adr = jnt_dofadr[jntid] + qpos = qpos_in[worldid] + qpos_next = qpos_out[worldid] + qvel = qvel_in[worldid] + + if jnttype == wp.static(JointType.FREE.value): + qpos_pos = wp.vec3(qpos[qpos_adr], qpos[qpos_adr + 1], qpos[qpos_adr + 2]) + qvel_lin = wp.vec3(qvel[dof_adr], qvel[dof_adr + 1], qvel[dof_adr + 2]) * qvel_scale_in + + qpos_new = qpos_pos + timestep * qvel_lin + + qpos_quat = wp.quat( + qpos[qpos_adr + 3], + qpos[qpos_adr + 4], + qpos[qpos_adr + 5], + qpos[qpos_adr + 6], + ) + qvel_ang = wp.vec3(qvel[dof_adr + 3], qvel[dof_adr + 4], qvel[dof_adr + 5]) * qvel_scale_in + + qpos_quat_new = math.quat_integrate(qpos_quat, qvel_ang, timestep) + + qpos_next[qpos_adr + 0] = qpos_new[0] + qpos_next[qpos_adr + 1] = qpos_new[1] + qpos_next[qpos_adr + 2] = qpos_new[2] + qpos_next[qpos_adr + 3] = qpos_quat_new[0] + qpos_next[qpos_adr + 4] = qpos_quat_new[1] + qpos_next[qpos_adr + 5] = qpos_quat_new[2] + qpos_next[qpos_adr + 6] = qpos_quat_new[3] + + elif jnttype == wp.static(JointType.BALL.value): + qpos_quat = wp.quat( + qpos[qpos_adr + 0], + qpos[qpos_adr + 1], + qpos[qpos_adr + 2], + qpos[qpos_adr + 3], + ) + qvel_ang = wp.vec3(qvel[dof_adr], qvel[dof_adr + 1], qvel[dof_adr + 2]) * qvel_scale_in + + qpos_quat_new = math.quat_integrate(qpos_quat, qvel_ang, timestep) + + qpos_next[qpos_adr + 0] = qpos_quat_new[0] + qpos_next[qpos_adr + 1] = qpos_quat_new[1] + qpos_next[qpos_adr + 2] = qpos_quat_new[2] + qpos_next[qpos_adr + 3] = qpos_quat_new[3] + + else: # if jnt_type in (JointType.HINGE, JointType.SLIDE): + qpos_next[qpos_adr] = qpos[qpos_adr] + timestep * qvel[dof_adr] * qvel_scale_in + + +@wp.kernel +def _next_velocity( + # Model: + opt_timestep: wp.array(dtype=float), + # Data in: + qvel_in: wp.array2d(dtype=float), + qacc_in: wp.array2d(dtype=float), + # In: + qacc_scale_in: float, + # Data out: + qvel_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + timestep = opt_timestep[worldid] + qvel_out[worldid, dofid] = qvel_in[worldid, dofid] + qacc_scale_in * qacc_in[worldid, dofid] * timestep + + +# TODO(team): kernel analyzer array slice? +@wp.func +def _next_act( + # Model: + opt_timestep: float, # kernel_analyzer: ignore + actuator_dyntype: int, # kernel_analyzer: ignore + actuator_dynprm: vec10f, # kernel_analyzer: ignore + actuator_actrange: wp.vec2, # kernel_analyzer: ignore + # Data In: + act_in: float, # kernel_analyzer: ignore + act_dot_in: float, # kernel_analyzer: ignore + # In: + act_dot_scale: float, + clamp: bool, +) -> float: + # advance actuation + if actuator_dyntype == wp.static(DynType.FILTEREXACT.value): + tau = wp.max(MJ_MINVAL, actuator_dynprm[0]) + act = act_in + act_dot_scale * act_dot_in * tau * (1.0 - wp.exp(-opt_timestep / tau)) + else: + act = act_in + act_dot_scale * act_dot_in * opt_timestep + + # clamp to actrange + if clamp: + act = wp.clamp(act, actuator_actrange[0], actuator_actrange[1]) + + return act + + +@wp.kernel +def _next_activation( + # Model: + opt_timestep: wp.array(dtype=float), + actuator_dyntype: wp.array(dtype=int), + actuator_actlimited: wp.array(dtype=bool), + actuator_dynprm: wp.array2d(dtype=vec10f), + actuator_actrange: wp.array2d(dtype=wp.vec2), + # Data in: + act_in: wp.array2d(dtype=float), + act_dot_in: wp.array2d(dtype=float), + # In: + act_dot_scale: float, + limit: bool, + # Data out: + act_out: wp.array2d(dtype=float), +): + worldid, actid = wp.tid() + act = _next_act( + opt_timestep[worldid], + actuator_dyntype[actid], + actuator_dynprm[worldid, actid], + actuator_actrange[worldid, actid], + act_in[worldid, actid], + act_dot_in[worldid, actid], + act_dot_scale, + limit and actuator_actlimited[actid], + ) + act_out[worldid, actid] = act + + +@wp.kernel +def _next_time( + # Model: + opt_timestep: wp.array(dtype=float), + # Data in: + nconmax_in: int, + njmax_in: int, + ncon_in: wp.array(dtype=int), + nefc_in: wp.array(dtype=int), + time_in: wp.array(dtype=float), + ncollision_in: wp.array(dtype=int), + # Data out: + time_out: wp.array(dtype=float), +): + worldid = wp.tid() + time_out[worldid] = time_in[worldid] + opt_timestep[worldid] + nefc = nefc_in[worldid] + + if nefc > njmax_in: + wp.printf("nefc overflow - please increase njmax to %u\n", nefc) + + if worldid == 0: + ncollision = ncollision_in[0] + if ncollision > nconmax_in: + wp.printf("ncollision overflow - please increase nconmax to %u\n", ncollision) + + if ncon_in[0] > nconmax_in: + wp.printf("ncon overflow - please increase nconmax to %u\n", ncon_in[0]) + + +def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None): + """Advance state and time given activation derivatives and acceleration.""" + + # TODO(team): can we assume static timesteps? + + # advance activations + if m.na: + wp.launch( + _next_activation, + dim=(d.nworld, m.na), + inputs=[ + m.opt.timestep, + m.actuator_dyntype, + m.actuator_actlimited, + m.actuator_dynprm, + m.actuator_actrange, + d.act, + d.act_dot, + 1.0, + True, + ], + outputs=[ + d.act, + ], + ) + + wp.launch( + _next_velocity, + dim=(d.nworld, m.nv), + inputs=[ + m.opt.timestep, + d.qvel, + qacc, + 1.0, + ], + outputs=[ + d.qvel, + ], + ) + + # advance positions with qvel if given, d.qvel otherwise (semi-implicit) + if qvel is not None: + qvel_in = qvel + else: + qvel_in = d.qvel + + wp.launch( + _next_position, + dim=(d.nworld, m.njnt), + inputs=[ + m.opt.timestep, + m.jnt_type, + m.jnt_qposadr, + m.jnt_dofadr, + d.qpos, + qvel_in, + 1.0, + ], + outputs=[ + d.qpos, + ], + ) + + wp.launch( + _next_time, + dim=(d.nworld,), + inputs=[ + m.opt.timestep, + d.nconmax, + d.njmax, + d.ncon, + d.nefc, + d.time, + d.ncollision, + ], + outputs=[ + d.time, + ], + ) + + +@wp.kernel +def _euler_damp_qfrc_sparse( + # Model: + opt_timestep: wp.array(dtype=float), + dof_Madr: wp.array(dtype=int), + dof_damping: wp.array2d(dtype=float), + # Data in: + qfrc_smooth_in: wp.array2d(dtype=float), + qfrc_constraint_in: wp.array2d(dtype=float), + # Data out: + qfrc_integration_out: wp.array2d(dtype=float), + qM_integration_out: wp.array3d(dtype=float), +): + worldid, tid = wp.tid() + timestep = opt_timestep[worldid] + + adr = dof_Madr[tid] + qM_integration_out[worldid, 0, adr] += timestep * dof_damping[worldid, tid] + qfrc_integration_out[worldid, tid] = qfrc_smooth_in[worldid, tid] + qfrc_constraint_in[worldid, tid] + + +def _euler_sparse(m: Model, d: Data): + wp.copy(d.qM_integration, d.qM) + wp.launch( + _euler_damp_qfrc_sparse, + dim=(d.nworld, m.nv), + inputs=[ + m.opt.timestep, + m.dof_Madr, + m.dof_damping, + d.qfrc_smooth, + d.qfrc_constraint, + ], + outputs=[ + d.qfrc_integration, + d.qM_integration, + ], + ) + smooth.factor_solve_i( + m, + d, + d.qM_integration, + d.qLD_integration, + d.qLDiagInv_integration, + d.qacc_integration, + d.qfrc_integration, + ) + + +@cache_kernel +def _tile_euler_dense(tile: TileSet): + @nested_kernel + def euler_dense( + # Model: + dof_damping: wp.array2d(dtype=float), + opt_timestep: wp.array(dtype=float), + # Data in: + qM_in: wp.array3d(dtype=float), + qfrc_smooth_in: wp.array2d(dtype=float), + qfrc_constraint_in: wp.array2d(dtype=float), + # In: + adr_in: wp.array(dtype=int), + # Data out: + qacc_integration_out: wp.array2d(dtype=float), + ): + worldid, nodeid = wp.tid() + timestep = opt_timestep[worldid] + TILE_SIZE = wp.static(tile.size) + + dofid = adr_in[nodeid] + M_tile = wp.tile_load(qM_in[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid)) + damping_tile = wp.tile_load(dof_damping[worldid], shape=(TILE_SIZE,), offset=(dofid,)) + damping_scaled = damping_tile * timestep + qm_integration_tile = wp.tile_diag_add(M_tile, damping_scaled) + + qfrc_smooth_tile = wp.tile_load(qfrc_smooth_in[worldid], shape=(TILE_SIZE,), offset=(dofid,)) + qfrc_constraint_tile = wp.tile_load(qfrc_constraint_in[worldid], shape=(TILE_SIZE,), offset=(dofid,)) + + qfrc_tile = qfrc_smooth_tile + qfrc_constraint_tile + + L_tile = wp.tile_cholesky(qm_integration_tile) + qacc_tile = wp.tile_cholesky_solve(L_tile, qfrc_tile) + wp.tile_store(qacc_integration_out[worldid], qacc_tile, offset=(dofid)) + + return euler_dense + + +@event_scope +def euler(m: Model, d: Data): + """Euler integrator, semi-implicit in velocity.""" + + # integrate damping implicitly + if not m.opt.disableflags & DisableBit.EULERDAMP.value: + if m.opt.is_sparse: + _euler_sparse(m, d) + else: + for tile in m.qM_tiles: + wp.launch_tiled( + _tile_euler_dense(tile), + dim=(d.nworld, tile.adr.size), + inputs=[m.dof_damping, m.opt.timestep, d.qM, d.qfrc_smooth, d.qfrc_constraint, tile.adr], + outputs=[d.qacc_integration], + block_dim=m.block_dim.euler_dense, + ) + + _advance(m, d, d.qacc_integration) + else: + _advance(m, d, d.qacc) + + +def _rk_perturb_state(m: Model, d: Data, scale: float): + # position + wp.launch( + _next_position, + dim=(d.nworld, m.njnt), + inputs=[m.opt.timestep, m.jnt_type, m.jnt_qposadr, m.jnt_dofadr, d.qpos_t0, d.qvel, scale], + outputs=[d.qpos], + ) + + # velocity + wp.launch( + _next_velocity, + dim=(d.nworld, m.nv), + inputs=[m.opt.timestep, d.qvel_t0, d.qacc, scale], + outputs=[d.qvel], + ) + + # activation + if m.na: + wp.launch( + _next_activation, + dim=(d.nworld, m.na), + inputs=[m.opt.timestep, d.act_t0, d.act_dot, scale, False], + outputs=[d.act], + ) + + +@wp.kernel +def _rk_accumulate_velocity_acceleration( + # Data in: + qvel_in: wp.array2d(dtype=float), + qacc_in: wp.array2d(dtype=float), + # In: + scale: float, + # Data out: + qvel_out: wp.array2d(dtype=float), + qacc_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + qvel_out[worldid, dofid] += scale * qvel_in[worldid, dofid] + qacc_out[worldid, dofid] += scale * qacc_in[worldid, dofid] + + +@wp.kernel +def _rk_accumulate_activation_velocity( + # Data in: + act_dot_in: wp.array2d(dtype=float), + # In: + scale: float, + # Data out: + act_dot_out: wp.array2d(dtype=float), +): + worldid, actid = wp.tid() + act_dot_out[worldid, actid] += scale * act_dot_in[worldid, actid] + + +def _rk_accumulate(m: Model, d: Data, scale: float): + """Computes one term of 1/6 k_1 + 1/3 k_2 + 1/3 k_3 + 1/6 k_4""" + + wp.launch( + _rk_accumulate_velocity_acceleration, + dim=(d.nworld, m.nv), + inputs=[d.qvel, d.qacc, scale], + outputs=[d.qvel_rk, d.qacc_rk], + ) + + if m.na: + wp.launch( + _rk_accumulate_activation_velocity, + dim=(d.nworld, m.na), + inputs=[d.act_dot, scale], + outputs=[d.act_dot_rk], + ) + + +@event_scope +def rungekutta4(m: Model, d: Data): + """Runge-Kutta explicit order 4 integrator.""" + + wp.copy(d.qpos_t0, d.qpos) + wp.copy(d.qvel_t0, d.qvel) + + d.qvel_rk.zero_() + d.qacc_rk.zero_() + d.act_dot_rk.zero_() + + if m.na: + wp.copy(d.act_t0, d.act) + + A, B = _RK4_A, _RK4_B + + _rk_accumulate(m, d, B[0]) + for i in range(3): + a, b = float(A[i][i]), B[i + 1] + _rk_perturb_state(m, d, a) + forward(m, d) + _rk_accumulate(m, d, b) + + wp.copy(d.qpos, d.qpos_t0) + wp.copy(d.qvel, d.qvel_t0) + if m.na: + wp.copy(d.act, d.act_t0) + wp.copy(d.act_dot, d.act_dot_rk) + _advance(m, d, d.qacc_rk, d.qvel_rk) + + +@event_scope +def implicit(m: Model, d: Data): + """Integrates fully implicit in velocity.""" + + # compile-time constants + passive_enabled = not m.opt.disableflags & DisableBit.PASSIVE.value + actuation_enabled = (not m.opt.disableflags & DisableBit.ACTUATION.value) and m.actuator_affine_bias_gain + + if passive_enabled or actuation_enabled: + derivative.deriv_smooth_vel(m, d) + smooth.factor_solve_i( + m, d, d.qM_integration, d.qLD_integration, d.qLDiagInv_integration, d.qacc_integration, d.qfrc_integration + ) + _advance(m, d, d.qacc_integration) + else: + _advance(m, d, d.qacc) + + +@event_scope +def fwd_position(m: Model, d: Data, factorize: bool = True): + """Position-dependent computations.""" + + smooth.kinematics(m, d) + smooth.com_pos(m, d) + smooth.camlight(m, d) + smooth.tendon(m, d) + smooth.crb(m, d) + smooth.tendon_armature(m, d) + if factorize: + smooth.factor_m(m, d) + if m.opt.run_collision_detection: + collision_driver.collision(m, d) + constraint.make_constraint(m, d) + smooth.transmission(m, d) + + +# TODO(team): sparse actuator_moment version +def _actuator_velocity(m: Model, d: Data): + NV = m.nv + + @kernel + def actuator_velocity( + # Data in: + qvel_in: wp.array2d(dtype=float), + actuator_moment_in: wp.array3d(dtype=float), + # Data out: + actuator_velocity_out: wp.array2d(dtype=float), + ): + worldid, actid = wp.tid() + moment_tile = wp.tile_load(actuator_moment_in[worldid, actid], shape=NV) + qvel_tile = wp.tile_load(qvel_in[worldid], shape=NV) + moment_qvel_tile = wp.tile_map(wp.mul, moment_tile, qvel_tile) + actuator_velocity_tile = wp.tile_reduce(wp.add, moment_qvel_tile) + actuator_velocity_out[worldid, actid] = actuator_velocity_tile[0] + + wp.launch_tiled( + actuator_velocity, + dim=(d.nworld, m.nu), + inputs=[ + d.qvel, + d.actuator_moment, + ], + outputs=[ + d.actuator_velocity, + ], + block_dim=m.block_dim.actuator_velocity, + ) + + +def _tendon_velocity(m: Model, d: Data): + NV = m.nv + + @kernel + def tendon_velocity( + # Data in: + qvel_in: wp.array2d(dtype=float), + ten_J_in: wp.array3d(dtype=float), + # Data out: + ten_velocity_out: wp.array2d(dtype=float), + ): + worldid, tenid = wp.tid() + ten_J_tile = wp.tile_load(ten_J_in[worldid, tenid], shape=NV) + qvel_tile = wp.tile_load(qvel_in[worldid], shape=NV) + ten_J_qvel_tile = wp.tile_map(wp.mul, ten_J_tile, qvel_tile) + ten_velocity_tile = wp.tile_reduce(wp.add, ten_J_qvel_tile) + ten_velocity_out[worldid, tenid] = ten_velocity_tile[0] + + wp.launch_tiled( + tendon_velocity, + dim=(d.nworld, m.ntendon), + inputs=[ + d.qvel, + d.ten_J, + ], + outputs=[ + d.ten_velocity, + ], + block_dim=m.block_dim.tendon_velocity, + ) + + +@event_scope +def fwd_velocity(m: Model, d: Data): + """Velocity-dependent computations.""" + + _actuator_velocity(m, d) + + if m.ntendon > 0: + # TODO(team): sparse version + _tendon_velocity(m, d) + + smooth.com_vel(m, d) + passive.passive(m, d) + smooth.rne(m, d) + smooth.tendon_bias(m, d, d.qfrc_bias) + + +@wp.kernel +def _actuator_force( + # Model: + na: int, + opt_timestep: wp.array(dtype=float), + actuator_dyntype: wp.array(dtype=int), + actuator_gaintype: wp.array(dtype=int), + actuator_biastype: wp.array(dtype=int), + actuator_actadr: wp.array(dtype=int), + actuator_actnum: wp.array(dtype=int), + actuator_ctrllimited: wp.array(dtype=bool), + actuator_forcelimited: wp.array(dtype=bool), + actuator_actlimited: wp.array(dtype=bool), + actuator_dynprm: wp.array2d(dtype=vec10f), + actuator_gainprm: wp.array2d(dtype=vec10f), + actuator_biasprm: wp.array2d(dtype=vec10f), + actuator_actearly: wp.array(dtype=bool), + actuator_ctrlrange: wp.array2d(dtype=wp.vec2), + actuator_forcerange: wp.array2d(dtype=wp.vec2), + actuator_actrange: wp.array2d(dtype=wp.vec2), + actuator_acc0: wp.array(dtype=float), + actuator_lengthrange: wp.array(dtype=wp.vec2), + # Data in: + act_in: wp.array2d(dtype=float), + ctrl_in: wp.array2d(dtype=float), + actuator_length_in: wp.array2d(dtype=float), + actuator_velocity_in: wp.array2d(dtype=float), + # In: + dsbl_clampctrl: int, + # Data out: + act_dot_out: wp.array2d(dtype=float), + actuator_force_out: wp.array2d(dtype=float), +): + worldid, uid = wp.tid() + + ctrl = ctrl_in[worldid, uid] + + if actuator_ctrllimited[uid] and not dsbl_clampctrl: + ctrlrange = actuator_ctrlrange[worldid, uid] + ctrl = wp.clamp(ctrl, ctrlrange[0], ctrlrange[1]) + ctrl_act = ctrl + + act_first = actuator_actadr[uid] + if na and act_first >= 0: + act_last = act_first + actuator_actnum[uid] - 1 + dyntype = actuator_dyntype[uid] + + if dyntype == int(DynType.INTEGRATOR.value): + act_dot = ctrl + elif dyntype == int(DynType.FILTER.value) or dyntype == int(DynType.FILTEREXACT.value): + dynprm = actuator_dynprm[worldid, uid] + act = act_in[worldid, act_last] + act_dot = (ctrl - act) / wp.max(dynprm[0], MJ_MINVAL) + elif dyntype == int(DynType.MUSCLE.value): + dynprm = actuator_dynprm[worldid, uid] + act = act_in[worldid, act_last] + act_dot = util_misc.muscle_dynamics(ctrl, act, dynprm) + else: # DynType.NONE + act_dot = 0.0 + + act_dot_out[worldid, act_last] = act_dot + + if actuator_actearly[uid]: + if dyntype == int(DynType.INTEGRATOR.value) or dyntype == int(DynType.NONE.value): + dynprm = actuator_dynprm[worldid, uid] + act = act_in[worldid, act_last] + + ctrl_act = _next_act( + opt_timestep[worldid], + dyntype, + dynprm, + actuator_actrange[worldid, uid], + act, + act_dot, + 1.0, + actuator_actlimited[uid], + ) + else: + ctrl_act = act_in[worldid, act_last] + + length = actuator_length_in[worldid, uid] + velocity = actuator_velocity_in[worldid, uid] + + # gain + gaintype = actuator_gaintype[uid] + gainprm = actuator_gainprm[worldid, uid] + + gain = 0.0 + if gaintype == int(GainType.FIXED.value): + gain = gainprm[0] + elif gaintype == int(GainType.AFFINE.value): + gain = gainprm[0] + gainprm[1] * length + gainprm[2] * velocity + elif gaintype == int(GainType.MUSCLE.value): + acc0 = actuator_acc0[uid] + lengthrange = actuator_lengthrange[uid] + gain = util_misc.muscle_gain(length, velocity, lengthrange, acc0, gainprm) + + # bias + biastype = actuator_biastype[uid] + biasprm = actuator_biasprm[worldid, uid] + + bias = 0.0 # BiasType.NONE + if biastype == int(BiasType.AFFINE.value): + bias = biasprm[0] + biasprm[1] * length + biasprm[2] * velocity + elif biastype == int(BiasType.MUSCLE.value): + acc0 = actuator_acc0[uid] + lengthrange = actuator_lengthrange[uid] + bias = util_misc.muscle_bias(length, lengthrange, acc0, biasprm) + + force = gain * ctrl_act + bias + + # TODO(team): tendon total force clamping + + if actuator_forcelimited[uid]: + forcerange = actuator_forcerange[worldid, uid] + force = wp.clamp(force, forcerange[0], forcerange[1]) + + actuator_force_out[worldid, uid] = force + + +@wp.kernel +def _tendon_actuator_force( + # Model: + actuator_trntype: wp.array(dtype=int), + actuator_trnid: wp.array(dtype=wp.vec2i), + # Data in: + actuator_force_in: wp.array2d(dtype=float), + # Data out: + ten_actfrc_out: wp.array2d(dtype=float), +): + worldid, actid = wp.tid() + + if actuator_trntype[actid] == int(TrnType.TENDON.value): + tenid = actuator_trnid[actid][0] + # TODO(team): only compute for tendons with force limits? + wp.atomic_add(ten_actfrc_out[worldid], tenid, actuator_force_in[worldid, actid]) + + +@wp.kernel +def _tendon_actuator_force_clamp( + # Model: + actuator_trntype: wp.array(dtype=int), + actuator_trnid: wp.array(dtype=wp.vec2i), + tendon_actfrclimited: wp.array(dtype=bool), + tendon_actfrcrange: wp.array2d(dtype=wp.vec2), + # Data in: + ten_actfrc_in: wp.array2d(dtype=float), + # Data out: + actuator_force_out: wp.array2d(dtype=float), +): + worldid, actid = wp.tid() + + if actuator_trntype[actid] == int(TrnType.TENDON.value): + tenid = actuator_trnid[actid][0] + if tendon_actfrclimited[tenid]: + ten_actfrc = ten_actfrc_in[worldid, tenid] + actfrcrange = tendon_actfrcrange[worldid, tenid] + + if ten_actfrc < actfrcrange[0]: + actuator_force_out[worldid, actid] *= actfrcrange[0] / ten_actfrc + elif ten_actfrc > actfrcrange[1]: + actuator_force_out[worldid, actid] *= actfrcrange[1] / ten_actfrc + + +def _qfrc_actuator(m: Model, d: Data): + NU = m.nu + + @wp.kernel + def qfrc_actuator( + # Model: + ngravcomp: int, + jnt_actfrclimited: wp.array(dtype=bool), + jnt_actfrcrange: wp.array2d(dtype=wp.vec2), + jnt_actgravcomp: wp.array(dtype=int), + dof_jntid: wp.array(dtype=int), + # Data in: + actuator_moment_in: wp.array3d(dtype=float), + qfrc_gravcomp_in: wp.array2d(dtype=float), + actuator_force_in: wp.array2d(dtype=float), + # Data out: + qfrc_actuator_out: wp.array2d(dtype=float), + ): + worldid, dofid = wp.tid() + + actuator_moment_tile = wp.tile_load(actuator_moment_in[worldid], shape=(NU, 1), offset=(0, dofid)) + actuator_moment_tile = wp.tile_squeeze(actuator_moment_tile, axis=(1,)) + actuator_force_tile = wp.tile_load(actuator_force_in[worldid], shape=NU) + actuator_moment_force_tile = wp.tile_map(wp.mul, actuator_moment_tile, actuator_force_tile) + qfrc_tile = wp.tile_reduce(wp.add, actuator_moment_force_tile) + qfrc = qfrc_tile[0] + + jntid = dof_jntid[dofid] + + # actuator-level gravity compensation, skip if added as passive force + if ngravcomp and jnt_actgravcomp[jntid]: + qfrc += qfrc_gravcomp_in[worldid, dofid] + + if jnt_actfrclimited[jntid]: + frcrange = jnt_actfrcrange[worldid, jntid] + qfrc = wp.clamp(qfrc, frcrange[0], frcrange[1]) + + qfrc_actuator_out[worldid, dofid] = qfrc + + wp.launch_tiled( + qfrc_actuator, + dim=(d.nworld, m.nv), + inputs=[ + m.ngravcomp, + m.jnt_actfrclimited, + m.jnt_actfrcrange, + m.jnt_actgravcomp, + m.dof_jntid, + d.actuator_moment, + d.qfrc_gravcomp, + d.actuator_force, + ], + outputs=[d.qfrc_actuator], + block_dim=m.block_dim.qfrc_actuator, + ) + + +@event_scope +def fwd_actuation(m: Model, d: Data): + """Actuation-dependent computations.""" + if not m.nu or (m.opt.disableflags & DisableBit.ACTUATION): + d.act_dot.zero_() + d.qfrc_actuator.zero_() + return + + wp.launch( + _actuator_force, + dim=(d.nworld, m.nu), + inputs=[ + m.na, + m.opt.timestep, + m.actuator_dyntype, + m.actuator_gaintype, + m.actuator_biastype, + m.actuator_actadr, + m.actuator_actnum, + m.actuator_ctrllimited, + m.actuator_forcelimited, + m.actuator_actlimited, + m.actuator_dynprm, + m.actuator_gainprm, + m.actuator_biasprm, + m.actuator_actearly, + m.actuator_ctrlrange, + m.actuator_forcerange, + m.actuator_actrange, + m.actuator_acc0, + m.actuator_lengthrange, + d.act, + d.ctrl, + d.actuator_length, + d.actuator_velocity, + m.opt.disableflags & DisableBit.CLAMPCTRL, + ], + outputs=[d.act_dot, d.actuator_force], + ) + + if m.ntendon: + d.ten_actfrc.zero_() + + wp.launch( + _tendon_actuator_force, + dim=(d.nworld, m.nu), + inputs=[ + m.actuator_trntype, + m.actuator_trnid, + d.actuator_force, + ], + outputs=[d.ten_actfrc], + ) + + wp.launch( + _tendon_actuator_force_clamp, + dim=(d.nworld, m.nu), + inputs=[ + m.actuator_trntype, + m.actuator_trnid, + m.tendon_actfrclimited, + m.tendon_actfrcrange, + d.ten_actfrc, + ], + outputs=[d.actuator_force], + ) + + _qfrc_actuator(m, d) + + +@wp.kernel +def _qfrc_smooth( + # Data in: + qfrc_applied_in: wp.array2d(dtype=float), + qfrc_bias_in: wp.array2d(dtype=float), + qfrc_passive_in: wp.array2d(dtype=float), + qfrc_actuator_in: wp.array2d(dtype=float), + # Data out: + qfrc_smooth_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + qfrc_smooth_out[worldid, dofid] = ( + qfrc_passive_in[worldid, dofid] + - qfrc_bias_in[worldid, dofid] + + qfrc_actuator_in[worldid, dofid] + + qfrc_applied_in[worldid, dofid] + ) + + +@event_scope +def fwd_acceleration(m: Model, d: Data, factorize: bool = False): + """Add up all non-constraint forces, compute qacc_smooth.""" + + wp.launch( + _qfrc_smooth, + dim=(d.nworld, m.nv), + inputs=[ + d.qfrc_applied, + d.qfrc_bias, + d.qfrc_passive, + d.qfrc_actuator, + ], + outputs=[ + d.qfrc_smooth, + ], + ) + xfrc_accumulate(m, d, d.qfrc_smooth) + + if factorize: + smooth.factor_solve_i(m, d, d.qM, d.qLD, d.qLDiagInv, d.qacc_smooth, d.qfrc_smooth) + else: + smooth.solve_m(m, d, d.qacc_smooth, d.qfrc_smooth) + + +@wp.kernel +def _zero_energy( + # Data out: + energy_out: wp.array(dtype=wp.vec2), +): + tid = wp.tid() + energy_out[tid] = wp.vec2(0.0, 0.0) + + +@event_scope +def forward(m: Model, d: Data): + """Forward dynamics.""" + energy = m.opt.enableflags & EnableBit.ENERGY + + fwd_position(m, d, factorize=False) + sensor.sensor_pos(m, d) + + if energy: + if m.sensor_e_potential == 0: # not computed by sensor + sensor.energy_pos(m, d) + else: + wp.launch( + _zero_energy, + dim=d.nworld, + inputs=[d.energy], + ) + + fwd_velocity(m, d) + sensor.sensor_vel(m, d) + + if energy: + if m.sensor_e_kinetic == 0: # not computed by sensor + sensor.energy_vel(m, d) + + fwd_actuation(m, d) + fwd_acceleration(m, d, factorize=True) + sensor.sensor_acc(m, d) + + solver.solve(m, d) + + +@event_scope +def step(m: Model, d: Data): + """Advance simulation.""" + forward(m, d) + + if m.opt.integrator == IntegratorType.EULER: + euler(m, d) + elif m.opt.integrator == IntegratorType.RK4: + rungekutta4(m, d) + elif m.opt.integrator == IntegratorType.IMPLICITFAST: + implicit(m, d) + else: + raise NotImplementedError(f"integrator {m.opt.integrator} not implemented.") diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py new file mode 100644 index 00000000..f18d4150 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py @@ -0,0 +1,285 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Tests for forward dynamics functions.""" + +import mujoco +import numpy as np +import warp as wp +from absl.testing import absltest +from absl.testing import parameterized + +import mujoco_warp as mjwarp + +from mujoco.mjx.third_party.mujoco_warp._src import test_util +from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType +from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import GainType +from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType + +# tolerance for difference between MuJoCo and mjwarp smooth calculations - mostly +# due to float precision +_TOLERANCE = 5e-5 + + +def _assert_eq(a, b, name): + tol = _TOLERANCE * 10 # avoid test noise + err_msg = f"mismatch: {name}" + np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol) + + +class ForwardTest(parameterized.TestCase): + # TODO(team): test sparse when actuator_moment and/or ten_J have sparse representation + @parameterized.product(xml=["humanoid/humanoid.xml", "pendula.xml"]) + def test_fwd_velocity(self, xml): + _, mjd, m, d = test_util.fixture(xml, kick=True) + + for arr in (d.actuator_velocity, d.qfrc_bias): + arr.zero_() + + mjwarp.fwd_velocity(m, d) + + _assert_eq(d.actuator_velocity.numpy()[0], mjd.actuator_velocity, "actuator_velocity") + _assert_eq(d.qfrc_bias.numpy()[0], mjd.qfrc_bias, "qfrc_bias") + + def test_fwd_velocity_tendon(self): + _, mjd, m, d = test_util.fixture("tendon/fixed.xml", sparse=False) + + d.ten_velocity.zero_() + mjwarp.fwd_velocity(m, d) + + _assert_eq(d.ten_velocity.numpy()[0], mjd.ten_velocity, "ten_velocity") + + @parameterized.parameters( + ("actuation/actuation.xml", True), + ("actuation/actuation.xml", False), + ("actuation/actuators.xml", True), + ("actuation/actuators.xml", False), + ("actuation/muscle.xml", True), + ("actuation/muscle.xml", False), + ) + def test_actuation(self, xml, actuation): + mjm, mjd, m, d = test_util.fixture(xml, actuation=actuation, keyframe=0) + + for arr in (d.qfrc_actuator, d.actuator_force, d.act_dot): + arr.zero_() + + mjwarp.fwd_actuation(m, d) + + _assert_eq(d.qfrc_actuator.numpy()[0], mjd.qfrc_actuator, "qfrc_actuator") + _assert_eq(d.actuator_force.numpy()[0], mjd.actuator_force, "actuator_force") + + if mjm.na: + _assert_eq(d.act_dot.numpy()[0], mjd.act_dot, "act_dot") + + # next activations + mujoco.mj_step(mjm, mjd) + mjwarp.step(m, d) + + _assert_eq(d.act.numpy()[0], mjd.act, "act") + + # TODO(team): test actearly + + @parameterized.parameters(True, False) + def test_clampctrl(self, clampctrl): + _, mjd, _, d = test_util.fixture( + xml=""" + + + + + + + + + + + + + + + """, + clampctrl=clampctrl, + keyframe=0, + ) + + _assert_eq(d.ctrl.numpy()[0], mjd.ctrl, "ctrl") + + def test_fwd_acceleration(self): + _, mjd, m, d = test_util.fixture("humanoid/humanoid.xml", kick=True) + + for arr in (d.qfrc_smooth, d.qacc_smooth): + arr.zero_() + + mjwarp.fwd_acceleration(m, d) + + _assert_eq(d.qfrc_smooth.numpy()[0], mjd.qfrc_smooth, "qfrc_smooth") + _assert_eq(d.qacc_smooth.numpy()[0], mjd.qacc_smooth, "qacc_smooth") + + @parameterized.parameters((True, True), (True, False), (False, True), (False, False)) + def test_euler(self, eulerdamp, sparse): + mjm, mjd, _, _ = test_util.fixture("pendula.xml", kick=True, eulerdamp=eulerdamp, sparse=sparse) + self.assertTrue((mjm.dof_damping > 0).any()) + + mjd.qvel[:] = 1.0 + mjd.qacc[:] = 1.0 + mujoco.mj_forward(mjm, mjd) + + m = mjwarp.put_model(mjm) + d = mjwarp.put_data(mjm, mjd) + + mujoco.mj_Euler(mjm, mjd) + mjwarp.euler(m, d) + + _assert_eq(d.qpos.numpy()[0], mjd.qpos, "qpos") + _assert_eq(d.qvel.numpy()[0], mjd.qvel, "qvel") + _assert_eq(d.act.numpy()[0], mjd.act, "act") + + def test_rungekutta4(self): + mjm, mjd, m, d = test_util.fixture( + xml=""" + + + + + + + + + + + + + + + + + """, + keyframe=0, + ) + + mjwarp.rungekutta4(m, d) + mujoco.mj_RungeKutta(mjm, mjd, 4) + + _assert_eq(d.qpos.numpy()[0], mjd.qpos, "qpos") + _assert_eq(d.qvel.numpy()[0], mjd.qvel, "qvel") + _assert_eq(d.time.numpy()[0], mjd.time, "time") + _assert_eq(d.xpos.numpy()[0], mjd.xpos, "xpos") + + # test rungekutta determinism + def rk_step() -> wp.array(dtype=wp.float32, ndim=2): + d.qpos = wp.ones_like(d.qpos) + d.qvel = wp.ones_like(d.qvel) + d.act = wp.ones_like(d.act) + mjwarp.rungekutta4(m, d) + return d.qpos + + _assert_eq(rk_step().numpy()[0], rk_step().numpy()[0], "qpos") + + @parameterized.product(actuation=[True, False], passive=[True, False], sparse=[True, False]) + def test_implicit(self, actuation, passive, sparse): + mjm, mjd, _, _ = test_util.fixture( + "pendula.xml", + integrator=IntegratorType.IMPLICITFAST, + actuation=actuation, + passive=passive, + sparse=sparse, + ) + + mjm.actuator_gainprm[:, 2] = np.random.uniform(low=0.01, high=10.0, size=mjm.actuator_gainprm[:, 2].shape) + + # change actuators to velocity/damper to cover all codepaths + mjm.actuator_gaintype[3] = GainType.AFFINE + mjm.actuator_gaintype[6] = GainType.AFFINE + mjm.actuator_biastype[0:3] = BiasType.AFFINE + mjm.actuator_biastype[4:6] = BiasType.AFFINE + mjm.actuator_biasprm[0:3, 2] = -1.0 + mjm.actuator_biasprm[4:6, 2] = -1.0 + mjm.actuator_ctrlrange[3:7] = 10.0 + mjm.actuator_gear[:] = 1.0 + + mjd.qvel = np.random.uniform(low=-0.01, high=0.01, size=mjd.qvel.shape) + mjd.ctrl = np.random.uniform(low=-0.1, high=0.1, size=mjd.ctrl.shape) + mjd.act = np.random.uniform(low=-0.1, high=0.1, size=mjd.act.shape) + mujoco.mj_forward(mjm, mjd) + + m = mjwarp.put_model(mjm) + d = mjwarp.put_data(mjm, mjd) + + mjwarp.implicit(m, d) + mujoco.mj_implicit(mjm, mjd) + + _assert_eq(d.qpos.numpy()[0], mjd.qpos, "qpos") + _assert_eq(d.act.numpy()[0], mjd.act, "act") + + def test_implicit_position(self): + mjm, mjd, m, d = test_util.fixture("actuation/position.xml", keyframe=0, integrator=IntegratorType.IMPLICITFAST, kick=True) + + mujoco.mj_implicit(mjm, mjd) + mjwarp.implicit(m, d) + + _assert_eq(d.qpos.numpy()[0], mjd.qpos, "qpos") + _assert_eq(d.qvel.numpy()[0], mjd.qvel, "qvel") + + def test_implicit_tendon_damping(self): + mjm, mjd, m, d = test_util.fixture("tendon/damping.xml", keyframe=0, integrator=IntegratorType.IMPLICITFAST, kick=True) + + mujoco.mj_implicit(mjm, mjd) + mjwarp.implicit(m, d) + + _assert_eq(d.qpos.numpy()[0], mjd.qpos, "qpos") + _assert_eq(d.qvel.numpy()[0], mjd.qvel, "qvel") + + @parameterized.product( + xml=("humanoid/humanoid.xml", "pendula.xml", "constraints.xml", "collision.xml"), graph_conditional=(True, False) + ) + def test_graph_capture(self, xml, graph_conditional): + # TODO(team): test more environments + if wp.get_device().is_cuda and wp.config.verify_cuda == False: + _, _, m, d = test_util.fixture(xml) + m.opt.graph_conditional = graph_conditional + + with wp.ScopedCapture() as capture: + mjwarp.step(m, d) + + # step a few times to ensure no errors at the step boundary + wp.capture_launch(capture.graph) + wp.capture_launch(capture.graph) + wp.capture_launch(capture.graph) + + self.assertTrue(d.time.numpy()[0] > 0.0) + + def test_forward_energy(self): + _, mjd, _, d = test_util.fixture("humanoid/humanoid.xml", kick=True, energy=True) + + _assert_eq(d.energy.numpy()[0][0], mjd.energy[0], "potential energy") + _assert_eq(d.energy.numpy()[0][1], mjd.energy[1], "kinetic energy") + + def test_tendon_actuator_force_limits(self): + for keyframe in range(7): + _, mjd, m, d = test_util.fixture("actuation/tendon_force_limit.xml", keyframe=keyframe) + + d.actuator_force.zero_() + + mjwarp.forward(m, d) + + _assert_eq(d.actuator_force.numpy()[0], mjd.actuator_force, "actuator_force") + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py new file mode 100644 index 00000000..3856bfc1 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py @@ -0,0 +1,147 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src import derivative +from mujoco.mjx.third_party.mujoco_warp._src import forward +from mujoco.mjx.third_party.mujoco_warp._src import sensor +from mujoco.mjx.third_party.mujoco_warp._src import smooth +from mujoco.mjx.third_party.mujoco_warp._src import solver +from mujoco.mjx.third_party.mujoco_warp._src import support +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model + + +@wp.kernel +def _qfrc_eulerdamp( + # Model: + opt_timestep: wp.array(dtype=float), + dof_damping: wp.array2d(dtype=float), + # Data in: + qacc_in: wp.array2d(dtype=float), + # Out: + qfrc_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + timestep = opt_timestep[worldid] + qfrc_out[worldid, dofid] += timestep * dof_damping[worldid, dofid] * qacc_in[worldid, dofid] + + +@wp.kernel +def _qfrc_inverse( + # Data in: + qfrc_bias_in: wp.array2d(dtype=float), + qfrc_passive_in: wp.array2d(dtype=float), + qfrc_constraint_in: wp.array2d(dtype=float), + # In: + Ma: wp.array2d(dtype=float), + # Data out: + qfrc_inverse_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + + qfrc_inverse = qfrc_bias_in[worldid, dofid] + qfrc_inverse += Ma[worldid, dofid] + qfrc_inverse -= qfrc_passive_in[worldid, dofid] + qfrc_inverse -= qfrc_constraint_in[worldid, dofid] + + qfrc_inverse_out[worldid, dofid] = qfrc_inverse + + +def discrete_acc(m: Model, d: Data, qacc: wp.array2d(dtype=float), qfrc: wp.array2d(dtype=float)): + """Convert discrete-time qacc to continuous-time qacc.""" + + if m.opt.integrator == IntegratorType.RK4: + raise NotImplementedError("discrete inverse dynamics is not supported by RK4 integrator") + elif m.opt.integrator == IntegratorType.EULER: + if m.opt.disableflags & DisableBit.EULERDAMP: + wp.copy(qacc, d.qacc) + return + + # TODO(team): qacc = d.qacc if (m.dof_damping == 0.0).all() + + # set qfrc = (d.qM + m.opt.timestep * diag(m.dof_damping)) * d.qacc + + # d.qM @ d.qacc + support.mul_m(m, d, qfrc, d.qacc, d.inverse_mul_m_skip) + + # qfrc += m.opt.timestep * m.dof_damping * d.qacc + wp.launch( + _qfrc_eulerdamp, + dim=(d.nworld, m.nv), + inputs=[m.opt.timestep, m.dof_damping, d.qacc], + outputs=[qfrc], + ) + elif m.opt.integrator == IntegratorType.IMPLICITFAST: + derivative.deriv_smooth_vel(m, d, flg_forward=False) + smooth.factor_solve_i(m, d, d.qM, d.qLD, d.qLDiagInv, qacc, qfrc) + else: + raise NotImplementedError(f"integrator {m.opt.integrator} not implemented.") + + # solve for qacc: qfrc = d.qM @ d.qacc + smooth.solve_m(m, d, qacc, qfrc) + + +def inv_constraint(m: Model, d: Data): + """Inverse constraint solver.""" + + # no constraints + if d.njmax == 0: + d.qfrc_constraint.zero_() + return + + # update + solver.create_context(m, d, grad=False) + + +def inverse(m: Model, d: Data): + """Inverse dynamics.""" + forward.fwd_position(m, d) + sensor.sensor_pos(m, d) + forward.fwd_velocity(m, d) + sensor.sensor_vel(m, d) + + invdiscrete = m.opt.enableflags & EnableBit.INVDISCRETE + if invdiscrete: + # save discrete-time qacc and compute continuous-time qacc + wp.copy(d.qacc_discrete, d.qacc) + discrete_acc(m, d, d.qacc, d.qfrc_integration) + + inv_constraint(m, d) + smooth.rne(m, d) + smooth.tendon_bias(m, d, d.qfrc_bias) + sensor.sensor_acc(m, d) + + support.mul_m(m, d, d.qfrc_inverse, d.qacc, d.inverse_mul_m_skip) + + wp.launch( + _qfrc_inverse, + dim=(d.nworld, m.nv), + inputs=[ + d.qfrc_bias, + d.qfrc_passive, + d.qfrc_constraint, + d.qfrc_inverse, + ], + outputs=[d.qfrc_inverse], + ) + + if invdiscrete: + # restore discrete-time qacc + wp.copy(d.qacc, d.qacc_discrete) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse_test.py new file mode 100644 index 00000000..e6a7a065 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse_test.py @@ -0,0 +1,169 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Tests for inverse dynamics.""" + +import mujoco +import numpy as np +import warp as wp +from absl.testing import absltest +from absl.testing import parameterized + +import mujoco_warp as mjwarp + +from mujoco.mjx.third_party.mujoco_warp._src import inverse +from mujoco.mjx.third_party.mujoco_warp._src import support +from mujoco.mjx.third_party.mujoco_warp._src import test_util +from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType + + +def _assert_eq(a, b, name): + tol = 5e-3 # avoid test noise + err_msg = f"mismatch: {name}" + np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol) + + +_XML = """ + + +""" + + +class InverseTest(parameterized.TestCase): + @parameterized.product( + integrator=[IntegratorType.EULER, IntegratorType.IMPLICITFAST], + invdiscrete=[True, False], + sparse=[True, False], + ) + def test_inverse(self, integrator, invdiscrete, sparse): + """Tests inverse dynamics.""" + mjm, mjd, m, d = test_util.fixture( + xml=_XML, + contact=False, + integrator=integrator, + kick=True, + applied=True, + nstep=10, + sparse=sparse, + ) + + # discrete qacc + if invdiscrete: + mjm.opt.enableflags |= mujoco.mjtEnableBit.mjENBL_INVDISCRETE + + # save state + qpos = mjd.qpos.copy() + qvel = mjd.qvel.copy() + + # call step, save new qvel + mujoco.mj_step(mjm, mjd) + qvel_next = mjd.qvel.copy() + + # reset the state, compute discrete-time (finite-differenced) qacc + mjd.qpos = qpos + mjd.qvel = qvel + qacc_fd = (qvel_next - qvel) / mjm.opt.timestep + + # call forward, overwrite qacc with qacc_fd + mujoco.mj_forward(mjm, mjd) + mjd.qacc = qacc_fd + + m = mjwarp.put_model(mjm) + d = mjwarp.put_data(mjm, mjd) + + qacc = d.qacc.numpy()[0].copy() + qfrc_constraint = d.qfrc_constraint.numpy()[0].copy() + + # qfrc_inverse = qfrc_applied + J.T @ xfrc_applied + qfrc_actuator + qfrc_xfrc_applied = wp.zeros((d.nworld, m.nv), dtype=float) + support.xfrc_accumulate(m, d, qfrc_xfrc_applied) + qfrc_inverse = d.qfrc_applied.numpy()[0] + d.qfrc_actuator.numpy()[0] + qfrc_xfrc_applied.numpy()[0] + + for arr in (d.qfrc_constraint, d.qfrc_inverse): + arr.zero_() + + mjwarp.inverse(m, d) + + _assert_eq(d.qfrc_constraint.numpy()[0], qfrc_constraint, "qfrc_constraint") + _assert_eq(d.qfrc_inverse.numpy()[0], qfrc_inverse, "qfrc_inverse") + _assert_eq(d.qacc.numpy()[0], qacc, "qacc") + + def test_discrete_acc_eulerdamp(self): + _, _, m, d = test_util.fixture( + xml=_XML, integrator=IntegratorType.EULER, eulerdamp=False, kick=True, applied=True, nstep=10 + ) + qacc = wp.zeros((1, m.nv), dtype=float) + qfrc = wp.zeros((1, m.nv), dtype=float) + inverse.discrete_acc(m, d, qacc, qfrc) + _assert_eq(qacc.numpy()[0], d.qacc.numpy()[0], "qacc") + + def test_discrete_acc_rk4(self): + _, _, m, d = test_util.fixture(xml=_XML, integrator=IntegratorType.RK4) + qacc = wp.zeros((1, m.nv), dtype=float) + qfrc = wp.zeros((1, m.nv), dtype=float) + + with self.assertRaises(NotImplementedError): + inverse.discrete_acc(m, d, qacc, qfrc) + + def test_inverse_tendon_armature(self): + """Tests inverse dynamics with tendon armature.""" + _, _, m, d = test_util.fixture( + "tendon/armature.xml", + constraint=False, + gravity=False, + kick=True, + applied=True, + nstep=10, + keyframe=0, + ) + + qacc = d.qacc.numpy()[0].copy() + qfrc_constraint = d.qfrc_constraint.numpy()[0].copy() + + # qfrc_inverse = qfrc_applied + J.T @ xfrc_applied + qfrc_actuator + qfrc_xfrc_applied = wp.zeros((d.nworld, m.nv), dtype=float) + support.xfrc_accumulate(m, d, qfrc_xfrc_applied) + qfrc_inverse = d.qfrc_applied.numpy()[0] + d.qfrc_actuator.numpy()[0] + qfrc_xfrc_applied.numpy()[0] + + for arr in (d.qfrc_constraint, d.qfrc_inverse): + arr.zero_() + + mjwarp.inverse(m, d) + + _assert_eq(d.qfrc_constraint.numpy()[0], qfrc_constraint, "qfrc_constraint") + _assert_eq(d.qfrc_inverse.numpy()[0], qfrc_inverse, "qfrc_inverse") + _assert_eq(d.qacc.numpy()[0], qacc, "qacc") + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py new file mode 100644 index 00000000..f2f2faa5 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py @@ -0,0 +1,1598 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +from typing import Optional, Tuple + +import mujoco +import numpy as np +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import conditional_graph_supported + +from mujoco.mjx.third_party.mujoco_warp._src import math +from mujoco.mjx.third_party.mujoco_warp._src import types + +# number of max iterations to run GJK/EPA +MJ_CCD_ITERATIONS = 12 + + +def _hfield_geom_pair(mjm: mujoco.MjModel) -> Tuple[int, np.array]: + geom1, geom2 = np.triu_indices(mjm.ngeom, k=1) + geom_type_hf = mujoco.mjtGeom.mjGEOM_HFIELD + has_hfield = (mjm.geom_type[geom1] == geom_type_hf) | (mjm.geom_type[geom2] == geom_type_hf) + nhfieldgeompair = np.sum(has_hfield) + geompair2hfgeompair = -1 * np.ones(mjm.ngeom * (mjm.ngeom - 1) // 2, dtype=int) + geompair2hfgeompair[has_hfield] = np.arange(nhfieldgeompair) + + return nhfieldgeompair, geompair2hfgeompair + + +def put_model(mjm: mujoco.MjModel) -> types.Model: + """ + Creates a model on device. + + Args: + mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host). + + Returns: + Model: The model containing kinematic and dynamic information (device). + """ + # check supported features + for field, field_types, field_str in ( + (mjm.actuator_trntype, types.TrnType, "Actuator transmission type"), + (mjm.actuator_dyntype, types.DynType, "Actuator dynamics type"), + (mjm.actuator_gaintype, types.GainType, "Gain type"), + (mjm.actuator_biastype, types.BiasType, "Bias type"), + (mjm.eq_type, types.EqType, "Equality constraint types"), + (mjm.geom_type, types.GeomType, "Geom type"), + (mjm.sensor_type, types.SensorType, "Sensor types"), + (mjm.wrap_type, types.WrapType, "Wrap types"), + ): + unsupported = ~np.isin(field, list(field_types)) + if unsupported.any(): + raise NotImplementedError(f"{field_str} {field[unsupported]} not supported.") + + plugin_id = [] + plugin_attr = [] + geom_plugin_index = np.full_like(mjm.geom_type, -1) + + if mjm.nplugin > 0: + for i in range(len(mjm.geom_plugin)): + if mjm.geom_plugin[i] != -1: + p = mjm.geom_plugin[i] + geom_plugin_index[i] = len(plugin_id) + plugin_id.append(mjm.plugin[p]) + start = mjm.plugin_attradr[p] + end = mjm.plugin_attradr[p + 1] if p + 1 < mjm.nplugin else len(mjm.plugin_attr) + values = mjm.plugin_attr[start:end] + attr_values = [] + current = [] + for v in values: + if v == 0: + if current: + s = "".join(chr(int(x)) for x in current) + attr_values.append(float(s)) + current = [] + else: + current.append(v) + # Pad with zeros if less than 3 + attr_values += [0.0] * (3 - len(attr_values)) + plugin_attr.append(attr_values[:3]) + + plugin_id = np.array(plugin_id) + plugin_attr = np.array(plugin_attr) + + if mjm.nflex > 1: + raise NotImplementedError("Only one flex is unsupported.") + + if ((mjm.flex_contype != 0) | (mjm.flex_conaffinity != 0)).any(): + raise NotImplementedError("Flex collisions are not implemented.") + + if mjm.geom_fluid.any(): + raise NotImplementedError("Ellipsoid fluid model not implemented.") + + # check options + for opt, opt_types, msg in ( + (mjm.opt.integrator, types.IntegratorType, "Integrator"), + (mjm.opt.cone, types.ConeType, "Cone"), + (mjm.opt.solver, types.SolverType, "Solver"), + ): + if opt not in set(opt_types): + raise NotImplementedError(f"{msg} {opt} is unsupported.") + + if mjm.opt.noslip_iterations > 0: + raise NotImplementedError(f"noslip solver not implemented.") + + # TODO(team): remove after _update_gradient for Newton uses tile operations for islands + nv_max = 60 + if mjm.nv > nv_max and mjm.opt.jacobian == mujoco.mjtJacobian.mjJAC_DENSE: + raise ValueError(f"Dense is unsupported for nv > {nv_max} (nv = {mjm.nv}).") + + is_sparse = mujoco.mj_isSparse(mjm) + + # calculate some fields that cannot be easily computed inline + nlsp = mjm.opt.ls_iterations # TODO(team): how to set nlsp? + + # unfortunately we must create Data in order to get some model fields like M_rownnz + mjd = mujoco.MjData(mjm) + + # dof lower triangle row and column indices (used in solver) + dof_tri_row, dof_tri_col = np.tril_indices(mjm.nv) + + # indices for sparse qM_fullm (used in solver) + qM_fullm_i, qM_fullm_j = [], [] + for i in range(mjm.nv): + j = i + while j > -1: + qM_fullm_i.append(i) + qM_fullm_j.append(j) + j = mjm.dof_parentid[j] + + # indices for sparse qM mul_m (used in support) + qM_mulm_i, qM_mulm_j, qM_madr_ij = [], [], [] + for i in range(mjm.nv): + madr_ij, j = mjm.dof_Madr[i], i + + while True: + madr_ij, j = madr_ij + 1, mjm.dof_parentid[j] + if j == -1: + break + qM_mulm_i.append(i) + qM_mulm_j.append(j) + qM_madr_ij.append(madr_ij) + + # body_tree is a list of body ids grouped by tree level + bodies, body_depth = {}, np.zeros(mjm.nbody, dtype=int) - 1 + for i in range(mjm.nbody): + body_depth[i] = body_depth[mjm.body_parentid[i]] + 1 + bodies.setdefault(body_depth[i], []).append(i) + body_tree = tuple(wp.array(bodies[i], dtype=int) for i in sorted(bodies)) + + # qLD_updates has dof tree ordering of qLD updates for sparse factor m + qLD_updates, dof_depth = {}, np.zeros(mjm.nv, dtype=int) - 1 + + for k in range(mjm.nv): + # skip diagonal rows + if mjd.M_rownnz[k] == 1: + continue + dof_depth[k] = dof_depth[mjm.dof_parentid[k]] + 1 + i = mjm.dof_parentid[k] + diag_k = mjd.M_rowadr[k] + mjd.M_rownnz[k] - 1 + Madr_ki = diag_k - 1 + while i > -1: + qLD_updates.setdefault(dof_depth[i], []).append((i, k, Madr_ki)) + i = mjm.dof_parentid[i] + Madr_ki -= 1 + + qLD_updates = tuple(wp.array(qLD_updates[i], dtype=wp.vec3i) for i in sorted(qLD_updates)) + + # qM_tiles records the block diagonal structure of qM + tile_corners = [i for i in range(mjm.nv) if mjm.dof_parentid[i] == -1] + tiles = {} + for i in range(len(tile_corners)): + tile_beg = tile_corners[i] + tile_end = mjm.nv if i == len(tile_corners) - 1 else tile_corners[i + 1] + tiles.setdefault(tile_end - tile_beg, []).append(tile_beg) + + qM_tiles = tuple(types.TileSet(adr=wp.array(tiles[sz], dtype=int), size=sz) for sz in sorted(tiles.keys())) + + # subtree_mass is a precalculated array used in smooth + subtree_mass = np.copy(mjm.body_mass) + # TODO(team): should this be [mjm.nbody - 1, 0) ? + for i in range(mjm.nbody - 1, -1, -1): + subtree_mass[mjm.body_parentid[i]] += subtree_mass[i] + + # actuator_moment tiles are grouped by dof size and number of actuators + tree_id = np.arange(len(tile_corners), dtype=np.int32) + num_trees = int(np.max(tree_id)) if len(tree_id) > 0 else 0 + bodyid = [] + for i in range(mjm.nu): + trntype = mjm.actuator_trntype[i] + if trntype == mujoco.mjtTrn.mjTRN_JOINT or trntype == mujoco.mjtTrn.mjTRN_JOINTINPARENT: + jntid = mjm.actuator_trnid[i, 0] + bodyid.append(mjm.jnt_bodyid[jntid]) + elif trntype == mujoco.mjtTrn.mjTRN_TENDON: + tenid = mjm.actuator_trnid[i, 0] + adr = mjm.tendon_adr[tenid] + if mjm.wrap_type[adr] == mujoco.mjtWrap.mjWRAP_JOINT: + ten_num = mjm.tendon_num[tenid] + for i in range(ten_num): + bodyid.append(mjm.jnt_bodyid[mjm.wrap_objid[adr + i]]) + else: + for i in range(mjm.nv): + bodyid.append(mjm.dof_bodyid[i]) + elif trntype == mujoco.mjtTrn.mjTRN_BODY: + pass + elif trntype == mujoco.mjtTrn.mjTRN_SITE: + siteid = mjm.actuator_trnid[i, 0] + bid = mjm.site_bodyid[siteid] + while bid > 0: + bodyid.append(bid) + bid = mjm.body_parentid[bid] + elif trntype == mujoco.mjtTrn.mjTRN_SLIDERCRANK: + for i in range(mjm.nv): + bodyid.append(mjm.dof_bodyid[i]) + else: + raise NotImplementedError(f"Transmission type {trntype} not implemented.") + tree = mjm.body_treeid[np.array(bodyid, dtype=int)] + counts, ids = np.histogram(tree, bins=np.arange(0, num_trees + 2)) + acts_per_tree = dict(zip(ids, counts)) + + tiles = {} + act_beg = 0 + for i in range(len(tile_corners)): + tile_beg = tile_corners[i] + tile_end = mjm.nv if i == len(tile_corners) - 1 else tile_corners[i + 1] + tree = int(tree_id[i]) + act_num = acts_per_tree[tree] + tiles.setdefault((tile_end - tile_beg, act_num), []).append((tile_beg, act_beg)) + act_beg += act_num + + actuator_moment_tiles_nv, actuator_moment_tiles_nu = tuple(), tuple() + + for (nv, nu), adr in sorted(tiles.items()): + adr_nv = wp.array([nv for nv, _ in adr], dtype=int) + adr_nu = wp.array([nu for _, nu in adr], dtype=int) + actuator_moment_tiles_nv += (types.TileSet(adr=adr_nv, size=nv),) + actuator_moment_tiles_nu += (types.TileSet(adr=adr_nu, size=nu),) + + # fixed tendon + tendon_jnt_adr = [] + wrap_jnt_adr = [] + for i in range(mjm.ntendon): + adr = mjm.tendon_adr[i] + if mjm.wrap_type[adr] == mujoco.mjtWrap.mjWRAP_JOINT: + tendon_num = mjm.tendon_num[i] + for j in range(tendon_num): + tendon_jnt_adr.append(i) + wrap_jnt_adr.append(adr + j) + + # spatial tendon + tendon_site_pair_adr = [] + tendon_geom_adr = [] + + ten_wrapadr_site = [0] + ten_wrapnum_site = [] + for i, tendon_num in enumerate(mjm.tendon_num): + adr = mjm.tendon_adr[i] + # sites + if (mjm.wrap_type[adr : adr + tendon_num] == mujoco.mjtWrap.mjWRAP_SITE).all(): + if i < mjm.ntendon: + ten_wrapadr_site.append(ten_wrapadr_site[-1] + tendon_num) + ten_wrapnum_site.append(tendon_num) + else: + if i < mjm.ntendon: + ten_wrapadr_site.append(ten_wrapadr_site[-1]) + ten_wrapnum_site.append(0) + + # geoms + for j in range(tendon_num): + wrap_type = mjm.wrap_type[adr + j] + if j < tendon_num - 1: + next_wrap_type = mjm.wrap_type[adr + j + 1] + if wrap_type == mujoco.mjtWrap.mjWRAP_SITE and next_wrap_type == mujoco.mjtWrap.mjWRAP_SITE: + tendon_site_pair_adr.append(i) + if wrap_type == mujoco.mjtWrap.mjWRAP_SPHERE or wrap_type == mujoco.mjtWrap.mjWRAP_CYLINDER: + tendon_geom_adr.append(i) + + wrap_site_adr = np.nonzero(mjm.wrap_type == mujoco.mjtWrap.mjWRAP_SITE)[0] + wrap_site_pair_adr = np.setdiff1d(wrap_site_adr[np.nonzero(np.diff(wrap_site_adr) == 1)[0]], mjm.tendon_adr[1:] - 1) + wrap_geom_adr = np.nonzero(np.isin(mjm.wrap_type, [mujoco.mjtWrap.mjWRAP_SPHERE, mujoco.mjtWrap.mjWRAP_CYLINDER]))[0] + + # pulley scaling + wrap_pulley_scale = np.ones(mjm.nwrap, dtype=float) + pulley_adr = np.nonzero(mjm.wrap_type == mujoco.mjtWrap.mjWRAP_PULLEY)[0] + for tadr, tnum in zip(mjm.tendon_adr, mjm.tendon_num): + for padr in pulley_adr: + if tadr <= padr < tadr + tnum: + wrap_pulley_scale[padr : tadr + tnum] = 1.0 / mjm.wrap_prm[padr] + + # mocap + mocap_bodyid = np.arange(mjm.nbody)[mjm.body_mocapid >= 0] + mocap_bodyid = mocap_bodyid[mjm.body_mocapid[mjm.body_mocapid >= 0].argsort()] + + # precalculated geom pairs + filterparent = not (mjm.opt.disableflags & types.DisableBit.FILTERPARENT.value) + + geom1, geom2 = np.triu_indices(mjm.ngeom, k=1) + nxn_geom_pair = np.stack((geom1, geom2), axis=1) + + bodyid1 = mjm.geom_bodyid[geom1] + bodyid2 = mjm.geom_bodyid[geom2] + contype1 = mjm.geom_contype[geom1] + contype2 = mjm.geom_contype[geom2] + conaffinity1 = mjm.geom_conaffinity[geom1] + conaffinity2 = mjm.geom_conaffinity[geom2] + weldid1 = mjm.body_weldid[bodyid1] + weldid2 = mjm.body_weldid[bodyid2] + weld_parentid1 = mjm.body_weldid[mjm.body_parentid[weldid1]] + weld_parentid2 = mjm.body_weldid[mjm.body_parentid[weldid2]] + + self_collision = weldid1 == weldid2 + parent_child_collision = ( + filterparent & (weldid1 != 0) & (weldid2 != 0) & ((weldid1 == weld_parentid2) | (weldid2 == weld_parentid1)) + ) + mask = np.array((contype1 & conaffinity2) | (contype2 & conaffinity1), dtype=bool) + exclude = np.isin((bodyid1 << 16) + bodyid2, mjm.exclude_signature) + + nxn_pairid = -1 * np.ones(len(geom1), dtype=int) + nxn_pairid[~(mask & ~self_collision & ~parent_child_collision & ~exclude)] = -2 + + # contact pairs + for i in range(mjm.npair): + pair_geom1 = mjm.pair_geom1[i] + pair_geom2 = mjm.pair_geom2[i] + + if pair_geom2 < pair_geom1: + pairid = np.int32(math.upper_tri_index(mjm.ngeom, int(pair_geom2), int(pair_geom1))) + else: + pairid = np.int32(math.upper_tri_index(mjm.ngeom, int(pair_geom1), int(pair_geom2))) + + nxn_pairid[pairid] = i + + include = nxn_pairid > -2 + nxn_pairid_filtered = nxn_pairid[include] + nxn_geom_pair_filtered = nxn_geom_pair[include] + + # count contact pair types + geom_type_pair_count = np.bincount( + [ + math.upper_trid_index(len(types.GeomType), int(mjm.geom_type[geom1[i]]), int(mjm.geom_type[geom2[i]])) + for i in np.arange(len(geom1)) + if nxn_pairid[i] > -2 + ], + minlength=len(types.GeomType) * (len(types.GeomType) + 1) // 2, + ) + + # Disable collisions if there are no potentially colliding pairs + if np.sum(geom_type_pair_count) == 0: + mjm.opt.disableflags |= types.DisableBit.CONTACT.value + + def create_nmodel_batched_array(mjm_array, dtype, expand_dim=True): + array = wp.array(mjm_array, dtype=dtype) + # add private attribute for JAX to determine which fields are batched + array._is_batched = True + if not expand_dim: + array.strides = (0,) + array.strides[1:] + return array + array.strides = (0,) + array.strides + array.ndim += 1 + array.shape = (1,) + array.shape + return array + + # rangefinder + is_rangefinder = mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER + sensor_rangefinder_adr = np.nonzero(is_rangefinder)[0] + rangefinder_sensor_adr = np.full(mjm.nsensor, -1) + rangefinder_sensor_adr[sensor_rangefinder_adr] = np.arange(len(sensor_rangefinder_adr)) + + # TODO(team): improve heuristic for selecting broadphase routine + if mjm.ngeom > 1000: + broadphase = types.BroadphaseType.SAP_SEGMENTED + elif mjm.ngeom > 100: + broadphase = types.BroadphaseType.SAP_TILE + else: + broadphase = types.BroadphaseType.NXN + + condim = np.concatenate((mjm.geom_condim, mjm.pair_dim)) + condim_max = np.max(condim) if len(condim) > 0 else 0 + + m = types.Model( + nq=mjm.nq, + nv=mjm.nv, + nu=mjm.nu, + na=mjm.na, + nbody=mjm.nbody, + njnt=mjm.njnt, + ngeom=mjm.ngeom, + nsite=mjm.nsite, + ncam=mjm.ncam, + nlight=mjm.nlight, + nflex=mjm.nflex, + nflexvert=mjm.nflexvert, + nflexedge=mjm.nflexedge, + nflexelem=mjm.nflexelem, + nflexelemdata=mjm.nflexelemdata, + nexclude=mjm.nexclude, + neq=mjm.neq, + nmocap=mjm.nmocap, + ngravcomp=mjm.ngravcomp, + nM=mjm.nM, + nC=mjm.nC, + ntendon=mjm.ntendon, + nwrap=mjm.nwrap, + nsensor=mjm.nsensor, + nsensordata=mjm.nsensordata, + nmeshvert=mjm.nmeshvert, + nmeshface=mjm.nmeshface, + nmeshgraph=mjm.nmeshgraph, + nmeshpoly=mjm.nmeshpoly, + nmeshpolyvert=mjm.nmeshpolyvert, + nmeshpolymap=mjm.nmeshpolymap, + nlsp=nlsp, + npair=mjm.npair, + opt=types.Option( + timestep=create_nmodel_batched_array(np.array(mjm.opt.timestep), dtype=float, expand_dim=False), + tolerance=create_nmodel_batched_array(np.array(mjm.opt.tolerance), dtype=float, expand_dim=False), + ls_tolerance=create_nmodel_batched_array(np.array(mjm.opt.ls_tolerance), dtype=float, expand_dim=False), + gravity=create_nmodel_batched_array(mjm.opt.gravity, dtype=wp.vec3, expand_dim=False), + magnetic=create_nmodel_batched_array(mjm.opt.magnetic, dtype=wp.vec3, expand_dim=False), + wind=create_nmodel_batched_array(mjm.opt.wind, dtype=wp.vec3, expand_dim=False), + has_fluid=bool(mjm.opt.wind.any() or mjm.opt.density or mjm.opt.viscosity), + density=create_nmodel_batched_array(np.array(mjm.opt.density), dtype=float, expand_dim=False), + viscosity=create_nmodel_batched_array(np.array(mjm.opt.viscosity), dtype=float, expand_dim=False), + cone=mjm.opt.cone, + solver=mjm.opt.solver, + iterations=mjm.opt.iterations, + ls_iterations=mjm.opt.ls_iterations, + integrator=mjm.opt.integrator, + disableflags=mjm.opt.disableflags, + enableflags=mjm.opt.enableflags, + impratio=create_nmodel_batched_array(np.array(mjm.opt.impratio), dtype=float, expand_dim=False), + is_sparse=bool(is_sparse), + ls_parallel=False, + gjk_iterations=MJ_CCD_ITERATIONS, + epa_iterations=MJ_CCD_ITERATIONS, + broadphase=int(broadphase), + broadphase_filter=int( + types.BroadphaseFilter.PLANE.value | types.BroadphaseFilter.SPHERE.value | types.BroadphaseFilter.OBB.value + ), + graph_conditional=True and conditional_graph_supported(), + sdf_initpoints=mjm.opt.sdf_initpoints, + sdf_iterations=mjm.opt.sdf_iterations, + run_collision_detection=True, + ), + stat=types.Statistic( + meaninertia=mjm.stat.meaninertia, + ), + qpos0=create_nmodel_batched_array(mjm.qpos0, dtype=float), + qpos_spring=create_nmodel_batched_array(mjm.qpos_spring, dtype=float), + qM_fullm_i=wp.array(qM_fullm_i, dtype=int), + qM_fullm_j=wp.array(qM_fullm_j, dtype=int), + qM_mulm_i=wp.array(qM_mulm_i, dtype=int), + qM_mulm_j=wp.array(qM_mulm_j, dtype=int), + qM_madr_ij=wp.array(qM_madr_ij, dtype=int), + qLD_updates=qLD_updates, + M_rownnz=wp.array(mjd.M_rownnz, dtype=int), + M_rowadr=wp.array(mjd.M_rowadr, dtype=int), + M_colind=wp.array(mjd.M_colind, dtype=int), + mapM2M=wp.array(mjd.mapM2M, dtype=int), + qM_tiles=qM_tiles, + body_tree=body_tree, + body_parentid=wp.array(mjm.body_parentid, dtype=int), + body_rootid=wp.array(mjm.body_rootid, dtype=int), + body_weldid=wp.array(mjm.body_weldid, dtype=int), + body_mocapid=wp.array(mjm.body_mocapid, dtype=int), + mocap_bodyid=wp.array(mocap_bodyid, dtype=int), + body_jntnum=wp.array(mjm.body_jntnum, dtype=int), + body_jntadr=wp.array(mjm.body_jntadr, dtype=int), + body_dofnum=wp.array(mjm.body_dofnum, dtype=int), + body_dofadr=wp.array(mjm.body_dofadr, dtype=int), + body_geomnum=wp.array(mjm.body_geomnum, dtype=int), + body_geomadr=wp.array(mjm.body_geomadr, dtype=int), + body_pos=create_nmodel_batched_array(mjm.body_pos, dtype=wp.vec3), + body_quat=create_nmodel_batched_array(mjm.body_quat, dtype=wp.quat), + body_ipos=create_nmodel_batched_array(mjm.body_ipos, dtype=wp.vec3), + body_iquat=create_nmodel_batched_array(mjm.body_iquat, dtype=wp.quat), + body_mass=create_nmodel_batched_array(mjm.body_mass, dtype=float), + body_subtreemass=create_nmodel_batched_array(mjm.body_subtreemass, dtype=float), + subtree_mass=create_nmodel_batched_array(subtree_mass, dtype=float), + body_inertia=create_nmodel_batched_array(mjm.body_inertia, dtype=wp.vec3), + body_invweight0=create_nmodel_batched_array(mjm.body_invweight0, dtype=wp.vec2), + body_contype=wp.array(mjm.body_contype, dtype=int), + body_conaffinity=wp.array(mjm.body_conaffinity, dtype=int), + body_gravcomp=create_nmodel_batched_array(mjm.body_gravcomp, dtype=float), + jnt_type=wp.array(mjm.jnt_type, dtype=int), + jnt_qposadr=wp.array(mjm.jnt_qposadr, dtype=int), + jnt_dofadr=wp.array(mjm.jnt_dofadr, dtype=int), + jnt_bodyid=wp.array(mjm.jnt_bodyid, dtype=int), + jnt_limited=wp.array(mjm.jnt_limited, dtype=int), + jnt_actfrclimited=wp.array(mjm.jnt_actfrclimited, dtype=bool), + jnt_solref=create_nmodel_batched_array(mjm.jnt_solref, dtype=wp.vec2), + jnt_solimp=create_nmodel_batched_array(mjm.jnt_solimp, dtype=types.vec5), + jnt_pos=create_nmodel_batched_array(mjm.jnt_pos, dtype=wp.vec3), + jnt_axis=create_nmodel_batched_array(mjm.jnt_axis, dtype=wp.vec3), + jnt_stiffness=create_nmodel_batched_array(mjm.jnt_stiffness, dtype=float), + jnt_range=create_nmodel_batched_array(mjm.jnt_range, dtype=wp.vec2), + jnt_actfrcrange=create_nmodel_batched_array(mjm.jnt_actfrcrange, dtype=wp.vec2), + jnt_margin=create_nmodel_batched_array(mjm.jnt_margin, dtype=float), + # these jnt_limited adrs are used in constraint.py + jnt_limited_slide_hinge_adr=wp.array( + np.nonzero( + mjm.jnt_limited & ((mjm.jnt_type == mujoco.mjtJoint.mjJNT_SLIDE) | (mjm.jnt_type == mujoco.mjtJoint.mjJNT_HINGE)) + )[0], + dtype=int, + ), + jnt_limited_ball_adr=wp.array( + np.nonzero(mjm.jnt_limited & (mjm.jnt_type == mujoco.mjtJoint.mjJNT_BALL))[0], + dtype=int, + ), + jnt_actgravcomp=wp.array(mjm.jnt_actgravcomp, dtype=int), + dof_bodyid=wp.array(mjm.dof_bodyid, dtype=int), + dof_jntid=wp.array(mjm.dof_jntid, dtype=int), + dof_parentid=wp.array(mjm.dof_parentid, dtype=int), + dof_Madr=wp.array(mjm.dof_Madr, dtype=int), + dof_armature=create_nmodel_batched_array(mjm.dof_armature, dtype=float), + dof_damping=create_nmodel_batched_array(mjm.dof_damping, dtype=float), + dof_invweight0=create_nmodel_batched_array(mjm.dof_invweight0, dtype=float), + dof_frictionloss=create_nmodel_batched_array(mjm.dof_frictionloss, dtype=float), + dof_solimp=create_nmodel_batched_array(mjm.dof_solimp, dtype=types.vec5), + dof_solref=create_nmodel_batched_array(mjm.dof_solref, dtype=wp.vec2), + dof_tri_row=wp.array(dof_tri_row, dtype=int), + dof_tri_col=wp.array(dof_tri_col, dtype=int), + geom_type=wp.array(mjm.geom_type, dtype=int), + geom_contype=wp.array(mjm.geom_contype, dtype=int), + geom_conaffinity=wp.array(mjm.geom_conaffinity, dtype=int), + geom_condim=wp.array(mjm.geom_condim, dtype=int), + geom_bodyid=wp.array(mjm.geom_bodyid, dtype=int), + geom_dataid=wp.array(mjm.geom_dataid, dtype=int), + geom_group=wp.array(mjm.geom_group, dtype=int), + geom_matid=create_nmodel_batched_array(mjm.geom_matid, dtype=int), + geom_priority=wp.array(mjm.geom_priority, dtype=int), + geom_solmix=create_nmodel_batched_array(mjm.geom_solmix, dtype=float), + geom_solref=create_nmodel_batched_array(mjm.geom_solref, dtype=wp.vec2), + geom_solimp=create_nmodel_batched_array(mjm.geom_solimp, dtype=types.vec5), + geom_size=create_nmodel_batched_array(mjm.geom_size, dtype=wp.vec3), + geom_aabb=wp.array2d(mjm.geom_aabb, dtype=wp.vec3), + geom_rbound=create_nmodel_batched_array(mjm.geom_rbound, dtype=float), + geom_pos=create_nmodel_batched_array(mjm.geom_pos, dtype=wp.vec3), + geom_quat=create_nmodel_batched_array(mjm.geom_quat, dtype=wp.quat), + geom_friction=create_nmodel_batched_array(mjm.geom_friction, dtype=wp.vec3), + geom_margin=create_nmodel_batched_array(mjm.geom_margin, dtype=float), + geom_gap=create_nmodel_batched_array(mjm.geom_gap, dtype=float), + geom_rgba=create_nmodel_batched_array(mjm.geom_rgba, dtype=wp.vec4), + site_type=wp.array(mjm.site_type, dtype=int), + site_bodyid=wp.array(mjm.site_bodyid, dtype=int), + site_size=wp.array(mjm.site_size, dtype=wp.vec3), + site_pos=create_nmodel_batched_array(mjm.site_pos, dtype=wp.vec3), + site_quat=create_nmodel_batched_array(mjm.site_quat, dtype=wp.quat), + cam_mode=wp.array(mjm.cam_mode, dtype=int), + cam_bodyid=wp.array(mjm.cam_bodyid, dtype=int), + cam_targetbodyid=wp.array(mjm.cam_targetbodyid, dtype=int), + cam_pos=create_nmodel_batched_array(mjm.cam_pos, dtype=wp.vec3), + cam_quat=create_nmodel_batched_array(mjm.cam_quat, dtype=wp.quat), + cam_poscom0=create_nmodel_batched_array(mjm.cam_poscom0, dtype=wp.vec3), + cam_pos0=create_nmodel_batched_array(mjm.cam_pos0, dtype=wp.vec3), + cam_mat0=create_nmodel_batched_array(mjm.cam_mat0, dtype=wp.mat33), + cam_fovy=wp.array(mjm.cam_fovy, dtype=float), + cam_resolution=wp.array(mjm.cam_resolution, dtype=wp.vec2i), + cam_sensorsize=wp.array(mjm.cam_sensorsize, dtype=wp.vec2), + cam_intrinsic=wp.array(mjm.cam_intrinsic, dtype=wp.vec4), + light_mode=wp.array(mjm.light_mode, dtype=int), + light_bodyid=wp.array(mjm.light_bodyid, dtype=int), + light_targetbodyid=wp.array(mjm.light_targetbodyid, dtype=int), + light_pos=create_nmodel_batched_array(mjm.light_pos, dtype=wp.vec3), + light_dir=create_nmodel_batched_array(mjm.light_dir, dtype=wp.vec3), + light_poscom0=create_nmodel_batched_array(mjm.light_poscom0, dtype=wp.vec3), + light_pos0=create_nmodel_batched_array(mjm.light_pos0, dtype=wp.vec3), + light_dir0=create_nmodel_batched_array(mjm.light_dir0, dtype=wp.vec3), + flex_dim=wp.array(mjm.flex_dim, dtype=int), + flex_vertadr=wp.array(mjm.flex_vertadr, dtype=int), + flex_vertnum=wp.array(mjm.flex_vertnum, dtype=int), + flex_edgeadr=wp.array(mjm.flex_edgeadr, dtype=int), + flex_elemedgeadr=wp.array(mjm.flex_elemedgeadr, dtype=int), + flex_vertbodyid=wp.array(mjm.flex_vertbodyid, dtype=int), + flex_edge=wp.array(mjm.flex_edge, dtype=wp.vec2i), + flex_edgeflap=wp.array(mjm.flex_edgeflap, dtype=wp.vec2i), + flex_elem=wp.array(mjm.flex_elem, dtype=int), + flex_elemedge=wp.array(mjm.flex_elemedge, dtype=int), + flexedge_length0=wp.array(mjm.flexedge_length0, dtype=float), + flex_stiffness=wp.array(mjm.flex_stiffness.flatten(), dtype=float), + flex_bending=wp.array(mjm.flex_bending, dtype=wp.mat44f), + flex_damping=wp.array(mjm.flex_damping, dtype=float), + mesh_vertadr=wp.array(mjm.mesh_vertadr, dtype=int), + mesh_vertnum=wp.array(mjm.mesh_vertnum, dtype=int), + mesh_vert=wp.array(mjm.mesh_vert, dtype=wp.vec3), + mesh_faceadr=wp.array(mjm.mesh_faceadr, dtype=int), + mesh_face=wp.array(mjm.mesh_face, dtype=wp.vec3i), + mesh_graphadr=wp.array(mjm.mesh_graphadr, dtype=int), + mesh_graph=wp.array(mjm.mesh_graph, dtype=int), + mesh_polynum=wp.array(mjm.mesh_polynum, dtype=int), + mesh_polyadr=wp.array(mjm.mesh_polyadr, dtype=int), + mesh_polynormal=wp.array(mjm.mesh_polynormal, dtype=wp.vec3), + mesh_polyvertadr=wp.array(mjm.mesh_polyvertadr, dtype=int), + mesh_polyvertnum=wp.array(mjm.mesh_polyvertnum, dtype=int), + mesh_polyvert=wp.array(mjm.mesh_polyvert, dtype=int), + mesh_polymapadr=wp.array(mjm.mesh_polymapadr, dtype=int), + mesh_polymapnum=wp.array(mjm.mesh_polymapnum, dtype=int), + mesh_polymap=wp.array(mjm.mesh_polymap, dtype=int), + nhfield=mjm.nhfield, + nhfielddata=mjm.nhfielddata, + hfield_adr=wp.array(mjm.hfield_adr, dtype=int), + hfield_nrow=wp.array(mjm.hfield_nrow, dtype=int), + hfield_ncol=wp.array(mjm.hfield_ncol, dtype=int), + hfield_size=wp.array(mjm.hfield_size, dtype=wp.vec4), + hfield_data=wp.array(mjm.hfield_data, dtype=float), + eq_type=wp.array(mjm.eq_type, dtype=int), + eq_obj1id=wp.array(mjm.eq_obj1id, dtype=int), + eq_obj2id=wp.array(mjm.eq_obj2id, dtype=int), + eq_objtype=wp.array(mjm.eq_objtype, dtype=int), + eq_active0=wp.array(mjm.eq_active0, dtype=bool), + eq_solref=create_nmodel_batched_array(mjm.eq_solref, dtype=wp.vec2), + eq_solimp=create_nmodel_batched_array(mjm.eq_solimp, dtype=types.vec5), + eq_data=create_nmodel_batched_array(mjm.eq_data, dtype=types.vec11), + # pre-compute indices of equality constraints + eq_connect_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.CONNECT.value)[0], dtype=int), + eq_wld_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.WELD.value)[0], dtype=int), + eq_jnt_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.JOINT.value)[0], dtype=int), + eq_ten_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.TENDON.value)[0], dtype=int), + actuator_moment_tiles_nv=actuator_moment_tiles_nv, + actuator_moment_tiles_nu=actuator_moment_tiles_nu, + actuator_trntype=wp.array(mjm.actuator_trntype, dtype=int), + actuator_dyntype=wp.array(mjm.actuator_dyntype, dtype=int), + actuator_gaintype=wp.array(mjm.actuator_gaintype, dtype=int), + actuator_biastype=wp.array(mjm.actuator_biastype, dtype=int), + actuator_trnid=wp.array(mjm.actuator_trnid, dtype=wp.vec2i), + actuator_actadr=wp.array(mjm.actuator_actadr, dtype=int), + actuator_actnum=wp.array(mjm.actuator_actnum, dtype=int), + actuator_ctrllimited=wp.array(mjm.actuator_ctrllimited, dtype=bool), + actuator_forcelimited=wp.array(mjm.actuator_forcelimited, dtype=bool), + actuator_actlimited=wp.array(mjm.actuator_actlimited, dtype=bool), + actuator_dynprm=create_nmodel_batched_array(mjm.actuator_dynprm, dtype=types.vec10f), + actuator_gainprm=create_nmodel_batched_array(mjm.actuator_gainprm, dtype=types.vec10f), + actuator_biasprm=create_nmodel_batched_array(mjm.actuator_biasprm, dtype=types.vec10f), + actuator_actearly=wp.array(mjm.actuator_actearly, dtype=bool), + actuator_ctrlrange=create_nmodel_batched_array(mjm.actuator_ctrlrange, dtype=wp.vec2), + actuator_forcerange=create_nmodel_batched_array(mjm.actuator_forcerange, dtype=wp.vec2), + actuator_actrange=create_nmodel_batched_array(mjm.actuator_actrange, dtype=wp.vec2), + actuator_gear=create_nmodel_batched_array(mjm.actuator_gear, dtype=wp.spatial_vector), + actuator_cranklength=wp.array(mjm.actuator_cranklength, dtype=float), + actuator_acc0=wp.array(mjm.actuator_acc0, dtype=float), + actuator_lengthrange=wp.array(mjm.actuator_lengthrange, dtype=wp.vec2), + exclude_signature=wp.array(mjm.exclude_signature, dtype=int), + # short-circuiting here allows us to skip a lot of code in implicit integration + actuator_affine_bias_gain=bool( + np.any(mjm.actuator_biastype == types.BiasType.AFFINE.value) + or np.any(mjm.actuator_gaintype == types.GainType.AFFINE.value) + ), + nxn_geom_pair=wp.array(nxn_geom_pair, dtype=wp.vec2i), + nxn_geom_pair_filtered=wp.array(nxn_geom_pair_filtered, dtype=wp.vec2i), + nxn_pairid=wp.array(nxn_pairid, dtype=int), + nxn_pairid_filtered=wp.array(nxn_pairid_filtered, dtype=int), + pair_dim=wp.array(mjm.pair_dim, dtype=int), + pair_geom1=wp.array(mjm.pair_geom1, dtype=int), + pair_geom2=wp.array(mjm.pair_geom2, dtype=int), + pair_solref=create_nmodel_batched_array(mjm.pair_solref, dtype=wp.vec2), + pair_solreffriction=create_nmodel_batched_array(mjm.pair_solreffriction, dtype=wp.vec2), + pair_solimp=create_nmodel_batched_array(mjm.pair_solimp, dtype=types.vec5), + pair_margin=create_nmodel_batched_array(mjm.pair_margin, dtype=float), + pair_gap=create_nmodel_batched_array(mjm.pair_gap, dtype=float), + pair_friction=create_nmodel_batched_array(mjm.pair_friction, dtype=types.vec5), + condim_max=condim_max, # TODO(team): get max after filtering, + tendon_adr=wp.array(mjm.tendon_adr, dtype=int), + tendon_num=wp.array(mjm.tendon_num, dtype=int), + tendon_limited=wp.array(mjm.tendon_limited, dtype=int), + tendon_limited_adr=wp.array(np.nonzero(mjm.tendon_limited)[0], dtype=int), + tendon_actfrclimited=wp.array(mjm.tendon_actfrclimited, dtype=bool), + tendon_solref_lim=create_nmodel_batched_array(mjm.tendon_solref_lim, dtype=wp.vec2f), + tendon_solimp_lim=create_nmodel_batched_array(mjm.tendon_solimp_lim, dtype=types.vec5), + tendon_solref_fri=create_nmodel_batched_array(mjm.tendon_solref_fri, dtype=wp.vec2f), + tendon_solimp_fri=create_nmodel_batched_array(mjm.tendon_solimp_fri, dtype=types.vec5), + tendon_range=create_nmodel_batched_array(mjm.tendon_range, dtype=wp.vec2f), + tendon_actfrcrange=create_nmodel_batched_array(mjm.tendon_actfrcrange, dtype=wp.vec2), + tendon_margin=create_nmodel_batched_array(mjm.tendon_margin, dtype=float), + tendon_stiffness=create_nmodel_batched_array(mjm.tendon_stiffness, dtype=float), + tendon_damping=create_nmodel_batched_array(mjm.tendon_damping, dtype=float), + tendon_armature=create_nmodel_batched_array(mjm.tendon_armature, dtype=float), + tendon_frictionloss=create_nmodel_batched_array(mjm.tendon_frictionloss, dtype=float), + tendon_lengthspring=create_nmodel_batched_array(mjm.tendon_lengthspring, dtype=wp.vec2), + tendon_length0=create_nmodel_batched_array(mjm.tendon_length0, dtype=float), + tendon_invweight0=create_nmodel_batched_array(mjm.tendon_invweight0, dtype=float), + wrap_objid=wp.array(mjm.wrap_objid, dtype=int), + wrap_prm=wp.array(mjm.wrap_prm, dtype=float), + wrap_type=wp.array(mjm.wrap_type, dtype=int), + tendon_jnt_adr=wp.array(tendon_jnt_adr, dtype=int), + tendon_site_pair_adr=wp.array(tendon_site_pair_adr, dtype=int), + tendon_geom_adr=wp.array(tendon_geom_adr, dtype=int), + ten_wrapadr_site=wp.array(ten_wrapadr_site, dtype=int), + ten_wrapnum_site=wp.array(ten_wrapnum_site, dtype=int), + wrap_jnt_adr=wp.array(wrap_jnt_adr, dtype=int), + wrap_site_adr=wp.array(wrap_site_adr, dtype=int), + wrap_site_pair_adr=wp.array(wrap_site_pair_adr, dtype=int), + wrap_geom_adr=wp.array(wrap_geom_adr, dtype=int), + wrap_pulley_scale=wp.array(wrap_pulley_scale, dtype=float), + sensor_type=wp.array(mjm.sensor_type, dtype=int), + sensor_datatype=wp.array(mjm.sensor_datatype, dtype=int), + sensor_objtype=wp.array(mjm.sensor_objtype, dtype=int), + sensor_objid=wp.array(mjm.sensor_objid, dtype=int), + sensor_reftype=wp.array(mjm.sensor_reftype, dtype=int), + sensor_refid=wp.array(mjm.sensor_refid, dtype=int), + sensor_dim=wp.array(mjm.sensor_dim, dtype=int), + sensor_adr=wp.array(mjm.sensor_adr, dtype=int), + sensor_cutoff=wp.array(mjm.sensor_cutoff, dtype=float), + sensor_pos_adr=wp.array( + np.nonzero( + (mjm.sensor_needstage == mujoco.mjtStage.mjSTAGE_POS) + & (mjm.sensor_type != mujoco.mjtSensor.mjSENS_JOINTLIMITPOS) + & (mjm.sensor_type != mujoco.mjtSensor.mjSENS_TENDONLIMITPOS) + )[0], + dtype=int, + ), + sensor_limitpos_adr=wp.array( + np.nonzero( + (mjm.sensor_type == mujoco.mjtSensor.mjSENS_JOINTLIMITPOS) | (mjm.sensor_type == mujoco.mjtSensor.mjSENS_TENDONLIMITPOS) + )[0], + dtype=int, + ), + sensor_vel_adr=wp.array( + np.nonzero( + (mjm.sensor_needstage == mujoco.mjtStage.mjSTAGE_VEL) + & ( + (mjm.sensor_type != mujoco.mjtSensor.mjSENS_JOINTLIMITVEL) + | (mjm.sensor_type != mujoco.mjtSensor.mjSENS_TENDONLIMITVEL) + ) + )[0], + dtype=int, + ), + sensor_limitvel_adr=wp.array( + np.nonzero( + (mjm.sensor_type == mujoco.mjtSensor.mjSENS_JOINTLIMITVEL) | (mjm.sensor_type == mujoco.mjtSensor.mjSENS_TENDONLIMITVEL) + )[0], + dtype=int, + ), + sensor_acc_adr=wp.array( + np.nonzero( + (mjm.sensor_needstage == mujoco.mjtStage.mjSTAGE_ACC) + & ( + (mjm.sensor_type != mujoco.mjtSensor.mjSENS_TOUCH) + | (mjm.sensor_type != mujoco.mjtSensor.mjSENS_JOINTLIMITFRC) + | (mjm.sensor_type != mujoco.mjtSensor.mjSENS_TENDONLIMITFRC) + | (mjm.sensor_type != mujoco.mjtSensor.mjSENS_TENDONACTFRC) + ) + )[0], + dtype=int, + ), + sensor_rangefinder_adr=wp.array(sensor_rangefinder_adr, dtype=int), + rangefinder_sensor_adr=wp.array(rangefinder_sensor_adr, dtype=int), + sensor_touch_adr=wp.array( + np.nonzero(mjm.sensor_type == mujoco.mjtSensor.mjSENS_TOUCH)[0], + dtype=int, + ), + sensor_limitfrc_adr=wp.array( + np.nonzero( + (mjm.sensor_type == mujoco.mjtSensor.mjSENS_JOINTLIMITFRC) | (mjm.sensor_type == mujoco.mjtSensor.mjSENS_TENDONLIMITFRC) + )[0], + dtype=int, + ), + sensor_e_potential=(mjm.sensor_type == mujoco.mjtSensor.mjSENS_E_POTENTIAL).any(), + sensor_e_kinetic=(mjm.sensor_type == mujoco.mjtSensor.mjSENS_E_KINETIC).any(), + sensor_tendonactfrc_adr=wp.array( + np.nonzero(mjm.sensor_type == mujoco.mjtSensor.mjSENS_TENDONACTFRC)[0], + dtype=int, + ), + sensor_subtree_vel=np.isin( + mjm.sensor_type, + [mujoco.mjtSensor.mjSENS_SUBTREELINVEL, mujoco.mjtSensor.mjSENS_SUBTREEANGMOM], + ).any(), + sensor_rne_postconstraint=np.isin( + mjm.sensor_type, + [ + mujoco.mjtSensor.mjSENS_ACCELEROMETER, + mujoco.mjtSensor.mjSENS_FORCE, + mujoco.mjtSensor.mjSENS_TORQUE, + mujoco.mjtSensor.mjSENS_FRAMELINACC, + mujoco.mjtSensor.mjSENS_FRAMEANGACC, + ], + ).any(), + sensor_rangefinder_bodyid=wp.array( + mjm.site_bodyid[mjm.sensor_objid[mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER]], dtype=int + ), + plugin=wp.array(plugin_id, dtype=int), + plugin_attr=wp.array(plugin_attr, dtype=wp.vec3f), + geom_plugin_index=wp.array(geom_plugin_index, dtype=int), + mat_rgba=create_nmodel_batched_array(mjm.mat_rgba, dtype=wp.vec4), + actuator_trntype_body_adr=wp.array(np.nonzero(mjm.actuator_trntype == mujoco.mjtTrn.mjTRN_BODY)[0], dtype=int), + geompair2hfgeompair=wp.array(_hfield_geom_pair(mjm)[1], dtype=int), + block_dim=types.BlockDim(), + geom_pair_type_count=tuple(geom_type_pair_count), + has_sdf_geom=bool(np.any(mjm.geom_type == mujoco.mjtGeom.mjGEOM_SDF)), + ) + + return m + + +def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: int = -1) -> types.Data: + """ + Creates a data object on device. + + Args: + mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host). + nworld (int, optional): Number of worlds. Defaults to 1. + nconmax (int, optional): Maximum number of contacts for all worlds. Defaults to -1. + njmax (int, optional): Maximum number of constraints for all worlds. Defaults to -1. + + Returns: + Data: The data object containing the current state and output arrays (device). + """ + # TODO(team): move to Model? + if nconmax == -1: + # TODO(team): heuristic for nconmax + nconmax = nworld * 20 + if njmax == -1: + # TODO(team): heuristic for njmax + njmax = 20 * 6 + condim = np.concatenate((mjm.geom_condim, mjm.pair_dim)) + condim_max = np.max(condim) if len(condim) > 0 else 0 + + if mujoco.mj_isSparse(mjm): + qM = wp.zeros((nworld, 1, mjm.nM), dtype=float) + qLD = wp.zeros((nworld, 1, mjm.nM), dtype=float) + else: + qM = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float) + qLD = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float) + + nrangefinder = sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER) + + return types.Data( + nworld=nworld, + nconmax=nconmax, + njmax=njmax, + solver_niter=wp.zeros(nworld, dtype=int), + ncon=wp.zeros(1, dtype=int), + ncon_hfield=wp.zeros((nworld, _hfield_geom_pair(mjm)[0]), dtype=int), # warp only + ne=wp.zeros(nworld, dtype=int), + ne_connect=wp.zeros(nworld, dtype=int), # warp only + ne_weld=wp.zeros(nworld, dtype=int), # warp only + ne_jnt=wp.zeros(nworld, dtype=int), # warp only + ne_ten=wp.zeros(nworld, dtype=int), # warp only + nf=wp.zeros(nworld, dtype=int), + nl=wp.zeros(nworld, dtype=int), + nefc=wp.zeros(nworld, dtype=int), + nsolving=wp.zeros(1, dtype=int), # warp only + time=wp.zeros(nworld, dtype=float), + energy=wp.zeros(nworld, dtype=wp.vec2), + qpos=wp.zeros((nworld, mjm.nq), dtype=float), + qvel=wp.zeros((nworld, mjm.nv), dtype=float), + act=wp.zeros((nworld, mjm.na), dtype=float), + qacc_warmstart=wp.zeros((nworld, mjm.nv), dtype=float), + qacc_discrete=wp.zeros((nworld, mjm.nv), dtype=float), + ctrl=wp.zeros((nworld, mjm.nu), dtype=float), + qfrc_applied=wp.zeros((nworld, mjm.nv), dtype=float), + xfrc_applied=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), + fluid_applied=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), + eq_active=wp.array(np.tile(mjm.eq_active0, (nworld, 1)), dtype=bool), + mocap_pos=wp.zeros((nworld, mjm.nmocap), dtype=wp.vec3), + mocap_quat=wp.zeros((nworld, mjm.nmocap), dtype=wp.quat), + qacc=wp.zeros((nworld, mjm.nv), dtype=float), + act_dot=wp.zeros((nworld, mjm.na), dtype=float), + xpos=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3), + xquat=wp.zeros((nworld, mjm.nbody), dtype=wp.quat), + xmat=wp.zeros((nworld, mjm.nbody), dtype=wp.mat33), + xipos=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3), + ximat=wp.zeros((nworld, mjm.nbody), dtype=wp.mat33), + xanchor=wp.zeros((nworld, mjm.njnt), dtype=wp.vec3), + xaxis=wp.zeros((nworld, mjm.njnt), dtype=wp.vec3), + geom_skip=wp.zeros(mjm.ngeom, dtype=bool), # warp only + geom_xpos=wp.zeros((nworld, mjm.ngeom), dtype=wp.vec3), + geom_xmat=wp.zeros((nworld, mjm.ngeom), dtype=wp.mat33), + site_xpos=wp.zeros((nworld, mjm.nsite), dtype=wp.vec3), + site_xmat=wp.zeros((nworld, mjm.nsite), dtype=wp.mat33), + cam_xpos=wp.zeros((nworld, mjm.ncam), dtype=wp.vec3), + cam_xmat=wp.zeros((nworld, mjm.ncam), dtype=wp.mat33), + light_xpos=wp.zeros((nworld, mjm.nlight), dtype=wp.vec3), + light_xdir=wp.zeros((nworld, mjm.nlight), dtype=wp.vec3), + subtree_com=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3), + cdof=wp.zeros((nworld, mjm.nv), dtype=wp.spatial_vector), + cinert=wp.zeros((nworld, mjm.nbody), dtype=types.vec10), + flexvert_xpos=wp.zeros((nworld, mjm.nflexvert), dtype=wp.vec3), + flexedge_length=wp.zeros((nworld, mjm.nflexedge), dtype=wp.float32), + flexedge_velocity=wp.zeros((nworld, mjm.nflexedge), dtype=wp.float32), + actuator_length=wp.zeros((nworld, mjm.nu), dtype=float), + actuator_moment=wp.zeros((nworld, mjm.nu, mjm.nv), dtype=float), + crb=wp.zeros((nworld, mjm.nbody), dtype=types.vec10), + qM=qM, + qLD=qLD, + qLDiagInv=wp.zeros((nworld, mjm.nv), dtype=float), + ten_velocity=wp.zeros((nworld, mjm.ntendon), dtype=float), + actuator_velocity=wp.zeros((nworld, mjm.nu), dtype=float), + cvel=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), + cdof_dot=wp.zeros((nworld, mjm.nv), dtype=wp.spatial_vector), + qfrc_bias=wp.zeros((nworld, mjm.nv), dtype=float), + qfrc_spring=wp.zeros((nworld, mjm.nv), dtype=float), + qfrc_damper=wp.zeros((nworld, mjm.nv), dtype=float), + qfrc_gravcomp=wp.zeros((nworld, mjm.nv), dtype=float), + qfrc_fluid=wp.zeros((nworld, mjm.nv), dtype=float), + qfrc_passive=wp.zeros((nworld, mjm.nv), dtype=float), + subtree_linvel=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3), + subtree_angmom=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3), + subtree_bodyvel=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), # warp only + actuator_force=wp.zeros((nworld, mjm.nu), dtype=float), + qfrc_actuator=wp.zeros((nworld, mjm.nv), dtype=float), + qfrc_smooth=wp.zeros((nworld, mjm.nv), dtype=float), + qacc_smooth=wp.zeros((nworld, mjm.nv), dtype=float), + qfrc_constraint=wp.zeros((nworld, mjm.nv), dtype=float), + qfrc_inverse=wp.zeros((nworld, mjm.nv), dtype=float), + contact=types.Contact( + dist=wp.zeros((nconmax,), dtype=float), + pos=wp.zeros((nconmax,), dtype=wp.vec3f), + frame=wp.zeros((nconmax,), dtype=wp.mat33f), + includemargin=wp.zeros((nconmax,), dtype=float), + friction=wp.zeros((nconmax,), dtype=types.vec5), + solref=wp.zeros((nconmax,), dtype=wp.vec2f), + solreffriction=wp.zeros((nconmax,), dtype=wp.vec2f), + solimp=wp.zeros((nconmax,), dtype=types.vec5), + dim=wp.zeros((nconmax,), dtype=int), + geom=wp.zeros((nconmax,), dtype=wp.vec2i), + efc_address=wp.zeros( + (nconmax, np.maximum(1, 2 * (condim_max - 1))), + dtype=int, + ), + worldid=wp.zeros((nconmax,), dtype=int), + ), + efc=types.Constraint( + type=wp.zeros((nworld, njmax), dtype=int), + id=wp.zeros((nworld, njmax), dtype=int), + J=wp.zeros((nworld, njmax, mjm.nv), dtype=float), + pos=wp.zeros((nworld, njmax), dtype=float), + margin=wp.zeros((nworld, njmax), dtype=float), + D=wp.zeros((nworld, njmax), dtype=float), + vel=wp.zeros((nworld, njmax), dtype=float), + aref=wp.zeros((nworld, njmax), dtype=float), + frictionloss=wp.zeros((nworld, njmax), dtype=float), + force=wp.zeros((nworld, njmax), dtype=float), + Jaref=wp.zeros((nworld, njmax), dtype=float), + Ma=wp.zeros((nworld, mjm.nv), dtype=float), + grad=wp.zeros((nworld, mjm.nv), dtype=float), + cholesky_L_tmp=wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float), + cholesky_y_tmp=wp.zeros((nworld, mjm.nv), dtype=float), + grad_dot=wp.zeros((nworld,), dtype=float), + Mgrad=wp.zeros((nworld, mjm.nv), dtype=float), + search=wp.zeros((nworld, mjm.nv), dtype=float), + search_dot=wp.zeros((nworld,), dtype=float), + gauss=wp.zeros((nworld,), dtype=float), + cost=wp.zeros((nworld,), dtype=float), + prev_cost=wp.zeros((nworld,), dtype=float), + active=wp.zeros((nworld, njmax), dtype=bool), + gtol=wp.zeros((nworld,), dtype=float), + mv=wp.zeros((nworld, mjm.nv), dtype=float), + jv=wp.zeros((nworld, njmax), dtype=float), + quad=wp.zeros((nworld, njmax), dtype=wp.vec3f), + quad_gauss=wp.zeros((nworld,), dtype=wp.vec3f), + h=wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float), + alpha=wp.zeros((nworld,), dtype=float), + prev_grad=wp.zeros((nworld, mjm.nv), dtype=float), + prev_Mgrad=wp.zeros((nworld, mjm.nv), dtype=float), + beta=wp.zeros((nworld,), dtype=float), + beta_num=wp.zeros((nworld,), dtype=float), + beta_den=wp.zeros((nworld,), dtype=float), + done=wp.zeros((nworld,), dtype=bool), + # linesearch + ls_done=wp.zeros((nworld,), dtype=bool), + p0=wp.zeros((nworld,), dtype=wp.vec3), + lo=wp.zeros((nworld,), dtype=wp.vec3), + lo_alpha=wp.zeros((nworld,), dtype=float), + hi=wp.zeros((nworld,), dtype=wp.vec3), + hi_alpha=wp.zeros((nworld,), dtype=float), + lo_next=wp.zeros((nworld,), dtype=wp.vec3), + lo_next_alpha=wp.zeros((nworld,), dtype=float), + hi_next=wp.zeros((nworld,), dtype=wp.vec3), + hi_next_alpha=wp.zeros((nworld,), dtype=float), + mid=wp.zeros((nworld,), dtype=wp.vec3), + mid_alpha=wp.zeros((nworld,), dtype=float), + cost_candidate=wp.zeros((nworld, mjm.opt.ls_iterations), dtype=float), + # elliptic cone + u=wp.zeros((nconmax,), dtype=types.vec6), + uu=wp.zeros((nconmax,), dtype=float), + uv=wp.zeros((nconmax,), dtype=float), + vv=wp.zeros((nconmax,), dtype=float), + condim=wp.zeros((nworld, njmax), dtype=int), + ), + # RK4 + qpos_t0=wp.zeros((nworld, mjm.nq), dtype=float), + qvel_t0=wp.zeros((nworld, mjm.nv), dtype=float), + act_t0=wp.zeros((nworld, mjm.na), dtype=float), + qvel_rk=wp.zeros((nworld, mjm.nv), dtype=float), + qacc_rk=wp.zeros((nworld, mjm.nv), dtype=float), + act_dot_rk=wp.zeros((nworld, mjm.na), dtype=float), + # euler + implicit integration + qfrc_integration=wp.zeros((nworld, mjm.nv), dtype=float), + qacc_integration=wp.zeros((nworld, mjm.nv), dtype=float), + act_vel_integration=wp.zeros((nworld, mjm.nu), dtype=float), + qM_integration=wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float), + qLD_integration=wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float), + qLDiagInv_integration=wp.zeros((nworld, mjm.nv), dtype=float), + # sweep-and-prune broadphase + sap_projection_lower=wp.zeros((nworld, mjm.ngeom, 2), dtype=float), + sap_projection_upper=wp.zeros((nworld, mjm.ngeom), dtype=float), + sap_sort_index=wp.zeros((nworld, mjm.ngeom, 2), dtype=int), + sap_range=wp.zeros((nworld, mjm.ngeom), dtype=int), + sap_cumulative_sum=wp.zeros((nworld, mjm.ngeom), dtype=int), + sap_segment_index=wp.array( + np.array([i * mjm.ngeom if i < nworld + 1 else 0 for i in range(2 * nworld)]).reshape((nworld, 2)), dtype=int + ), + # collision driver + collision_pair=wp.zeros((nconmax,), dtype=wp.vec2i), + collision_hftri_index=wp.zeros((nconmax,), dtype=int), + collision_pairid=wp.zeros((nconmax,), dtype=int), + collision_worldid=wp.zeros((nconmax,), dtype=int), + ncollision=wp.zeros((1,), dtype=int), + # narrowphase (EPA polytope) + epa_vert=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=wp.vec3), + epa_vert1=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=wp.vec3), + epa_vert2=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=wp.vec3), + epa_vert_index1=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=int), + epa_vert_index2=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=int), + epa_face=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=wp.vec3i), + epa_pr=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=wp.vec3), + epa_norm2=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=float), + epa_index=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=int), + epa_map=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=int), + epa_horizon=wp.zeros(shape=(nconmax, 6 * MJ_CCD_ITERATIONS), dtype=int), + # rne_postconstraint + cacc=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), + cfrc_int=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), + cfrc_ext=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), + # tendon + ten_length=wp.zeros((nworld, mjm.ntendon), dtype=float), + ten_J=wp.zeros((nworld, mjm.ntendon, mjm.nv), dtype=float), + ten_Jdot=wp.zeros((nworld, mjm.ntendon, mjm.nv), dtype=float), + ten_bias_coef=wp.zeros((nworld, mjm.ntendon), dtype=float), + ten_wrapadr=wp.zeros((nworld, mjm.ntendon), dtype=int), + ten_wrapnum=wp.zeros((nworld, mjm.ntendon), dtype=int), + ten_actfrc=wp.zeros((nworld, mjm.ntendon), dtype=float), + wrap_obj=wp.zeros((nworld, mjm.nwrap), dtype=wp.vec2i), + wrap_xpos=wp.zeros((nworld, mjm.nwrap), dtype=wp.spatial_vector), + wrap_geom_xpos=wp.zeros((nworld, mjm.nwrap), dtype=wp.spatial_vector), + # sensors + sensordata=wp.zeros((nworld, mjm.nsensordata), dtype=float), + sensor_rangefinder_pnt=wp.zeros((nworld, nrangefinder), dtype=wp.vec3), + sensor_rangefinder_vec=wp.zeros((nworld, nrangefinder), dtype=wp.vec3), + sensor_rangefinder_dist=wp.zeros((nworld, nrangefinder), dtype=float), + sensor_rangefinder_geomid=wp.zeros((nworld, nrangefinder), dtype=int), + # ray + ray_bodyexclude=wp.zeros(1, dtype=int), + ray_dist=wp.zeros((nworld, 1), dtype=float), + ray_geomid=wp.zeros((nworld, 1), dtype=int), + # mul_m + energy_vel_mul_m_skip=wp.zeros((nworld,), dtype=bool), + inverse_mul_m_skip=wp.zeros((nworld,), dtype=bool), + # actuator + actuator_trntype_body_ncon=wp.zeros((nworld, np.sum(mjm.actuator_trntype == mujoco.mjtTrn.mjTRN_BODY)), dtype=int), + ) + + +def put_data( + mjm: mujoco.MjModel, + mjd: mujoco.MjData, + nworld: Optional[int] = None, + nconmax: Optional[int] = None, + njmax: Optional[int] = None, +) -> types.Data: + """ + Moves data from host to a device. + + Args: + mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host). + mjd (mujoco.MjData): The data object containing current state and output arrays (host). + nworld (int, optional): The number of worlds. Defaults to 1. + nconmax (int, optional): The maximum number of contacts for all worlds. Defaults to -1. + njmax (int, optional): The maximum number of constraints for all worlds. Defaults to -1. + + Returns: + Data: The data object containing the current state and output arrays (device). + """ + # TODO(team): move nconmax and njmax to Model? + # TODO(team): decide what to do about uninitialized warp-only fields created by put_data + # we need to ensure these are only workspace fields and don't carry state + + nworld = nworld or 1 + # TODO(team): better heuristic for nconmax + nconmax = nconmax or max(512, mjd.ncon * nworld) + # TODO(team): better heuristic for njmax + njmax = njmax or max(5, mjd.nefc) + + if nworld < 1: + raise ValueError("nworld must be >= 1") + + if nconmax < 1: + raise ValueError("nconmax must be >= 1") + + if njmax < 1: + raise ValueError("njmax must be >= 1") + + if nworld * mjd.ncon > nconmax: + raise ValueError(f"nconmax overflow (nconmax must be >= {nworld * mjd.ncon})") + + if mjd.nefc > njmax: + raise ValueError(f"njmax overflow (njmax must be >= {mjd.nefc})") + + # calculate some fields that cannot be easily computed inline: + if mujoco.mj_isSparse(mjm): + qM = np.expand_dims(mjd.qM, axis=0) + qLD = np.expand_dims(mjd.qLD, axis=0) + qM_integration = np.zeros((1, mjm.nM), dtype=float) + qLD_integration = np.zeros((1, mjm.nM), dtype=float) + efc_J = np.zeros((mjd.nefc, mjm.nv)) + mujoco.mju_sparse2dense(efc_J, mjd.efc_J, mjd.efc_J_rownnz, mjd.efc_J_rowadr, mjd.efc_J_colind) + ten_J = np.zeros((mjm.ntendon, mjm.nv)) + mujoco.mju_sparse2dense( + ten_J, + mjd.ten_J.reshape(-1), + mjd.ten_J_rownnz, + mjd.ten_J_rowadr, + mjd.ten_J_colind.reshape(-1), + ) + else: + qM = np.zeros((mjm.nv, mjm.nv)) + mujoco.mj_fullM(mjm, qM, mjd.qM) + if (mjd.qM == 0.0).all() or (mjd.qLD == 0.0).all(): + qLD = np.zeros((mjm.nv, mjm.nv)) + else: + qLD = np.linalg.cholesky(qM) + qM_integration = np.zeros((mjm.nv, mjm.nv), dtype=float) + qLD_integration = np.zeros((mjm.nv, mjm.nv), dtype=float) + efc_J = mjd.efc_J.reshape((mjd.nefc, mjm.nv)) + ten_J = mjd.ten_J.reshape((mjm.ntendon, mjm.nv)) + + # TODO(taylorhowell): sparse actuator_moment + actuator_moment = np.zeros((mjm.nu, mjm.nv)) + mujoco.mju_sparse2dense( + actuator_moment, + mjd.actuator_moment, + mjd.moment_rownnz, + mjd.moment_rowadr, + mjd.moment_colind, + ) + + condim = np.concatenate((mjm.geom_condim, mjm.pair_dim)) + condim_max = np.max(condim) if len(condim) > 0 else 0 + contact_efc_address = np.zeros((nconmax, np.maximum(1, 2 * (condim_max - 1))), dtype=int) + for i in range(nworld): + for j in range(mjd.ncon): + condim = mjd.contact.dim[j] + efc_address = mjd.contact.efc_address[j] + if efc_address == -1: + continue + if condim == 1: + nconvar = 1 + else: + nconvar = condim if mjm.opt.cone == mujoco.mjtCone.mjCONE_ELLIPTIC else 2 * (condim - 1) + for k in range(nconvar): + contact_efc_address[i * mjd.ncon + j, k] = mjd.nefc * i + efc_address + k + + contact_worldid = np.pad(np.repeat(np.arange(nworld), mjd.ncon), (0, nconmax - nworld * mjd.ncon)) + + ne_connect = int(3 * np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_CONNECT) & mjd.eq_active)) + ne_weld = int(6 * np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_WELD) & mjd.eq_active)) + ne_jnt = int(np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_JOINT) & mjd.eq_active)) + ne_ten = int(np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_TENDON) & mjd.eq_active)) + + efc_type_fill = np.zeros((nworld, njmax)) + efc_id_fill = np.zeros((nworld, njmax)) + efc_J_fill = np.zeros((nworld, njmax, mjm.nv)) + efc_D_fill = np.zeros((nworld, njmax)) + efc_vel_fill = np.zeros((nworld, njmax)) + efc_pos_fill = np.zeros((nworld, njmax)) + efc_aref_fill = np.zeros((nworld, njmax)) + efc_frictionloss_fill = np.zeros((nworld, njmax)) + efc_force_fill = np.zeros((nworld, njmax)) + efc_margin_fill = np.zeros((nworld, njmax)) + + nefc = mjd.nefc + efc_type_fill[:, :nefc] = np.tile(mjd.efc_type, (nworld, 1)) + efc_id_fill[:, :nefc] = np.tile(mjd.efc_id, (nworld, 1)) + efc_J_fill[:, :nefc, :] = np.tile(efc_J, (nworld, 1, 1)) + efc_D_fill[:, :nefc] = np.tile(mjd.efc_D, (nworld, 1)) + efc_vel_fill[:, :nefc] = np.tile(mjd.efc_vel, (nworld, 1)) + efc_pos_fill[:, :nefc] = np.tile(mjd.efc_pos, (nworld, 1)) + efc_aref_fill[:, :nefc] = np.tile(mjd.efc_aref, (nworld, 1)) + efc_frictionloss_fill[:, :nefc] = np.tile(mjd.efc_frictionloss, (nworld, 1)) + efc_force_fill[:, :nefc] = np.tile(mjd.efc_force, (nworld, 1)) + efc_margin_fill[:, :nefc] = np.tile(mjd.efc_margin, (nworld, 1)) + + nrangefinder = sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER) + + # some helper functions to simplify the data field definitions below + + def arr(x, dtype=None): + if not isinstance(x, np.ndarray): + x = np.array(x) + if dtype is None: + if np.issubdtype(x.dtype, np.integer): + dtype = wp.int32 + elif np.issubdtype(x.dtype, np.floating): + dtype = wp.float32 + elif np.issubdtype(x.dtype, bool): + dtype = wp.bool + else: + raise ValueError(f"Unsupported dtype: {x.dtype}") + wp_array = {1: wp.array, 2: wp.array2d, 3: wp.array3d}[x.ndim] + return wp_array(x, dtype=dtype) + + def tile(x, dtype=None): + return arr(np.tile(x, (nworld,) + (1,) * len(x.shape)), dtype) + + def padtile(x, length, dtype=None): + x = np.repeat(x, nworld, axis=0) + width = ((0, length - x.shape[0]),) + ((0, 0),) * (x.ndim - 1) + return arr(np.pad(x, width), dtype) + + return types.Data( + nworld=nworld, + nconmax=nconmax, + njmax=njmax, + solver_niter=tile(mjd.solver_niter[0]), + ncon=arr([mjd.ncon * nworld]), + ncon_hfield=wp.zeros((nworld, _hfield_geom_pair(mjm)[0]), dtype=int), # warp only + ne=wp.full(shape=(nworld), value=mjd.ne), + ne_connect=wp.full(shape=(nworld), value=ne_connect), + ne_weld=wp.full(shape=(nworld), value=ne_weld), + ne_jnt=wp.full(shape=(nworld), value=ne_jnt), + ne_ten=wp.full(shape=(nworld), value=ne_ten), + nf=wp.full(shape=(nworld), value=mjd.nf), + nl=wp.full(shape=(nworld), value=mjd.nl), + nefc=wp.full(shape=(nworld), value=mjd.nefc), + nsolving=arr([nworld]), + time=arr(mjd.time * np.ones(nworld)), + energy=tile(mjd.energy, dtype=wp.vec2), + qpos=tile(mjd.qpos), + qvel=tile(mjd.qvel), + act=tile(mjd.act), + qacc_warmstart=tile(mjd.qacc_warmstart), + qacc_discrete=wp.zeros((nworld, mjm.nv), dtype=float), + ctrl=tile(mjd.ctrl), + qfrc_applied=tile(mjd.qfrc_applied), + xfrc_applied=tile(mjd.xfrc_applied, dtype=wp.spatial_vector), + fluid_applied=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), + eq_active=tile(mjd.eq_active.astype(bool)), + mocap_pos=tile(mjd.mocap_pos, dtype=wp.vec3), + mocap_quat=tile(mjd.mocap_quat, dtype=wp.quat), + qacc=tile(mjd.qacc), + act_dot=tile(mjd.act_dot), + xpos=tile(mjd.xpos, dtype=wp.vec3), + xquat=tile(mjd.xquat, dtype=wp.quat), + xmat=tile(mjd.xmat, dtype=wp.mat33), + xipos=tile(mjd.xipos, dtype=wp.vec3), + ximat=tile(mjd.ximat, dtype=wp.mat33), + xanchor=tile(mjd.xanchor, dtype=wp.vec3), + xaxis=tile(mjd.xaxis, dtype=wp.vec3), + geom_skip=wp.zeros(mjm.ngeom, dtype=bool), # warp only + geom_xpos=tile(mjd.geom_xpos, dtype=wp.vec3), + geom_xmat=tile(mjd.geom_xmat, dtype=wp.mat33), + site_xpos=tile(mjd.site_xpos, dtype=wp.vec3), + site_xmat=tile(mjd.site_xmat, dtype=wp.mat33), + cam_xpos=tile(mjd.cam_xpos, dtype=wp.vec3), + cam_xmat=tile(mjd.cam_xmat, dtype=wp.mat33), + light_xpos=tile(mjd.light_xpos, dtype=wp.vec3), + light_xdir=tile(mjd.light_xdir, dtype=wp.vec3), + subtree_com=tile(mjd.subtree_com, dtype=wp.vec3), + cdof=tile(mjd.cdof, dtype=wp.spatial_vector), + cinert=tile(mjd.cinert, dtype=types.vec10), + flexvert_xpos=tile(mjd.flexvert_xpos, dtype=wp.vec3), + flexedge_length=tile(mjd.flexedge_length), + flexedge_velocity=tile(mjd.flexedge_velocity), + actuator_length=tile(mjd.actuator_length), + actuator_moment=tile(actuator_moment), + crb=tile(mjd.crb, dtype=types.vec10), + qM=tile(qM), + qLD=tile(qLD), + qLDiagInv=tile(mjd.qLDiagInv), + ten_velocity=tile(mjd.ten_velocity), + actuator_velocity=tile(mjd.actuator_velocity), + cvel=tile(mjd.cvel, dtype=wp.spatial_vector), + cdof_dot=tile(mjd.cdof_dot, dtype=wp.spatial_vector), + qfrc_bias=tile(mjd.qfrc_bias), + qfrc_spring=tile(mjd.qfrc_spring), + qfrc_damper=tile(mjd.qfrc_damper), + qfrc_gravcomp=tile(mjd.qfrc_gravcomp), + qfrc_fluid=tile(mjd.qfrc_fluid), + qfrc_passive=tile(mjd.qfrc_passive), + subtree_linvel=tile(mjd.subtree_linvel, dtype=wp.vec3), + subtree_angmom=tile(mjd.subtree_angmom, dtype=wp.vec3), + subtree_bodyvel=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), + actuator_force=tile(mjd.actuator_force), + qfrc_actuator=tile(mjd.qfrc_actuator), + qfrc_smooth=tile(mjd.qfrc_smooth), + qacc_smooth=tile(mjd.qacc_smooth), + qfrc_constraint=tile(mjd.qfrc_constraint), + qfrc_inverse=tile(mjd.qfrc_inverse), + contact=types.Contact( + dist=padtile(mjd.contact.dist, nconmax), + pos=padtile(mjd.contact.pos, nconmax, dtype=wp.vec3), + frame=padtile(mjd.contact.frame, nconmax, dtype=wp.mat33), + includemargin=padtile(mjd.contact.includemargin, nconmax), + friction=padtile(mjd.contact.friction, nconmax, dtype=types.vec5), + solref=padtile(mjd.contact.solref, nconmax, dtype=wp.vec2f), + solreffriction=padtile(mjd.contact.solreffriction, nconmax, dtype=wp.vec2f), + solimp=padtile(mjd.contact.solimp, nconmax, dtype=types.vec5), + dim=padtile(mjd.contact.dim, nconmax), + geom=padtile(mjd.contact.geom, nconmax, dtype=wp.vec2i), + efc_address=arr(contact_efc_address), + worldid=arr(contact_worldid), + ), + efc=types.Constraint( + type=wp.array2d(efc_type_fill, dtype=int), + id=wp.array2d(efc_id_fill, dtype=int), + J=wp.array3d(efc_J_fill, dtype=float), + pos=wp.array2d(efc_pos_fill, dtype=float), + margin=wp.array2d(efc_margin_fill, dtype=float), + D=wp.array2d(efc_D_fill, dtype=float), + vel=wp.array2d(efc_vel_fill, dtype=float), + aref=wp.array2d(efc_aref_fill, dtype=float), + frictionloss=wp.array2d(efc_frictionloss_fill, dtype=float), + force=wp.array2d(efc_force_fill, dtype=float), + Jaref=wp.empty(shape=(nworld, njmax), dtype=float), + Ma=wp.empty(shape=(nworld, mjm.nv), dtype=float), + grad=wp.empty(shape=(nworld, mjm.nv), dtype=float), + cholesky_L_tmp=wp.empty(shape=(nworld, mjm.nv, mjm.nv), dtype=float), + cholesky_y_tmp=wp.empty(shape=(nworld, mjm.nv), dtype=float), + grad_dot=wp.empty(shape=(nworld,), dtype=float), + Mgrad=wp.empty(shape=(nworld, mjm.nv), dtype=float), + search=wp.empty(shape=(nworld, mjm.nv), dtype=float), + search_dot=wp.empty(shape=(nworld,), dtype=float), + gauss=wp.empty(shape=(nworld,), dtype=float), + cost=wp.empty(shape=(nworld,), dtype=float), + prev_cost=wp.empty(shape=(nworld,), dtype=float), + active=wp.empty(shape=(nworld, njmax), dtype=bool), + gtol=wp.empty(shape=(nworld,), dtype=float), + mv=wp.empty(shape=(nworld, mjm.nv), dtype=float), + jv=wp.empty(shape=(nworld, njmax), dtype=float), + quad=wp.empty(shape=(nworld, njmax), dtype=wp.vec3f), + quad_gauss=wp.empty(shape=(nworld,), dtype=wp.vec3f), + h=wp.empty(shape=(nworld, mjm.nv, mjm.nv), dtype=float), + alpha=wp.empty(shape=(nworld,), dtype=float), + prev_grad=wp.empty(shape=(nworld, mjm.nv), dtype=float), + prev_Mgrad=wp.empty(shape=(nworld, mjm.nv), dtype=float), + beta=wp.empty(shape=(nworld,), dtype=float), + beta_num=wp.empty(shape=(nworld,), dtype=float), + beta_den=wp.empty(shape=(nworld,), dtype=float), + done=wp.empty(shape=(nworld,), dtype=bool), + ls_done=wp.zeros(shape=(nworld,), dtype=bool), + p0=wp.empty(shape=(nworld,), dtype=wp.vec3), + lo=wp.empty(shape=(nworld,), dtype=wp.vec3), + lo_alpha=wp.empty(shape=(nworld,), dtype=float), + hi=wp.empty(shape=(nworld,), dtype=wp.vec3), + hi_alpha=wp.empty(shape=(nworld,), dtype=float), + lo_next=wp.empty(shape=(nworld,), dtype=wp.vec3), + lo_next_alpha=wp.empty(shape=(nworld,), dtype=float), + hi_next=wp.empty(shape=(nworld,), dtype=wp.vec3), + hi_next_alpha=wp.empty(shape=(nworld,), dtype=float), + mid=wp.empty(shape=(nworld,), dtype=wp.vec3), + mid_alpha=wp.empty(shape=(nworld,), dtype=float), + cost_candidate=wp.empty(shape=(nworld, mjm.opt.ls_iterations), dtype=float), + # TODO(team): skip allocation if not elliptic + u=wp.empty((nconmax,), dtype=types.vec6), + uu=wp.empty((nconmax,), dtype=float), + uv=wp.empty((nconmax,), dtype=float), + vv=wp.empty((nconmax,), dtype=float), + condim=wp.empty((nworld, njmax), dtype=int), + ), + # TODO(team): skip allocation if integrator != RK4 + qpos_t0=wp.empty((nworld, mjm.nq), dtype=float), + qvel_t0=wp.empty((nworld, mjm.nv), dtype=float), + act_t0=wp.empty((nworld, mjm.na), dtype=float), + qvel_rk=wp.empty((nworld, mjm.nv), dtype=float), + qacc_rk=wp.empty((nworld, mjm.nv), dtype=float), + act_dot_rk=wp.empty((nworld, mjm.na), dtype=float), + # TODO(team): skip allocation if integrator != euler | implicit + qfrc_integration=wp.zeros((nworld, mjm.nv), dtype=float), + qacc_integration=wp.zeros((nworld, mjm.nv), dtype=float), + act_vel_integration=wp.zeros((nworld, mjm.nu), dtype=float), + qM_integration=tile(qM_integration), + qLD_integration=tile(qLD_integration), + qLDiagInv_integration=wp.zeros((nworld, mjm.nv), dtype=float), + # TODO(team): skip allocation if broadphase != sap + sap_projection_lower=wp.zeros((nworld, mjm.ngeom, 2), dtype=float), + sap_projection_upper=wp.zeros((nworld, mjm.ngeom), dtype=float), + sap_sort_index=wp.zeros((nworld, mjm.ngeom, 2), dtype=int), + sap_range=wp.zeros((nworld, mjm.ngeom), dtype=int), + sap_cumulative_sum=wp.zeros((nworld, mjm.ngeom), dtype=int), + sap_segment_index=arr(np.array([i * mjm.ngeom if i < nworld + 1 else 0 for i in range(2 * nworld)]).reshape((nworld, 2))), + # collision driver + collision_pair=wp.empty(nconmax, dtype=wp.vec2i), + collision_hftri_index=wp.empty(nconmax, dtype=int), + collision_pairid=wp.empty(nconmax, dtype=int), + collision_worldid=wp.empty(nconmax, dtype=int), + ncollision=wp.zeros(1, dtype=int), + # narrowphase (EPA polytope) + epa_vert=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=wp.vec3), + epa_vert1=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=wp.vec3), + epa_vert2=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=wp.vec3), + epa_vert_index1=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=int), + epa_vert_index2=wp.zeros(shape=(nconmax, 5 + MJ_CCD_ITERATIONS), dtype=int), + epa_face=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=wp.vec3i), + epa_pr=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=wp.vec3), + epa_norm2=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=float), + epa_index=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=int), + epa_map=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=int), + epa_horizon=wp.zeros(shape=(nconmax, 6 * MJ_CCD_ITERATIONS), dtype=int), + # rne_postconstraint but also smooth + cacc=tile(mjd.cacc, dtype=wp.spatial_vector), + cfrc_int=tile(mjd.cfrc_int, dtype=wp.spatial_vector), + cfrc_ext=tile(mjd.cfrc_ext, dtype=wp.spatial_vector), + # tendon + ten_length=tile(mjd.ten_length), + ten_J=tile(ten_J), + ten_Jdot=wp.zeros((nworld, mjm.ntendon, mjm.nv), dtype=float), + ten_bias_coef=wp.zeros((nworld, mjm.ntendon), dtype=float), + ten_wrapadr=tile(mjd.ten_wrapadr), + ten_wrapnum=tile(mjd.ten_wrapnum), + ten_actfrc=wp.zeros((nworld, mjm.ntendon), dtype=float), + wrap_obj=tile(mjd.wrap_obj, dtype=wp.vec2i), + wrap_xpos=tile(mjd.wrap_xpos, dtype=wp.spatial_vector), + wrap_geom_xpos=wp.zeros((nworld, mjm.nwrap), dtype=wp.spatial_vector), + # sensors + sensordata=tile(mjd.sensordata), + sensor_rangefinder_pnt=wp.zeros((nworld, nrangefinder), dtype=wp.vec3), + sensor_rangefinder_vec=wp.zeros((nworld, nrangefinder), dtype=wp.vec3), + sensor_rangefinder_dist=wp.zeros((nworld, nrangefinder), dtype=float), + sensor_rangefinder_geomid=wp.zeros((nworld, nrangefinder), dtype=int), + # ray + ray_bodyexclude=wp.zeros(1, dtype=int), + ray_dist=wp.zeros((nworld, 1), dtype=float), + ray_geomid=wp.zeros((nworld, 1), dtype=int), + # mul_m + energy_vel_mul_m_skip=wp.zeros((nworld,), dtype=bool), + inverse_mul_m_skip=wp.zeros((nworld,), dtype=bool), + # actuator + actuator_trntype_body_ncon=wp.zeros((nworld, np.sum(mjm.actuator_trntype == mujoco.mjtTrn.mjTRN_BODY)), dtype=int), + ) + + +def get_data_into( + result: mujoco.MjData, + mjm: mujoco.MjModel, + d: types.Data, +): + """Gets data from a device into an existing mujoco.MjData. + + Args: + result (mujoco.MjData): The data object containing the current state and output arrays + (host). + mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host). + d (Data): The data object containing the current state and output arrays (device). + """ + if d.nworld > 1: + raise NotImplementedError("only nworld == 1 supported for now") + + result.solver_niter[0] = d.solver_niter.numpy()[0] + + ncon = d.ncon.numpy()[0] + nefc = d.nefc.numpy()[0] + + if ncon != result.ncon or nefc != result.nefc: + mujoco._functions._realloc_con_efc(result, ncon=ncon, nefc=nefc) + + result.time = d.time.numpy()[0] + result.energy = d.energy.numpy()[0] + result.ne = d.ne.numpy()[0] + result.qpos[:] = d.qpos.numpy()[0] + result.qvel[:] = d.qvel.numpy()[0] + result.qacc_warmstart = d.qacc_warmstart.numpy()[0] + result.qfrc_applied = d.qfrc_applied.numpy()[0] + result.mocap_pos = d.mocap_pos.numpy()[0] + result.mocap_quat = d.mocap_quat.numpy()[0] + result.qacc = d.qacc.numpy()[0] + result.xanchor = d.xanchor.numpy()[0] + result.xaxis = d.xaxis.numpy()[0] + result.xmat = d.xmat.numpy().reshape((-1, 9)) + result.xpos = d.xpos.numpy()[0] + result.xquat = d.xquat.numpy()[0] + result.xipos = d.xipos.numpy()[0] + result.ximat = d.ximat.numpy().reshape((-1, 9)) + result.subtree_com = d.subtree_com.numpy()[0] + result.geom_xpos = d.geom_xpos.numpy()[0] + result.geom_xmat = d.geom_xmat.numpy().reshape((-1, 9)) + result.site_xpos = d.site_xpos.numpy()[0] + result.site_xmat = d.site_xmat.numpy().reshape((-1, 9)) + result.cam_xpos = d.cam_xpos.numpy()[0] + result.cam_xmat = d.cam_xmat.numpy().reshape((-1, 9)) + result.light_xpos = d.light_xpos.numpy()[0] + result.light_xdir = d.light_xdir.numpy()[0] + result.cinert = d.cinert.numpy()[0] + result.flexvert_xpos = d.flexvert_xpos.numpy()[0] + result.flexedge_length = d.flexedge_length.numpy()[0] + result.flexedge_velocity = d.flexedge_velocity.numpy()[0] + result.cdof = d.cdof.numpy()[0] + result.crb = d.crb.numpy()[0] + result.qLDiagInv = d.qLDiagInv.numpy()[0] + result.ctrl = d.ctrl.numpy()[0] + result.ten_velocity = d.ten_velocity.numpy()[0] + result.actuator_velocity = d.actuator_velocity.numpy()[0] + result.actuator_force = d.actuator_force.numpy()[0] + result.actuator_length = d.actuator_length.numpy()[0] + mujoco.mju_dense2sparse( + result.actuator_moment, + d.actuator_moment.numpy()[0], + result.moment_rownnz, + result.moment_rowadr, + result.moment_colind, + ) + result.cvel = d.cvel.numpy()[0] + result.cdof_dot = d.cdof_dot.numpy()[0] + result.qfrc_bias = d.qfrc_bias.numpy()[0] + result.qfrc_fluid = d.qfrc_fluid.numpy()[0] + result.qfrc_passive = d.qfrc_passive.numpy()[0] + result.subtree_linvel = d.subtree_linvel.numpy()[0] + result.subtree_angmom = d.subtree_angmom.numpy()[0] + result.qfrc_spring = d.qfrc_spring.numpy()[0] + result.qfrc_damper = d.qfrc_damper.numpy()[0] + result.qfrc_gravcomp = d.qfrc_gravcomp.numpy()[0] + result.qfrc_fluid = d.qfrc_fluid.numpy()[0] + result.qfrc_actuator = d.qfrc_actuator.numpy()[0] + result.qfrc_smooth = d.qfrc_smooth.numpy()[0] + result.qfrc_constraint = d.qfrc_constraint.numpy()[0] + result.qfrc_inverse = d.qfrc_inverse.numpy()[0] + result.qacc_smooth = d.qacc_smooth.numpy()[0] + result.act = d.act.numpy()[0] + result.act_dot = d.act_dot.numpy()[0] + + result.contact.dist[:] = d.contact.dist.numpy()[:ncon] + result.contact.pos[:] = d.contact.pos.numpy()[:ncon] + result.contact.frame[:] = d.contact.frame.numpy()[:ncon].reshape((-1, 9)) + result.contact.includemargin[:] = d.contact.includemargin.numpy()[:ncon] + result.contact.friction[:] = d.contact.friction.numpy()[:ncon] + result.contact.solref[:] = d.contact.solref.numpy()[:ncon] + result.contact.solreffriction[:] = d.contact.solreffriction.numpy()[:ncon] + result.contact.solimp[:] = d.contact.solimp.numpy()[:ncon] + result.contact.dim[:] = d.contact.dim.numpy()[:ncon] + result.contact.efc_address[:] = d.contact.efc_address.numpy()[:ncon, 0] + + if mujoco.mj_isSparse(mjm): + result.qM[:] = d.qM.numpy()[0, 0] + result.qLD[:] = d.qLD.numpy()[0, 0] + # TODO(team): set efc_J after fix to _realloc_con_efc lands + # efc_J = d.efc_J.numpy()[0, :nefc] + # mujoco.mju_dense2sparse( + # result.efc_J, efc_J, result.efc_J_rownnz, result.efc_J_rowadr, result.efc_J_colind + # ) + else: + qM = d.qM.numpy() + adr = 0 + for i in range(mjm.nv): + j = i + while j >= 0: + result.qM[adr] = qM[0, i, j] + j = mjm.dof_parentid[j] + adr += 1 + mujoco.mj_factorM(mjm, result) + # TODO(team): set efc_J after fix to _realloc_con_efc lands + # if nefc > 0: + # result.efc_J[:nefc * mjm.nv] = d.efc_J.numpy()[:nefc].flatten() + result.xfrc_applied[:] = d.xfrc_applied.numpy()[0] + result.eq_active[:] = d.eq_active.numpy()[0] + + # TODO(team): set these efc_* fields after fix to _realloc_con_efc + # Safely copy only up to the minimum of the destination and source sizes + # n = min(result.efc_D.shape[0], d.efc.D.numpy()[:nefc].shape[0]) + # result.efc_D[:n] = d.efc.D.numpy()[:nefc][:n] + # n_pos = min(result.efc_pos.shape[0], d.efc.pos.numpy()[:nefc].shape[0]) + # result.efc_pos[:n_pos] = d.efc.pos.numpy()[:nefc][:n_pos] + + # n_aref = min(result.efc_aref.shape[0], d.efc.aref.numpy()[:nefc].shape[0]) + # result.efc_aref[:n_aref] = d.efc.aref.numpy()[:nefc][:n_aref] + + # n_force = min(result.efc_force.shape[0], d.efc.force.numpy()[:nefc].shape[0]) + # result.efc_force[:n_force] = d.efc.force.numpy()[:nefc][:n_force] + + # n_margin = min(result.efc_margin.shape[0], d.efc.margin.numpy()[:nefc].shape[0]) + # result.efc_margin[:n_margin] = d.efc.margin.numpy()[:nefc][:n_margin] + + result.cacc[:] = d.cacc.numpy()[0] + result.cfrc_int[:] = d.cfrc_int.numpy()[0] + result.cfrc_ext[:] = d.cfrc_ext.numpy()[0] + + # TODO: other efc_ fields, anything else missing + + # tendon + result.ten_length[:] = d.ten_length.numpy()[0] + result.ten_J[:] = d.ten_J.numpy()[0] + result.ten_wrapadr[:] = d.ten_wrapadr.numpy()[0] + result.ten_wrapnum[:] = d.ten_wrapnum.numpy()[0] + result.wrap_obj[:] = d.wrap_obj.numpy()[0] + result.wrap_xpos[:] = d.wrap_xpos.numpy()[0] + + # sensors + result.sensordata[:] = d.sensordata.numpy() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py new file mode 100644 index 00000000..33503268 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py @@ -0,0 +1,354 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Tests for io functions.""" + +import dataclasses +import typing +from typing import Any, Dict, Optional, Union + +import mujoco +import numpy as np +import warp as wp +from absl.testing import absltest + +import mujoco_warp as mjwarp + +from mujoco.mjx.third_party.mujoco_warp._src import test_util + + +def _dims_match(test_obj, d1: Any, d2: Any, prefix: str = ""): + """Checks that two dataclasses have fields with the same leading dims.""" + fields1, fields2 = dataclasses.fields(d1), dataclasses.fields(d2) + for f1, f2 in zip(fields1, fields2): + full_name = prefix + f1.name + a1, a2 = getattr(d1, f1.name), getattr(d2, f2.name) + if dataclasses.is_dataclass(a1) or dataclasses.is_dataclass(a2): + _dims_match(test_obj, a1, a2, prefix + f1.name + ".") + continue + + if isinstance(f1.type, wp.types.array) or isinstance(f2.type, wp.types.array): + s1, s2 = a1.shape, a2.shape + test_obj.assertEqual(len(s1), len(s2), full_name + f" dims mismatch. Got {s1} and {s2}.") + test_obj.assertSequenceAlmostEqual(s1, s2, full_name + f" dims mismatch. Got {s1} and {s2}.") + + +def _get_np_scalar_type(val: Any) -> Optional[Union[bool, int, float]]: + """Returns the python type from a numpy scalar.""" + is_np_scalar = list(isinstance(val, t) for t in (np.integer, np.floating, np.bool_)) + if any(is_np_scalar): + return [int, float, bool][is_np_scalar.index(True)] + + +def _check_type_matches_annotation(test_obj, obj: Any, prefix: str = ""): + """Checks that dataclass annotations match the runtime types.""" + assert dataclasses.is_dataclass(obj), prefix + " must be dataclass." + msg = "Type of {val_type} does not match annotation {type_} for field {prefix}{field_name}" + + for field in dataclasses.fields(obj): + field_name = field.name + val = getattr(obj, field_name) + val_type = type(val) + type_ = field.type + + if dataclasses.is_dataclass(val): + test_obj.assertTrue(dataclasses.is_dataclass(type_), msg.format(**locals())) + _check_type_matches_annotation(test_obj, val, prefix + field_name + ".") + continue + + np_scalar_type = _get_np_scalar_type(val) + if np_scalar_type: + test_obj.assertIsInstance(np_scalar_type(val), type_, msg.format(**locals())) + continue + + if isinstance(type_, wp.types.array): + test_obj.assertIsInstance(val, wp.types.array, msg.format(**locals())) + continue + + origin_type = typing.get_origin(type_) + if tuple in (val_type, origin_type): + test_obj.assertEqual(val_type, origin_type, msg.format(**locals())) + field_name += ".tuple[]" + type_ = typing.get_args(type_)[0] + + items = val + for val in items: + val_type = type(val) + if dataclasses.is_dataclass(val): + _check_type_matches_annotation(test_obj, val, prefix + field_name) + continue + + np_scalar_type = _get_np_scalar_type(val) + if np_scalar_type: + test_obj.assertIsInstance(np_scalar_type(val), type_, msg.format(**locals())) + continue + + if isinstance(type_, wp.types.array): + test_obj.assertIsInstance(val, wp.types.array, msg.format(**locals())) + continue + + test_obj.assertEqual(type(val), type_, msg.format(**locals())) + continue + + test_obj.assertEqual(type(val), field.type, msg.format(**locals())) + + +def _check_annotation_compat( + annotations: Dict[str, Any], prefix: str = "", in_cls: bool = False, in_tuple: bool = False +) -> Dict[str, Any]: + """Checks that dataclass annotations match criteria for JAX API compat.""" + for k, v in annotations.items(): + full_key = f"{prefix}{k}" + info = f"Found {v} for annotation {full_key}." + + if v in (int, bool, float): + continue + + if isinstance(v, wp.types.array): + continue + + if v in wp.types.vector_types: + raise AssertionError(f"Vector types are not allowed. {info}") + + if typing.get_origin(v) == tuple and (in_cls or in_tuple): + raise AssertionError(f"Nested args in Model/Data must not be tuple. {info}") + + if typing.get_origin(v) == tuple: + tuple_args = typing.get_args(v) + if len(tuple_args) != 2 and tuple_args[1] != ...: + raise AssertionError(f"Tuple args must be variadic. {info}") + + _check_annotation_compat( + {"[]": tuple_args[0]}, + prefix=f"{full_key}.tuple", + in_cls=in_cls, + in_tuple=True, + ) + continue + + if hasattr(v, "__class__") and in_cls: + raise AssertionError(f"Nested object args in Model/Data are not allowed. {info}") + + if hasattr(v, "__class__") and not dataclasses.is_dataclass(v): + raise AssertionError(f"Args that are objects must be dataclass. {info}") + + if hasattr(v, "__class__") and not v.__module__.startswith("mujoco_warp"): + raise AssertionError(f"dataclass args must be within the mujoco_warp module. {info}") + + if hasattr(v, "__class__"): + _check_annotation_compat(v.__annotations__, prefix=f"{full_key}{v.__name__}.", in_cls=True, in_tuple=in_tuple) + continue + + raise AssertionError(f"Model/Data annotation is not allowed. {info}") + + +def _leading_dims_scale_w_nworld(test_obj, d1: Any, d2: Any, nworld1: int, nworld2: int, prefix: str = ""): + """Checks that dataclass fields that scale with nworld have leading dim nworld.""" + msg = "Arrays that scale with nworld should have leading dim nworld." + fields1, fields2 = dataclasses.fields(d1), dataclasses.fields(d2) + for f1, f2 in zip(fields1, fields2): + full_name = prefix + f1.name + a1, a2 = getattr(d1, f1.name), getattr(d2, f2.name) + if dataclasses.is_dataclass(a1) or dataclasses.is_dataclass(a2): + _leading_dims_scale_w_nworld(test_obj, a1, a2, nworld1, nworld2, prefix + f1.name + ".") + continue + + if isinstance(f1.type, wp.types.array) or isinstance(f2.type, wp.types.array): + s1, s2 = a1.shape[0], a2.shape[0] + if s1 == s2: + continue + test_obj.assertEqual(s2, nworld2, full_name + f" has leading dim {s2} with nworld={nworld2}. {msg}") + test_obj.assertEqual(s1, nworld1, full_name + f" has leading dim {s1} with nworld={nworld1}. {msg}") + + +class IOTest(absltest.TestCase): + def test_make_put_data(self): + """Tests that make_data and put_data are producing the same shapes for all arrays.""" + mjm, _, _, d = test_util.fixture("pendula.xml") + md = mjwarp.make_data(mjm, nconmax=512, njmax=512) + + # same number of fields + self.assertEqual(len(d.__dict__), len(md.__dict__)) + + # test shapes for all arrays + for attr, val in md.__dict__.items(): + if isinstance(val, wp.array): + self.assertEqual(val.shape, getattr(d, attr).shape, f"{attr} shape mismatch") + + # TODO(team): sensors + + def test_get_data_into_m(self): + mjm = mujoco.MjModel.from_xml_string(""" + + + + + + + + + + + + + """) + + mjd = mujoco.MjData(mjm) + mujoco.mj_forward(mjm, mjd) + + mjd_ref = mujoco.MjData(mjm) + mujoco.mj_forward(mjm, mjd_ref) + + m = mjwarp.put_model(mjm) + d = mjwarp.put_data(mjm, mjd) + + mjd.qLD.fill(-123) + mjd.qM.fill(-123) + + mjwarp.get_data_into(mjd, mjm, d) + np.testing.assert_allclose(mjd.qLD, mjd_ref.qLD) + np.testing.assert_allclose(mjd.qM, mjd_ref.qM) + + def test_ellipsoid_fluid_model(self): + with self.assertRaises(NotImplementedError): + mjm = mujoco.MjModel.from_xml_string( + """ + + + """ + ) + mjwarp.put_model(mjm) + + def test_jacobian_auto(self): + mjm = mujoco.MjModel.from_xml_string(""" + + + """) + mjwarp.put_model(mjm) + + def test_put_data_qLD(self): + mjm = mujoco.MjModel.from_xml_string(""" + + + + + + + + + """) + mjd = mujoco.MjData(mjm) + d = mjwarp.put_data(mjm, mjd) + self.assertTrue((d.qLD.numpy() == 0.0).all()) + + mujoco.mj_forward(mjm, mjd) + mjd.qM[:] = 0.0 + d = mjwarp.put_data(mjm, mjd) + self.assertTrue((d.qLD.numpy() == 0.0).all()) + + mujoco.mj_forward(mjm, mjd) + mjd.qLD[:] = 0.0 + d = mjwarp.put_data(mjm, mjd) + self.assertTrue((d.qLD.numpy() == 0.0).all()) + + def test_noslip_solver(self): + with self.assertRaises(NotImplementedError): + test_util.fixture( + xml=""" + + + """ + ) + + def test_put_model_nworld_array(self): + """Tests that put_model arrays with nworld leading dim have `_is_batched`.""" + mjm, *_ = test_util.fixture("pendula.xml") + m1 = mjwarp.put_model(mjm) + + self.assertTrue(hasattr(m1.geom_pos, "_is_batched")) + self.assertEqual(m1.geom_pos.shape[0], 1) + self.assertEqual(m1.geom_pos.strides[0], 0) + self.assertLen(m1.geom_pos.strides, m1.geom_pos.ndim) + self.assertTrue(hasattr(m1.opt.gravity, "_is_batched")) + self.assertEqual(m1.opt.gravity.shape[0], 1) + self.assertEqual(m1.opt.gravity.strides[0], 0) + self.assertLen(m1.opt.gravity.strides, m1.opt.gravity.ndim) + self.assertFalse(hasattr(m1.body_parentid, "_is_batched")) + self.assertGreater(m1.body_parentid.shape[0], 0) + self.assertGreater(m1.body_parentid.strides[0], 0) + self.assertLen(m1.body_parentid.strides, m1.body_parentid.ndim) + + def test_put_data_nworld_array(self): + """Tests that put_data arrays that scale with nworld have leading dim nworld.""" + mjm, mjd, _, _ = test_util.fixture("pendula.xml") + d1 = mjwarp.put_data(mjm, mjd, nworld=1, nconmax=1_000, njmax=1_000) + dn = mjwarp.put_data(mjm, mjd, nworld=133, nconmax=1_000, njmax=1_000) + _leading_dims_scale_w_nworld(self, d1, dn, 1, 133) + + def test_make_data_nworld_array(self): + """Tests that make_data arrays that scale with nworld have leading dim nworld.""" + mjm, *_ = test_util.fixture("pendula.xml") + d1 = mjwarp.make_data(mjm, nworld=1, nconmax=1_000, njmax=1_000) + dn = mjwarp.make_data(mjm, nworld=133, nconmax=1_000, njmax=1_000) + _leading_dims_scale_w_nworld(self, d1, dn, 1, 133) + + def test_public_api_jax_compat(self): + """Tests that annotations meet a set of criteria for JAX compat.""" + _check_annotation_compat(mjwarp.Model.__annotations__, "Model.") + _check_annotation_compat(mjwarp.Data.__annotations__, "Data.") + + def test_types_match_annotations(self): + """Tests that the types of dataclass fields match the annotations.""" + mjm, _, m, d = test_util.fixture("pendula.xml") + + _check_type_matches_annotation(self, m, "Model.") + _check_type_matches_annotation(self, d, "Data.") + + d = mjwarp.make_data(mjm, nworld=2) + _check_type_matches_annotation(self, d, "Data.") + + def test_make_put_data_dims_match(self): + """Tests that make_data and put_data have matching dimensions.""" + mjm, mjd, _, _ = test_util.fixture("pendula.xml") + dm2 = mjwarp.make_data(mjm, nworld=2, nconmax=13, njmax=42) + dm3 = mjwarp.make_data(mjm, nworld=3, nconmax=13, njmax=42) + + dp2 = mjwarp.put_data(mjm, mjd, nworld=2, nconmax=13, njmax=42) + dp3 = mjwarp.put_data(mjm, mjd, nworld=3, nconmax=13, njmax=42) + + _dims_match(self, dm2, dp2) + _dims_match(self, dm3, dp3) + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/jax_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/jax_test.py new file mode 100644 index 00000000..529a2288 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/jax_test.py @@ -0,0 +1,101 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import os + +import warp as wp +from absl.testing import absltest +from absl.testing import parameterized + +import mujoco_warp as mjwarp +from mujoco.mjx.third_party.mujoco_warp._src.test_util import fixture + +# TODO(team): JAX test is temporary, remove after we land MJX:Warp + + +class JAXTest(parameterized.TestCase): + @parameterized.parameters("humanoid/humanoid.xml", "pendula.xml") + def test_jax(self, xml): + os.environ["XLA_FLAGS"] = "--xla_gpu_graph_min_graph_size=1" + # Force JAX to allocate memory on demand and deallocate when not needed (slow) + os.environ["XLA_PYTHON_CLIENT_ALLOCATOR"] = "platform" + + try: + import jax + except ImportError: + self.skipTest("JAX not installed") + + from jax import numpy as jp + from warp.jax_experimental.ffi import jax_callable + + if jax.default_backend() != "gpu": + self.skipTest("JAX default backend is not GPU") + + NWORLDS = 2 + NCONTACTS = 16 + UNROLL_LENGTH = 1 + + mjm, _, m, d = fixture( + xml, + nworld=NWORLDS, + nconmax=NWORLDS * NCONTACTS, + njmax=NWORLDS * NCONTACTS * 4, + iterations=1, + ls_iterations=4, + kick=True, + ) + + # Disable CUDA graph conditional + m.opt.graph_conditional = False + + def warp_step( + qpos_in: wp.array(dtype=wp.float32, ndim=2), + qvel_in: wp.array(dtype=wp.float32, ndim=2), + qpos_out: wp.array(dtype=wp.float32, ndim=2), + qvel_out: wp.array(dtype=wp.float32, ndim=2), + ): + wp.copy(d.qpos, qpos_in) + wp.copy(d.qvel, qvel_in) + mjwarp.step(m, d) + wp.copy(qpos_out, d.qpos) + wp.copy(qvel_out, d.qvel) + + def unroll(qpos, qvel): + def step(carry, _): + qpos, qvel = carry + qpos, qvel = warp_step_fn(qpos, qvel) + return (qpos, qvel), None + + (qpos, qvel), _ = jax.lax.scan(step, (qpos, qvel), length=UNROLL_LENGTH) + + return qpos, qvel + + warp_step_fn = jax_callable( + warp_step, + num_outputs=2, + output_dims={"qpos_out": (NWORLDS, mjm.nq), "qvel_out": (NWORLDS, mjm.nv)}, + graph_compatible=True, + ) + + jax_qpos = jp.tile(jp.array(m.qpos0.numpy()), (NWORLDS, 1)) + jax_qvel = jp.zeros((NWORLDS, m.nv)) + + jax_unroll_fn = jax.jit(unroll).lower(jax_qpos, jax_qvel).compile() + jax_unroll_fn(jax_qpos, jax_qvel) + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py new file mode 100644 index 00000000..99cc0b20 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py @@ -0,0 +1,289 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +from typing import Any, Tuple + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src import types + + +@wp.func +def mul_quat(u: wp.quat, v: wp.quat) -> wp.quat: + return wp.quat( + u[0] * v[0] - u[1] * v[1] - u[2] * v[2] - u[3] * v[3], + u[0] * v[1] + u[1] * v[0] + u[2] * v[3] - u[3] * v[2], + u[0] * v[2] - u[1] * v[3] + u[2] * v[0] + u[3] * v[1], + u[0] * v[3] + u[1] * v[2] - u[2] * v[1] + u[3] * v[0], + ) + + +@wp.func +def quat_mul_axis(q: wp.quat, axis: wp.vec3f) -> wp.quat: + """Multiplies a quaternion and an axis.""" + return wp.quat( + -q[1] * axis[0] - q[2] * axis[1] - q[3] * axis[2], + q[0] * axis[0] + q[2] * axis[2] - q[3] * axis[1], + q[0] * axis[1] + q[3] * axis[0] - q[1] * axis[2], + q[0] * axis[2] + q[1] * axis[1] - q[2] * axis[0], + ) + + +@wp.func +def rot_vec_quat(vec: wp.vec3, quat: wp.quat) -> wp.vec3: + s, u = quat[0], wp.vec3(quat[1], quat[2], quat[3]) + r = 2.0 * (wp.dot(u, vec) * u) + (s * s - wp.dot(u, u)) * vec + r = r + 2.0 * s * wp.cross(u, vec) + return r + + +@wp.func +def axis_angle_to_quat(axis: wp.vec3, angle: float) -> wp.quat: + s, c = wp.sin(angle * 0.5), wp.cos(angle * 0.5) + axis = axis * s + return wp.quat(c, axis[0], axis[1], axis[2]) + + +@wp.func +def quat_to_mat(quat: wp.quat) -> wp.mat33: + """Converts a quaternion into a 9-dimensional rotation matrix.""" + vec = wp.vec4(quat[0], quat[1], quat[2], quat[3]) + q = wp.outer(vec, vec) + + return wp.mat33( + q[0, 0] + q[1, 1] - q[2, 2] - q[3, 3], + 2.0 * (q[1, 2] - q[0, 3]), + 2.0 * (q[1, 3] + q[0, 2]), + 2.0 * (q[1, 2] + q[0, 3]), + q[0, 0] - q[1, 1] + q[2, 2] - q[3, 3], + 2.0 * (q[2, 3] - q[0, 1]), + 2.0 * (q[1, 3] - q[0, 2]), + 2.0 * (q[2, 3] + q[0, 1]), + q[0, 0] - q[1, 1] - q[2, 2] + q[3, 3], + ) + + +@wp.func +def quat_inv(quat: wp.quat) -> wp.quat: + return wp.quat(quat[0], -quat[1], -quat[2], -quat[3]) + + +@wp.func +def inert_vec(i: types.vec10, v: wp.spatial_vector) -> wp.spatial_vector: + """mju_mulInertVec: multiply 6D vector (rotation, translation) by 6D inertia matrix.""" + return wp.spatial_vector( + i[0] * v[0] + i[3] * v[1] + i[4] * v[2] - i[8] * v[4] + i[7] * v[5], + i[3] * v[0] + i[1] * v[1] + i[5] * v[2] + i[8] * v[3] - i[6] * v[5], + i[4] * v[0] + i[5] * v[1] + i[2] * v[2] - i[7] * v[3] + i[6] * v[4], + i[8] * v[1] - i[7] * v[2] + i[9] * v[3], + i[6] * v[2] - i[8] * v[0] + i[9] * v[4], + i[7] * v[0] - i[6] * v[1] + i[9] * v[5], + ) + + +@wp.func +def motion_cross(u: wp.spatial_vector, v: wp.spatial_vector) -> wp.spatial_vector: + """Cross product of two motions.""" + + u0 = wp.vec3(u[0], u[1], u[2]) + u1 = wp.vec3(u[3], u[4], u[5]) + v0 = wp.vec3(v[0], v[1], v[2]) + v1 = wp.vec3(v[3], v[4], v[5]) + + ang = wp.cross(u0, v0) + vel = wp.cross(u1, v0) + wp.cross(u0, v1) + + return wp.spatial_vector(ang, vel) + + +@wp.func +def motion_cross_force(v: wp.spatial_vector, f: wp.spatial_vector) -> wp.spatial_vector: + """Cross product of a motion and a force.""" + + v0 = wp.vec3(v[0], v[1], v[2]) + v1 = wp.vec3(v[3], v[4], v[5]) + f0 = wp.vec3(f[0], f[1], f[2]) + f1 = wp.vec3(f[3], f[4], f[5]) + + ang = wp.cross(v0, f0) + wp.cross(v1, f1) + vel = wp.cross(v0, f1) + + return wp.spatial_vector(ang, vel) + + +@wp.func +def quat_to_vel(quat: wp.quat) -> wp.vec3: + axis = wp.vec3(quat[1], quat[2], quat[3]) + sin_a_2 = wp.norm_l2(axis) + + if sin_a_2 == 0.0: + return wp.vec3(0.0) + + speed = 2.0 * wp.atan2(sin_a_2, quat[0]) + # when axis-angle is larger than pi, rotation is in the opposite direction + if speed > wp.pi: + speed -= 2.0 * wp.pi + + return axis * speed / sin_a_2 + + +@wp.func +def quat_sub(qa: wp.quat, qb: wp.quat) -> wp.vec3: + """Subtract quaternions, express as 3D velocity: qb*quat(res) = qa.""" + # qdif = neg(qb)*qa + qneg = wp.quat(qb[0], -qb[1], -qb[2], -qb[3]) + qdif = mul_quat(qneg, qa) + + # convert to 3D velocity + return quat_to_vel(qdif) + + +@wp.func +def quat_integrate(q: wp.quat, v: wp.vec3, dt: float) -> wp.quat: + """Integrates a quaternion given angular velocity and dt.""" + norm_ = wp.length(v) + v = wp.normalize(v) # does that need proper zero gradient handling? + angle = dt * norm_ + + q_res = axis_angle_to_quat(v, angle) + q = wp.normalize(q) + q_res = mul_quat(q, q_res) + + return wp.normalize(q_res) + + +@wp.func +def orthogonals(a: wp.vec3): + y = wp.vec3(0.0, 1.0, 0.0) + z = wp.vec3(0.0, 0.0, 1.0) + b = wp.where((-0.5 < a[1]) and (a[1] < 0.5), y, z) + b = b - a * wp.dot(a, b) + b = wp.normalize(b) + if wp.length(a) == 0.0: + b = wp.vec3(0.0, 0.0, 0.0) + c = wp.cross(a, b) + + return b, c + + +@wp.func +def orthonormal(normal: wp.vec3) -> wp.vec3: + if wp.abs(normal[0]) < wp.abs(normal[1]) and wp.abs(normal[0]) < wp.abs(normal[2]): + dir = wp.vec3(1.0 - normal[0] * normal[0], -normal[0] * normal[1], -normal[0] * normal[2]) + elif wp.abs(normal[1]) < wp.abs(normal[2]): + dir = wp.vec3(-normal[1] * normal[0], 1.0 - normal[1] * normal[1], -normal[1] * normal[2]) + else: + dir = wp.vec3(-normal[2] * normal[0], -normal[2] * normal[1], 1.0 - normal[2] * normal[2]) + dir, _ = gjk_normalize(dir) + return dir + + +@wp.func +def gjk_normalize(a: wp.vec3): + norm = wp.length(a) + if norm > 1e-8 and norm < 1e12: + return a / norm, True + return a, False + + +@wp.func +def make_frame(a: wp.vec3): + a = wp.normalize(a) + b, c = orthogonals(a) + + # fmt: off + return wp.mat33( + a.x, a.y, a.z, + b.x, b.y, b.z, + c.x, c.y, c.z + ) + # fmt: on + + +@wp.func +def normalize_with_norm(x: Any): + norm = wp.length(x) + if norm == 0.0: + return x, 0.0 + return x / norm, norm + + +@wp.func +def closest_segment_point(a: wp.vec3, b: wp.vec3, pt: wp.vec3) -> wp.vec3: + """Returns the closest point on the a-b line segment to a point pt.""" + ab = b - a + t = wp.dot(pt - a, ab) / (wp.dot(ab, ab) + 1e-6) + return a + wp.clamp(t, 0.0, 1.0) * ab + + +@wp.func +def closest_segment_point_and_dist(a: wp.vec3, b: wp.vec3, pt: wp.vec3) -> Tuple[wp.vec3, float]: + """Returns closest point on the line segment and the distance squared.""" + closest = closest_segment_point(a, b, pt) + dist = wp.dot((pt - closest), (pt - closest)) + return closest, dist + + +@wp.func +def closest_segment_to_segment_points(a0: wp.vec3, a1: wp.vec3, b0: wp.vec3, b1: wp.vec3) -> Tuple[wp.vec3, wp.vec3]: + """Returns closest points between two line segments.""" + + dir_a, len_a = normalize_with_norm(a1 - a0) + dir_b, len_b = normalize_with_norm(b1 - b0) + + half_len_a = len_a * 0.5 + half_len_b = len_b * 0.5 + a_mid = a0 + dir_a * half_len_a + b_mid = b0 + dir_b * half_len_b + + trans = a_mid - b_mid + + dira_dot_dirb = wp.dot(dir_a, dir_b) + dira_dot_trans = wp.dot(dir_a, trans) + dirb_dot_trans = wp.dot(dir_b, trans) + denom = 1.0 - dira_dot_dirb * dira_dot_dirb + + orig_t_a = (-dira_dot_trans + dira_dot_dirb * dirb_dot_trans) / (denom + 1e-6) + orig_t_b = dirb_dot_trans + orig_t_a * dira_dot_dirb + t_a = wp.clamp(orig_t_a, -half_len_a, half_len_a) + t_b = wp.clamp(orig_t_b, -half_len_b, half_len_b) + + best_a = a_mid + dir_a * t_a + best_b = b_mid + dir_b * t_b + + new_a, d1 = closest_segment_point_and_dist(a0, a1, best_b) + new_b, d2 = closest_segment_point_and_dist(b0, b1, best_a) + if d1 < d2: + return new_a, best_b + return best_a, new_b + + +@wp.func +def safe_div(x: float, y: float) -> float: + return x / wp.where(y != 0.0, y, types.MJ_MINVAL) + + +@wp.func +def upper_tri_index(n: int, i: int, j: int) -> int: + """Returns index of a_ij = a_ji in upper triangular matrix (excluding diagonal).""" + return (i * (2 * n - i - 3)) // 2 + j - 1 + + +@wp.func +def upper_trid_index(n: int, i: int, j: int) -> int: + """Returns index of a_ij = a_ji in upper triangular matrix (including diagonal).""" + if j < i: + i, j = j, i + return (i * (2 * n - i - 1)) // 2 + j diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math_test.py new file mode 100644 index 00000000..3f1de2c7 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math_test.py @@ -0,0 +1,131 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import warp as wp +from absl.testing import absltest + +from mujoco.mjx.third_party.mujoco_warp._src.math import closest_segment_to_segment_points +from mujoco.mjx.third_party.mujoco_warp._src.math import upper_tri_index +from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index + + +class ClosestSegmentSegmentPointsTest(absltest.TestCase): + """Tests for closest segment-to-segment points.""" + + def test_closest_segments_points(self): + """Test closest points between two segments.""" + a0 = wp.vec3([0.73432405, 0.12372768, 0.20272314]) + a1 = wp.vec3([1.10600128, 0.88555209, 0.65209485]) + b0 = wp.vec3([0.85599262, 0.61736299, 0.9843583]) + b1 = wp.vec3([1.84270939, 0.92891793, 1.36343326]) + + best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1) + self.assertSequenceAlmostEqual(best_a, [1.09063, 0.85404, 0.63351], 5) + self.assertSequenceAlmostEqual(best_b, [0.99596, 0.66156, 1.03813], 5) + + def test_intersecting_segments(self): + """Tests segments that intersect.""" + a0, a1 = wp.vec3([0.0, 0.0, -1.0]), wp.vec3([0.0, 0.0, 1.0]) + b0, b1 = wp.vec3([-1.0, 0.0, 0.0]), wp.vec3([1.0, 0.0, 0.0]) + + best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1) + self.assertSequenceAlmostEqual(best_a, [0.0, 0.0, 0.0], 5) + self.assertSequenceAlmostEqual(best_b, [0.0, 0.0, 0.0], 5) + + def test_intersecting_lines(self): + """Tests that intersecting lines get clipped.""" + a0, a1 = wp.vec3([0.2, 0.2, 0.0]), wp.vec3([1.0, 1.0, 0.0]) + b0, b1 = wp.vec3([0.2, 0.4, 0.0]), wp.vec3([1.0, 2.0, 0.0]) + + best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1) + self.assertSequenceAlmostEqual(best_a, [0.3, 0.3, 0.0], 2) + self.assertSequenceAlmostEqual(best_b, [0.2, 0.4, 0.0], 2) + + def test_parallel_segments(self): + """Tests that parallel segments have closest points at the midpoint.""" + a0, a1 = wp.vec3([0.0, 0.0, -1.0]), wp.vec3([0.0, 0.0, 1.0]) + b0, b1 = wp.vec3([1.0, 0.0, -1.0]), wp.vec3([1.0, 0.0, 1.0]) + + best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1) + self.assertSequenceAlmostEqual(best_a, [0.0, 0.0, 0.0], 5) + self.assertSequenceAlmostEqual(best_b, [1.0, 0.0, 0.0], 5) + + def test_parallel_offset_segments(self): + """Tests that offset parallel segments are close at segment endpoints.""" + a0, a1 = wp.vec3([0.0, 0.0, -1.0]), wp.vec3([0.0, 0.0, 1.0]) + b0, b1 = wp.vec3([1.0, 0.0, 1.0]), wp.vec3([1.0, 0.0, 3.0]) + + best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1) + self.assertSequenceAlmostEqual(best_a, [0.0, 0.0, 1.0], 5) + self.assertSequenceAlmostEqual(best_b, [1.0, 0.0, 1.0], 5) + + def test_zero_length_segments(self): + """Test that zero length segments don't return NaNs.""" + a0, a1 = wp.vec3([0.0, 0.0, -1.0]), wp.vec3([0.0, 0.0, -1.0]) + b0, b1 = wp.vec3([1.0, 0.0, 0.1]), wp.vec3([1.0, 0.0, 0.1]) + + best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1) + self.assertSequenceAlmostEqual(best_a, [0.0, 0.0, -1.0], 5) + self.assertSequenceAlmostEqual(best_b, [1.0, 0.0, 0.1], 5) + + def test_overlapping_segments(self): + """Tests that perfectly overlapping segments intersect at the midpoints.""" + a0, a1 = wp.vec3([0.0, 0.0, -1.0]), wp.vec3([0.0, 0.0, 1.0]) + b0, b1 = wp.vec3([0.0, 0.0, -1.0]), wp.vec3([0.0, 0.0, 1.0]) + + best_a, best_b = closest_segment_to_segment_points(a0, a1, b0, b1) + self.assertSequenceAlmostEqual(best_a, [0.0, 0.0, 0.0], 5) + self.assertSequenceAlmostEqual(best_b, [0.0, 0.0, 0.0], 5) + + def test_upper_tri_index2(self): + """Tests upper_tri_index with size 2""" + arr = [] + for i in range(2): + for j in range(i + 1, 2): + arr.append(upper_tri_index(2, i, j)) + self.assertEqual(arr, list(range(0, 1))) + + def test_upper_tri_index10(self): + """Tests upper_tri_index with size 10""" + arr = [] + for i in range(10): + for j in range(i + 1, 10): + arr.append(upper_tri_index(10, i, j)) + self.assertEqual(arr, list(range(0, 45))) + + def test_upper_trid_index1(self): + """Tests upper_trid_index with size 1""" + arr = [] + for i in range(1): + for j in range(i, 1): + arr.append(upper_trid_index(1, i, j)) + self.assertEqual(arr, list(range(0, 1))) + + def test_upper_trid_index10(self): + """Tests upper_trid_index with size 10""" + arr = [] + for i in range(10): + for j in range(i, 10): + arr.append(upper_trid_index(10, i, j)) + self.assertEqual(arr, list(range(0, 55))) + + def test_upper_trid_index10(self): + """Tests upper_trid_index works with symmetric matrix""" + self.assertEqual(upper_trid_index(10, 1, 5), upper_trid_index(10, 5, 1)) + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py new file mode 100644 index 00000000..72b00d90 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py @@ -0,0 +1,617 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src import math +from mujoco.mjx.third_party.mujoco_warp._src import support +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import JointType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope + + +@wp.kernel +def _spring_damper_dof_passive( + # Model: + qpos_spring: wp.array2d(dtype=float), + jnt_type: wp.array(dtype=int), + jnt_qposadr: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + jnt_stiffness: wp.array2d(dtype=float), + dof_damping: wp.array2d(dtype=float), + # Data in: + qpos_in: wp.array2d(dtype=float), + qvel_in: wp.array2d(dtype=float), + # Data out: + qfrc_spring_out: wp.array2d(dtype=float), + qfrc_damper_out: wp.array2d(dtype=float), +): + worldid, jntid = wp.tid() + dofid = jnt_dofadr[jntid] + stiffness = jnt_stiffness[worldid, jntid] + damping = dof_damping[worldid, dofid] + + has_stiffness = stiffness != 0.0 + has_damping = damping != 0.0 + + if not has_stiffness: + qfrc_spring_out[worldid, dofid] = 0.0 + + if not has_damping: + qfrc_damper_out[worldid, dofid] = 0.0 + + if not (has_stiffness or has_damping): + return + + jnttype = jnt_type[jntid] + qposid = jnt_qposadr[jntid] + + if jnttype == wp.static(JointType.FREE.value): + # spring + if has_stiffness: + dif = wp.vec3( + qpos_in[worldid, qposid + 0] - qpos_spring[worldid, qposid + 0], + qpos_in[worldid, qposid + 1] - qpos_spring[worldid, qposid + 1], + qpos_in[worldid, qposid + 2] - qpos_spring[worldid, qposid + 2], + ) + qfrc_spring_out[worldid, dofid + 0] = -stiffness * dif[0] + qfrc_spring_out[worldid, dofid + 1] = -stiffness * dif[1] + qfrc_spring_out[worldid, dofid + 2] = -stiffness * dif[2] + rot = wp.quat( + qpos_in[worldid, qposid + 3], + qpos_in[worldid, qposid + 4], + qpos_in[worldid, qposid + 5], + qpos_in[worldid, qposid + 6], + ) + rot = wp.normalize(rot) + ref = wp.quat( + qpos_spring[worldid, qposid + 3], + qpos_spring[worldid, qposid + 4], + qpos_spring[worldid, qposid + 5], + qpos_spring[worldid, qposid + 6], + ) + dif = math.quat_sub(rot, ref) + qfrc_spring_out[worldid, dofid + 3] = -stiffness * dif[0] + qfrc_spring_out[worldid, dofid + 4] = -stiffness * dif[1] + qfrc_spring_out[worldid, dofid + 5] = -stiffness * dif[2] + + # damper + if has_damping: + qfrc_damper_out[worldid, dofid + 0] = -damping * qvel_in[worldid, dofid + 0] + qfrc_damper_out[worldid, dofid + 1] = -damping * qvel_in[worldid, dofid + 1] + qfrc_damper_out[worldid, dofid + 2] = -damping * qvel_in[worldid, dofid + 2] + qfrc_damper_out[worldid, dofid + 3] = -damping * qvel_in[worldid, dofid + 3] + qfrc_damper_out[worldid, dofid + 4] = -damping * qvel_in[worldid, dofid + 4] + qfrc_damper_out[worldid, dofid + 5] = -damping * qvel_in[worldid, dofid + 5] + elif jnttype == wp.static(JointType.BALL.value): + # spring + if has_stiffness: + rot = wp.quat( + qpos_in[worldid, qposid + 0], + qpos_in[worldid, qposid + 1], + qpos_in[worldid, qposid + 2], + qpos_in[worldid, qposid + 3], + ) + rot = wp.normalize(rot) + ref = wp.quat( + qpos_spring[worldid, qposid + 0], + qpos_spring[worldid, qposid + 1], + qpos_spring[worldid, qposid + 2], + qpos_spring[worldid, qposid + 3], + ) + dif = math.quat_sub(rot, ref) + qfrc_spring_out[worldid, dofid + 0] = -stiffness * dif[0] + qfrc_spring_out[worldid, dofid + 1] = -stiffness * dif[1] + qfrc_spring_out[worldid, dofid + 2] = -stiffness * dif[2] + + # damper + if has_damping: + qfrc_damper_out[worldid, dofid + 0] = -damping * qvel_in[worldid, dofid + 0] + qfrc_damper_out[worldid, dofid + 1] = -damping * qvel_in[worldid, dofid + 1] + qfrc_damper_out[worldid, dofid + 2] = -damping * qvel_in[worldid, dofid + 2] + else: # mjJNT_SLIDE, mjJNT_HINGE + # spring + if has_stiffness: + fdif = qpos_in[worldid, qposid] - qpos_spring[worldid, qposid] + qfrc_spring_out[worldid, dofid] = -stiffness * fdif + + # damper + if has_damping: + qfrc_damper_out[worldid, dofid] = -damping * qvel_in[worldid, dofid] + + +@wp.kernel +def _spring_damper_tendon_passive( + # Model: + tendon_stiffness: wp.array2d(dtype=float), + tendon_damping: wp.array2d(dtype=float), + tendon_lengthspring: wp.array2d(dtype=wp.vec2), + # Data in: + ten_velocity_in: wp.array2d(dtype=float), + ten_length_in: wp.array2d(dtype=float), + ten_J_in: wp.array3d(dtype=float), + # Data out: + qfrc_spring_out: wp.array2d(dtype=float), + qfrc_damper_out: wp.array2d(dtype=float), +): + worldid, tenid, dofid = wp.tid() + + stiffness = tendon_stiffness[worldid, tenid] + damping = tendon_damping[worldid, tenid] + + if stiffness == 0.0 and damping == 0.0: + return + + J = ten_J_in[worldid, tenid, dofid] + + if stiffness: + # compute spring force along tendon + length = ten_length_in[worldid, tenid] + lengthspring = tendon_lengthspring[worldid, tenid] + lower = lengthspring[0] + upper = lengthspring[1] + + if length > upper: + frc_spring = stiffness * (upper - length) + elif length < lower: + frc_spring = stiffness * (lower - length) + else: + frc_spring = 0.0 + + # transform to joint torque + wp.atomic_add(qfrc_spring_out[worldid], dofid, J * frc_spring) + + if damping: + # compute damper linear force along tendon + frc_damper = -damping * ten_velocity_in[worldid, tenid] + + # transform to joint torque + wp.atomic_add(qfrc_damper_out[worldid], dofid, J * frc_damper) + + +@wp.kernel +def _gravity_force( + # Model: + opt_gravity: wp.array(dtype=wp.vec3), + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + body_mass: wp.array2d(dtype=float), + body_gravcomp: wp.array2d(dtype=float), + dof_bodyid: wp.array(dtype=int), + # Data in: + xipos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + # Data out: + qfrc_gravcomp_out: wp.array2d(dtype=float), +): + worldid, bodyid, dofid = wp.tid() + bodyid += 1 # skip world body + gravcomp = body_gravcomp[worldid, bodyid] + gravity = opt_gravity[worldid] + + if gravcomp: + force = -gravity * body_mass[worldid, bodyid] * gravcomp + + pos = xipos_in[worldid, bodyid] + jac, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos, bodyid, dofid, worldid) + + wp.atomic_add(qfrc_gravcomp_out[worldid], dofid, wp.dot(jac, force)) + + +@wp.kernel +def _box_fluid( + # Model: + opt_wind: wp.array(dtype=wp.vec3), + opt_density: wp.array(dtype=float), + opt_viscosity: wp.array(dtype=float), + body_rootid: wp.array(dtype=int), + body_mass: wp.array2d(dtype=float), + body_inertia: wp.array2d(dtype=wp.vec3), + # Data in: + xipos_in: wp.array2d(dtype=wp.vec3), + ximat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + # Data out: + fluid_applied_out: wp.array2d(dtype=wp.spatial_vector), +): + """Fluid forces based on inertia-box approximation.""" + + worldid, bodyid = wp.tid() + wind = opt_wind[worldid] + density = opt_density[worldid] + viscosity = opt_viscosity[worldid] + + # map from CoM-centered to local body-centered 6D velocity + + # body-inertial + pos = xipos_in[worldid, bodyid] + rot = ximat_in[worldid, bodyid] + rotT = wp.transpose(rot) + + # transform velocity + cvel = cvel_in[worldid, bodyid] + torque = wp.spatial_top(cvel) + force = wp.spatial_bottom(cvel) + subtree_com = subtree_com_in[worldid, body_rootid[bodyid]] + dif = pos - subtree_com + force -= wp.cross(dif, torque) + + lvel_torque = rotT @ torque + lvel_force = rotT @ force + + if wind[0] or wind[1] or wind[2]: + # subtract translational component from body velocity + lvel_force -= rotT @ wind + + lfrc_torque = wp.vec3(0.0) + lfrc_force = wp.vec3(0.0) + + has_viscosity = viscosity > 0.0 + has_density = density > 0.0 + + if has_viscosity or has_density: + inertia = body_inertia[worldid, bodyid] + mass = body_mass[worldid, bodyid] + scl = 6.0 / mass + box0 = wp.sqrt(wp.max(MJ_MINVAL, inertia[1] + inertia[2] - inertia[0]) * scl) + box1 = wp.sqrt(wp.max(MJ_MINVAL, inertia[0] + inertia[2] - inertia[1]) * scl) + box2 = wp.sqrt(wp.max(MJ_MINVAL, inertia[0] + inertia[1] - inertia[2]) * scl) + + if has_viscosity: + # diameter of sphere approximation + diam = (box0 + box1 + box2) / 3.0 + + # angular viscosity + lfrc_torque = -lvel_torque * wp.pow(diam, 3.0) * wp.pi * viscosity + + # linear viscosity + lfrc_force = -3.0 * lvel_force * diam * wp.pi * viscosity + + if has_density: + # force + lfrc_force -= wp.vec3( + 0.5 * density * box1 * box2 * wp.abs(lvel_force[0]) * lvel_force[0], + 0.5 * density * box0 * box2 * wp.abs(lvel_force[1]) * lvel_force[1], + 0.5 * density * box0 * box1 * wp.abs(lvel_force[2]) * lvel_force[2], + ) + + # torque + scl = density / 64.0 + box0_pow4 = wp.pow(box0, 4.0) + box1_pow4 = wp.pow(box1, 4.0) + box2_pow4 = wp.pow(box2, 4.0) + lfrc_torque -= wp.vec3( + box0 * (box1_pow4 + box2_pow4) * wp.abs(lvel_torque[0]) * lvel_torque[0] * scl, + box1 * (box0_pow4 + box2_pow4) * wp.abs(lvel_torque[1]) * lvel_torque[1] * scl, + box2 * (box0_pow4 + box1_pow4) * wp.abs(lvel_torque[2]) * lvel_torque[2] * scl, + ) + + # rotate to global orientation: lfrc -> bfrc + torque = rot @ lfrc_torque + force = rot @ lfrc_force + + fluid_applied_out[worldid, bodyid] = wp.spatial_vector(force, torque) + + +def _fluid(m: Model, d: Data): + wp.launch( + _box_fluid, + dim=(d.nworld, m.nbody), + inputs=[ + m.opt.wind, + m.opt.density, + m.opt.viscosity, + m.body_rootid, + m.body_mass, + m.body_inertia, + d.xipos, + d.ximat, + d.subtree_com, + d.cvel, + ], + outputs=[ + d.fluid_applied, + ], + ) + + # TODO(team): ellipsoid fluid model + + support.apply_ft(m, d, d.fluid_applied, d.qfrc_fluid, False) + + +@wp.kernel +def _qfrc_passive( + # Model: + opt_has_fluid: bool, + jnt_actgravcomp: wp.array(dtype=int), + dof_jntid: wp.array(dtype=int), + # Data in: + qfrc_spring_in: wp.array2d(dtype=float), + qfrc_damper_in: wp.array2d(dtype=float), + qfrc_gravcomp_in: wp.array2d(dtype=float), + qfrc_fluid_in: wp.array2d(dtype=float), + # In: + gravcomp: bool, + # Data out: + qfrc_passive_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + qfrc_passive = qfrc_spring_in[worldid, dofid] + qfrc_passive += qfrc_damper_in[worldid, dofid] + + # add gravcomp unless added by actuators + if gravcomp and not jnt_actgravcomp[dof_jntid[dofid]]: + qfrc_passive += qfrc_gravcomp_in[worldid, dofid] + + # add fluid force + if opt_has_fluid: + qfrc_passive += qfrc_fluid_in[worldid, dofid] + + qfrc_passive_out[worldid, dofid] = qfrc_passive + + +@wp.kernel +def _flex_elasticity( + # Model: + opt_timestep: wp.array(dtype=float), + body_dofadr: wp.array(dtype=int), + flex_dim: wp.array(dtype=int), + flex_vertadr: wp.array(dtype=int), + flex_edgeadr: wp.array(dtype=int), + flex_elemedgeadr: wp.array(dtype=int), + flex_vertbodyid: wp.array(dtype=int), + flex_elem: wp.array(dtype=int), + flex_elemedge: wp.array(dtype=int), + flexedge_length0: wp.array(dtype=float), + flex_stiffness: wp.array(dtype=float), + flex_damping: wp.array(dtype=float), + # Data in: + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), + flexedge_length_in: wp.array2d(dtype=float), + flexedge_velocity_in: wp.array2d(dtype=float), + # Data out: + qfrc_spring_out: wp.array2d(dtype=float), +): + worldid, elemid = wp.tid() + timestep = opt_timestep[worldid] + f = 0 # TODO(quaglino): this should become a function of t + + dim = flex_dim[f] + nvert = dim + 1 + nedge = nvert * (nvert - 1) / 2 + edges = wp.where( + dim == 3, + wp.mat(0, 1, 1, 2, 2, 0, 2, 3, 0, 3, 1, 3, shape=(6, 2), dtype=int), + wp.mat(1, 2, 2, 0, 0, 1, 0, 0, 0, 0, 0, 0, shape=(6, 2), dtype=int), + ) + kD = flex_damping[f] / timestep + + gradient = wp.mat(0.0, shape=(6, 6)) + for e in range(nedge): + vert0 = flex_elem[(dim + 1) * elemid + edges[e, 0]] + vert1 = flex_elem[(dim + 1) * elemid + edges[e, 1]] + xpos0 = flexvert_xpos_in[worldid, vert0] + xpos1 = flexvert_xpos_in[worldid, vert1] + for i in range(3): + gradient[e, 0 + i] = xpos0[i] - xpos1[i] + gradient[e, 3 + i] = xpos1[i] - xpos0[i] + + elongation = wp.spatial_vectorf(0.0) + for e in range(nedge): + idx = flex_elemedge[flex_elemedgeadr[f] + elemid * nedge + e] + vel = flexedge_velocity_in[worldid, flex_edgeadr[f] + idx] + deformed = flexedge_length_in[worldid, flex_edgeadr[f] + idx] + reference = flexedge_length0[flex_edgeadr[f] + idx] + previous = deformed - vel * timestep + elongation[e] = deformed * deformed - reference * reference + (deformed * deformed - previous * previous) * kD + + metric = wp.mat(0.0, shape=(6, 6)) + id = int(0) + for ed1 in range(nedge): + for ed2 in range(ed1, nedge): + metric[ed1, ed2] = flex_stiffness[21 * elemid + id] + metric[ed2, ed1] = flex_stiffness[21 * elemid + id] + id += 1 + + force = wp.mat(0.0, shape=(6, 3)) + for ed1 in range(nedge): + for ed2 in range(nedge): + for i in range(2): + for x in range(3): + force[edges[ed2, i], x] -= elongation[ed1] * gradient[ed2, 3 * i + x] * metric[ed1, ed2] + + for v in range(nvert): + vert = flex_elem[(dim + 1) * elemid + v] + bodyid = flex_vertbodyid[flex_vertadr[f] + vert] + for x in range(3): + wp.atomic_add(qfrc_spring_out, worldid, body_dofadr[bodyid] + x, force[v, x]) + + +@wp.kernel +def _flex_bending( + # Model: + body_dofadr: wp.array(dtype=int), + flex_dim: wp.array(dtype=int), + flex_vertadr: wp.array(dtype=int), + flex_edgeadr: wp.array(dtype=int), + flex_vertbodyid: wp.array(dtype=int), + flex_edge: wp.array(dtype=wp.vec2i), + flex_edgeflap: wp.array(dtype=wp.vec2i), + flex_bending: wp.array(dtype=wp.mat44f), + # Data in: + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), + # Data out: + qfrc_spring_out: wp.array2d(dtype=float), +): + worldid, edgeid = wp.tid() + nvert = 4 + f = 0 # TODO(quaglino): this should become a function of t + + if flex_dim[f] != 2: + return + + v = wp.vec4i( + flex_edge[edgeid + flex_edgeadr[f]][0], + flex_edge[edgeid + flex_edgeadr[f]][1], + flex_edgeflap[edgeid + flex_edgeadr[f]][0], + flex_edgeflap[edgeid + flex_edgeadr[f]][1], + ) + + if v[3] == -1: + return + + force = wp.mat(0.0, shape=(nvert, 3)) + for i in range(nvert): + for j in range(nvert): + for x in range(3): + force[i, x] -= flex_bending[edgeid][i, j] * flexvert_xpos_in[worldid, v[j]][x] + + for i in range(nvert): + bodyid = flex_vertbodyid[flex_vertadr[f] + v[i]] + for x in range(3): + wp.atomic_add(qfrc_spring_out, worldid, body_dofadr[bodyid] + x, force[i, x]) + + +@event_scope +def passive(m: Model, d: Data): + """Adds all passive forces.""" + + if m.opt.disableflags & DisableBit.PASSIVE: + d.qfrc_spring.zero_() + d.qfrc_damper.zero_() + d.qfrc_gravcomp.zero_() + d.qfrc_fluid.zero_() + d.qfrc_passive.zero_() + return + + wp.launch( + _spring_damper_dof_passive, + dim=(d.nworld, m.njnt), + inputs=[ + m.qpos_spring, + m.jnt_type, + m.jnt_qposadr, + m.jnt_dofadr, + m.jnt_stiffness, + m.dof_damping, + d.qpos, + d.qvel, + ], + outputs=[d.qfrc_spring, d.qfrc_damper], + ) + + if m.ntendon: + wp.launch( + _spring_damper_tendon_passive, + dim=(d.nworld, m.ntendon, m.nv), + inputs=[ + m.tendon_stiffness, + m.tendon_damping, + m.tendon_lengthspring, + d.ten_velocity, + d.ten_length, + d.ten_J, + ], + outputs=[ + d.qfrc_spring, + d.qfrc_damper, + ], + ) + + wp.launch( + _flex_elasticity, + dim=(d.nworld, m.nflexelem), + inputs=[ + m.opt.timestep, + m.body_dofadr, + m.flex_dim, + m.flex_vertadr, + m.flex_edgeadr, + m.flex_elemedgeadr, + m.flex_vertbodyid, + m.flex_elem, + m.flex_elemedge, + m.flexedge_length0, + m.flex_stiffness, + m.flex_damping, + d.flexvert_xpos, + d.flexedge_length, + d.flexedge_velocity, + ], + outputs=[d.qfrc_spring], + ) + wp.launch( + _flex_bending, + dim=(d.nworld, m.nflexedge), + inputs=[ + m.body_dofadr, + m.flex_dim, + m.flex_vertadr, + m.flex_edgeadr, + m.flex_vertbodyid, + m.flex_edge, + m.flex_edgeflap, + m.flex_bending, + d.flexvert_xpos, + ], + outputs=[d.qfrc_spring], + ) + + gravcomp = m.ngravcomp and not (m.opt.disableflags & DisableBit.GRAVITY) + + if gravcomp: + d.qfrc_gravcomp.zero_() + wp.launch( + _gravity_force, + dim=(d.nworld, m.nbody - 1, m.nv), + inputs=[ + m.opt.gravity, + m.body_parentid, + m.body_rootid, + m.body_mass, + m.body_gravcomp, + m.dof_bodyid, + d.xipos, + d.subtree_com, + d.cdof, + ], + outputs=[d.qfrc_gravcomp], + ) + + if m.opt.has_fluid: + _fluid(m, d) + + wp.launch( + _qfrc_passive, + dim=(d.nworld, m.nv), + inputs=[ + m.opt.has_fluid, + m.jnt_actgravcomp, + m.dof_jntid, + d.qfrc_spring, + d.qfrc_damper, + d.qfrc_gravcomp, + d.qfrc_fluid, + gravcomp, + ], + outputs=[ + d.qfrc_passive, + ], + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive_test.py new file mode 100644 index 00000000..f914d268 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive_test.py @@ -0,0 +1,144 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Tests for passive force functions.""" + +import numpy as np +import warp as wp +from absl.testing import absltest +from absl.testing import parameterized + +import mujoco_warp as mjwarp + +from mujoco.mjx.third_party.mujoco_warp._src import test_util + +# tolerance for difference between MuJoCo and MJWarp passive force calculations - mostly +# due to float precision +_TOLERANCE = 5e-5 + + +def _assert_eq(a, b, name): + tol = _TOLERANCE * 10 # avoid test noise + err_msg = f"mismatch: {name}" + np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol) + + +class PassiveTest(parameterized.TestCase): + @parameterized.product(passive=[True, False], gravity=[True, False]) + def test_passive(self, passive, gravity): + """Tests passive.""" + _, mjd, m, d = test_util.fixture("pendula.xml", passive=passive, gravity=gravity, kick=True, applied=True) + + for arr in (d.qfrc_spring, d.qfrc_damper, d.qfrc_gravcomp, d.qfrc_passive): + arr.zero_() + + mjwarp.passive(m, d) + + _assert_eq(d.qfrc_spring.numpy()[0], mjd.qfrc_spring, "qfrc_spring") + _assert_eq(d.qfrc_damper.numpy()[0], mjd.qfrc_damper, "qfrc_damper") + _assert_eq(d.qfrc_gravcomp.numpy()[0], mjd.qfrc_gravcomp, "qfrc_gravcomp") + _assert_eq(d.qfrc_passive.numpy()[0], mjd.qfrc_passive, "qfrc_passive") + + @parameterized.parameters( + (1, 0, 0, 0, 0), + (0, 1, 0, 0, 0), + (0, 0, 1, 0, 0), + (0, 0, 0, 1, 0), + (0, 0, 0, 0, 1), + (1, 1, 1, 1, 1), + ) + def test_fluid(self, density, viscosity, wind0, wind1, wind2): + """Tests fluid model.""" + + _, mjd, m, d = test_util.fixture( + xml=f""" + + + """, + keyframe=0, + ) + + for arr in (d.qfrc_passive, d.qfrc_fluid): + arr.zero_() + + mjwarp.passive(m, d) + + _assert_eq(d.qfrc_passive.numpy()[0], mjd.qfrc_passive, "qfrc_passive") + _assert_eq(d.qfrc_fluid.numpy()[0], mjd.qfrc_fluid, "qfrc_fluid") + + @parameterized.parameters((True, True), (True, False), (False, True), (False, False)) + def test_gravcomp(self, sparse, gravity): + """Tests gravity compensation.""" + + _, mjd, m, d = test_util.fixture( + xml=""" + + + + + + + + + + + + + + + + + + + + + + + + + + + + """, + gravity=gravity, + sparse=sparse, + ) + + for arr in (d.qfrc_passive, d.qfrc_gravcomp, d.qfrc_actuator): + arr.zero_() + + mjwarp.passive(m, d) + mjwarp.fwd_actuation(m, d) + + _assert_eq(d.qfrc_passive.numpy()[0], mjd.qfrc_passive, "qfrc_passive") + _assert_eq(d.qfrc_gravcomp.numpy()[0], mjd.qfrc_gravcomp, "qfrc_gravcomp") + _assert_eq(d.qfrc_actuator.numpy()[0], mjd.qfrc_actuator, "qfrc_actuator") + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py new file mode 100644 index 00000000..114ce4a8 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py @@ -0,0 +1,903 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +from typing import Tuple + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import vec6 + + +@wp.func +def _ray_map(pos: wp.vec3, mat: wp.mat33, pnt: wp.vec3, vec: wp.vec3) -> Tuple[wp.vec3, wp.vec3]: + """Maps ray to local geom frame coordinates. + + Args: + pos: position of geom frame + mat: orientation of geom frame + pnt: starting point of ray in world coordinates + vec: direction of ray in world coordinates + + Returns: + 3D point and 3D direction in local geom frame + """ + + matT = wp.transpose(mat) + lpnt = matT @ (pnt - pos) + lvec = matT @ vec + + return lpnt, lvec + + +@wp.func +def _ray_eliminate( + # Model: + body_weldid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + geom_group: wp.array(dtype=int), + geom_matid: wp.array(dtype=int), # kernel_analyzer: ignore + geom_rgba: wp.array(dtype=wp.vec4), # kernel_analyzer: ignore + mat_rgba: wp.array(dtype=wp.vec4), # kernel_analyzer: ignore + # In: + geomid: int, + geomgroup: vec6, + flg_static: bool, + bodyexclude: int, +) -> bool: + """Eliminate ray.""" + bodyid = geom_bodyid[geomid] + matid = geom_matid[geomid] + + # body exclusion + if bodyid == bodyexclude: + return True + + # invisible geom exclusion + if matid < 0 and geom_rgba[geomid][3] == 0.0: + return True + + # invisible material exclusion + if matid >= 0: + if mat_rgba[matid][3] == 0.0: + return True + + # static exclusion + if not flg_static and body_weldid[bodyid] == 0: + return True + + # no geomgroup inclusion + if ( + geomgroup[0] == -1 + and geomgroup[1] == -1 + and geomgroup[2] == -1 + and geomgroup[3] == -1 + and geomgroup[4] == -1 + and geomgroup[5] == -1 + ): + return False + + # group inclusion/exclusion + groupid = wp.min(5, wp.max(0, geom_group[geomid])) + + return geomgroup[groupid] == 0 + + +@wp.func +def _ray_quad(a: float, b: float, c: float) -> Tuple[float, wp.vec2]: + """Compute solutions from quadratic: a*x^2 + 2*b*x + c = 0.""" + det = b * b - a * c + if det < MJ_MINVAL: + return wp.inf, wp.vec2(wp.inf, wp.inf) + det = wp.sqrt(det) + + # compute the two solutions + den = 1.0 / a + x0 = (-b - det) * den + x1 = (-b + det) * den + x = wp.vec2(x0, x1) + + # finalize result + if x0 >= 0.0: + return x0, x + elif x1 >= 0.0: + return x1, x + else: + return wp.inf, x + + +@wp.func +def _ray_triangle(v0: wp.vec3, v1: wp.vec3, v2: wp.vec3, pnt: wp.vec3, vec: wp.vec3, b0: wp.vec3, b1: wp.vec3) -> float: + """Returns the distance at which a ray intersects with a triangle.""" + dif0 = v0 - pnt + dif1 = v1 - pnt + dif2 = v2 - pnt + + # project difference vectors in normal plane + planar_00 = wp.dot(dif0, b0) + planar_01 = wp.dot(dif0, b1) + planar_10 = wp.dot(dif1, b0) + planar_11 = wp.dot(dif1, b1) + planar_20 = wp.dot(dif2, b0) + planar_21 = wp.dot(dif2, b1) + + # reject if on the same side of any coordinate axis + if ( + (planar_00 > 0.0 and planar_10 > 0.0 and planar_20 > 0.0) + or (planar_00 < 0.0 and planar_10 < 0.0 and planar_20 < 0.0) + or (planar_01 > 0.0 and planar_11 > 0.0 and planar_21 > 0.0) + or (planar_01 < 0.0 and planar_11 < 0.0 and planar_21 < 0.0) + ): + return float(wp.inf) + + # determine if origin is inside planar projection of triangle + # A = (p0-p2, p1-p2), b = -p2, solve A*t = b + A00 = planar_00 - planar_20 + A10 = planar_10 - planar_20 + A01 = planar_01 - planar_21 + A11 = planar_11 - planar_21 + + b = wp.vec2(-planar_20, -planar_21) + + det = A00 * A11 - A10 * A01 + if wp.abs(det) < MJ_MINVAL: + return float(wp.inf) + + t0 = (A11 * b[0] - A10 * b[1]) / det + t1 = (-A01 * b[0] + A00 * b[1]) / det + + # check if outside + if t0 < 0.0 or t1 < 0.0 or t0 + t1 > 1.0: + return float(wp.inf) + + # intersect ray with plane of triangle + dif0 = v0 - v2 + dif1 = v1 - v2 + dif2 = pnt - v2 + nrm = wp.cross(dif0, dif1) # normal to triangle plane + denom = wp.dot(vec, nrm) + if wp.abs(denom) < MJ_MINVAL: + return float(wp.inf) + + dist = -wp.dot(dif2, nrm) / denom + return wp.where(dist >= 0.0, dist, float(wp.inf)) + + +@wp.func +def _ray_plane(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float: + """Returns the distance at which a ray intersects with a plane.""" + + # map to local frame + lpnt, lvec = _ray_map(pos, mat, pnt, vec) + + # z-vec not pointing towards front face: reject + if lvec[2] > -MJ_MINVAL: + return wp.inf + + # intersection with plane + x = -lpnt[2] / lvec[2] + if x < 0.0: + return wp.inf + + p = wp.vec2( + lpnt[0] + x * lvec[0], + lpnt[1] + x * lvec[1], + ) + + # accept only within rendered rectangle + if (size[0] <= 0.0 or wp.abs(p[0]) <= size[0]) and (size[1] <= 0.0 or wp.abs(p[1]) <= size[1]): + return x + else: + return wp.inf + + +@wp.func +def _ray_sphere(pos: wp.vec3, dist_sqr: float, pnt: wp.vec3, vec: wp.vec3) -> float: + """Returns the distance at which a ray intersects with a sphere.""" + dif = pnt - pos + + a = wp.dot(vec, vec) + b = wp.dot(vec, dif) + c = wp.dot(dif, dif) - dist_sqr + + sol, _ = _ray_quad(a, b, c) + return sol + + +@wp.func +def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float: + """Returns the distance at which a ray intersects with a capsule.""" + + # bounding sphere test + ssz = size[0] + size[1] + if _ray_sphere(pos, ssz * ssz, pnt, vec) < 0.0: + return wp.inf + + # map to local frame + lpnt, lvec = _ray_map(pos, mat, pnt, vec) + + # init solution + x = -1.0 + + # cylinder round side: (x * lvec + lpnt)' * (x * lvec + lpnt) = size[0] * size[0] + sq_size0 = size[0] * size[0] + a = lvec[0] * lvec[0] + lvec[1] * lvec[1] + b = lvec[0] * lpnt[0] + lvec[1] * lpnt[1] + c = lpnt[0] * lpnt[0] + lpnt[1] * lpnt[1] - sq_size0 + + # solve a * x^2 + 2 * b * x + c = 0 + sol, xx = _ray_quad(a, b, c) + + # make sure round solution is between flat sides + if sol >= 0.0 and wp.abs(lpnt[2] + sol * vec[2]) <= size[1]: + if x < 0.0 or sol < x: + x = sol + + # top cap + ldif = wp.vec3(lpnt[0], lpnt[1], lpnt[2] - size[1]) + a += lvec[2] * lvec[2] + b = wp.dot(lvec, ldif) + c = wp.dot(ldif, ldif) - sq_size0 + _, xx = _ray_quad(a, b, c) + + # accept only top half of sphere + for i in range(2): + if xx[i] >= 0.0 and lpnt[2] + xx[i] * lvec[2] >= size[1]: + if x < 0.0 or xx[i] < x: + x = xx[i] + + # bottom cap + ldif = wp.vec3(ldif[0], ldif[1], lpnt[2] + size[1]) + b = wp.dot(lvec, ldif) + c = wp.dot(ldif, ldif) - sq_size0 + _, xx = _ray_quad(a, b, c) + + # accept only bottom half of sphere + for i in range(2): + if xx[i] >= 0.0 and lpnt[2] + xx[i] * lvec[2] <= -size[1]: + if x < 0.0 or xx[i] < x: + x = xx[i] + + return x + + +@wp.func +def _ray_ellipsoid(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float: + """Returns the distance at which a ray intersects with an ellipsoid.""" + + # map to local frame + lpnt, lvec = _ray_map(pos, mat, pnt, vec) + + # invert size^2 + s = wp.vec3( + 1.0 / (size[0] * size[0]), + 1.0 / (size[1] * size[1]), + 1.0 / (size[2] * size[2]), + ) + + # (x * lvec + lpnt)' * diag(1 / size^2) * (x * lvec + lpnt) = 1 + slvec = wp.cw_mul(s, lvec) + a = wp.dot(slvec, lvec) + b = wp.dot(slvec, lpnt) + c = wp.dot(wp.cw_mul(s, lpnt), lpnt) - 1.0 + + # solve a * x^2 + 2 * b * x + c = 0 + sol, _ = _ray_quad(a, b, c) + return sol + + +@wp.func +def _ray_cylinder(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float: + """Returns the distance at which a ray intersects with a cylinder.""" + # bounding sphere test + ssz = size[0] * size[0] + size[1] * size[1] + if _ray_sphere(pos, ssz, pnt, vec) < 0.0: + return wp.inf + + # map to local frame + lpnt, lvec = _ray_map(pos, mat, pnt, vec) + + # init solution + x = wp.inf + + # flat sides + if wp.abs(lvec[2]) > MJ_MINVAL: + for side in range(-1, 2, 2): + # solution of: lpnt[2] + x * lvec[2] = side * height_size + sol = (float(side) * size[1] - lpnt[2]) / lvec[2] + + # process if non-negative + if sol >= 0.0: + # intersection with horizontal face + p = wp.vec2( + lpnt[0] + sol * lvec[0], + lpnt[1] + sol * lvec[1], + ) + + # accept within radius + if wp.dot(p, p) <= size[0] * size[0]: + if x < 0.0 or sol < x: + x = sol + + # (x * lvec + lpnt)' * (x * lvec + lpnt) = size[0] * size[0] + a = lvec[0] * lvec[0] + lvec[1] * lvec[1] + b = lvec[0] * lpnt[0] + lvec[1] * lpnt[1] + c = lpnt[0] * lpnt[0] + lpnt[1] * lpnt[1] - size[0] * size[0] + + # solve a * x^2 + 2 * b * x + c = 0 + sol, _ = _ray_quad(a, b, c) + + # make sure round solution is between flat sides + if sol >= 0.0 and wp.abs(lpnt[2] + sol * lvec[2]) <= size[1]: + if x < 0.0 or sol < x: + x = sol + + return x + + +_IFACE = wp.types.matrix((3, 2), dtype=int)(1, 2, 0, 2, 0, 1) + + +@wp.func +def _ray_box(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, vec6]: + """Returns the distance at which a ray intersects with a box.""" + all = vec6(-1.0, -1.0, -1.0, -1.0, -1.0, -1.0) + + # bounding sphere test + ssz = wp.dot(size, size) + if _ray_sphere(pos, ssz, pnt, vec) < 0.0: + return wp.inf, all + + # map to local frame + lpnt, lvec = _ray_map(pos, mat, pnt, vec) + + # init solution + x = wp.inf + + # loop over axes with non-zero vec + for i in range(3): + if wp.abs(lvec[i]) > MJ_MINVAL: + for side in range(-1, 2, 2): + # solution of: lpnt[i] + x * lvec[i] = side * size[i] + sol = (float(side) * size[i] - lpnt[i]) / lvec[i] + + # process if non-negative + if sol >= 0.0: + id0 = _IFACE[i][0] + id1 = _IFACE[i][1] + + # intersection with face + p0 = lpnt[id0] + sol * lvec[id0] + p1 = lpnt[id1] + sol * lvec[id1] + + # accept within rectangle + if (wp.abs(p0) <= size[id0]) and (wp.abs(p1) <= size[id1]): + # update + if (x < 0.0) or (sol < x): + x = sol + + # save in all + all[2 * i + (side + 1) / 2] = sol + + return x, all + + +@wp.func +def _ray_hfield( + # Model: + geom_type: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + hfield_adr: wp.array(dtype=int), + hfield_nrow: wp.array(dtype=int), + hfield_ncol: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + hfield_data: wp.array(dtype=float), + # In: + pos: wp.vec3, + mat: wp.mat33, + pnt: wp.vec3, + vec: wp.vec3, + id: int, +): + # check geom type + if geom_type[id] != int(GeomType.HFIELD.value): + return wp.inf + + # hfield id and dimensions + hid = geom_dataid[id] + nrow = hfield_nrow[hid] + ncol = hfield_ncol[hid] + + size = hfield_size[hid] + adr = hfield_adr[hid] + + mat_col = wp.vec3(mat[0, 2], mat[1, 2], mat[2, 2]) + + # compute size and pos of base box + base_scale = size[3] * 0.5 + base_size = wp.vec3(size[0], size[1], base_scale) + base_pos = pos + mat_col * base_scale + + # compute size and pos of top box + top_scale = size[2] * 0.5 + top_size = wp.vec3(size[0], size[1], top_scale) + top_pos = pos + mat_col * top_scale + + # init: intersection with base box + x, _ = _ray_box(base_pos, mat, base_size, pnt, vec) + + # check top box: done if no intersection + top_intersect, all = _ray_box(top_pos, mat, top_size, pnt, vec) + + if top_intersect < 0.0: + return x + + # map to local frame + lpnt, lvec = _ray_map(pos, mat, pnt, vec) + + # construct basis vectors of normal plane + b0 = wp.vec3(1.0, 1.0, 1.0) + + if wp.abs(lvec[0]) >= wp.abs(lvec[1]) and wp.abs(lvec[0]) >= wp.abs(lvec[2]): + b0[0] = 0.0 + elif wp.abs(lvec[1]) >= wp.abs(lvec[2]): + b0[1] = 0.0 + else: + b0[2] = 0.0 + b1 = b0 + lvec * -wp.dot(lvec, b0) / wp.dot(lvec, lvec) + b1 = wp.normalize(b1) + + b2 = wp.cross(b1, lvec) + b2 = wp.normalize(b2) + + # find ray segment intersecting top box + seg = wp.vec2(0.0, top_intersect) + for i in range(6): + if all[i] > seg[1]: + seg[0] = top_intersect + seg[1] = all[i] + + # project segment endpoints in horizontal plane, discretize + dx = (2.0 * size[0]) / float(ncol - 1) + dy = (2.0 * size[1]) / float(nrow - 1) + SX = wp.vec2((lpnt[0] * seg[0] * lvec[0] + size[0]) / dx, (lpnt[0] * seg[1] * lvec[0] + size[0]) / dx) + SY = wp.vec2((lpnt[1] + seg[0] * lvec[1] + size[1]) / dy, (lpnt[1] + seg[1] * lvec[1] + size[1]) / dy) + + # compute ranges, with +1 padding + cmin = wp.max(0, int(wp.floor(wp.min(SX[0], SX[1])) - 1.0)) + cmax = wp.min(ncol - 1, int(wp.ceil(wp.max(SX[0], SX[1])) + 1.0)) + rmin = wp.max(0, int(wp.floor(wp.min(SY[0], SY[1])) - 1.0)) + rmax = wp.min(nrow - 1, int(wp.ceil(wp.max(SY[0], SY[1])) + 1.0)) + + # check triangles within bounds + for r in range(rmin, rmax): + for c in range(cmin, cmax): + # first triangle + v0 = wp.vec3(dx * float(c) - size[0], dy * float(r) - size[1], hfield_data[adr + r * ncol + c] * size[2]) + v1 = wp.vec3( + dx * float(c + 1) - size[0], dy * float(r + 1) - size[1], hfield_data[adr + (r + 1) * ncol + (c + 1)] * size[2] + ) + v2 = wp.vec3(dx * float(c + 1) - size[0], dy * float(r) - size[1], hfield_data[adr + r * ncol + (c + 1)] * size[2]) + sol = _ray_triangle(v0, v1, v2, pnt, vec, b0, b1) + if sol >= 0.0 and (x < 0.0 or sol < x): + x = sol + + # second triangle + v0 = wp.vec3(dx * float(c) - size[0], dy * float(r) - size[1], hfield_data[adr + r * ncol + c] * size[2]) + v1 = wp.vec3( + dx * float(c + 1) - size[0], dy * float(r + 1) - size[1], hfield_data[adr + (r + 1) * ncol + (c + 1)] * size[2] + ) + v2 = wp.vec3(dx * float(c) - size[0], dy * float(r + 1) - size[1], hfield_data[adr + (r + 1) * ncol + c] * size[2]) + sol = _ray_triangle(v0, v1, v2, pnt, vec, b0, b1) + if sol >= 0.0 and (x < 0.0 or sol < x): + x = sol + + # check viable sides of top box + for i in range(4): + if all[i] >= 0.0 and (all[i] < x or x < 0.0): + # normalized height of intersection point + z = (lpnt[2] + all[i] * lvec[2]) / size[2] + + # rectangle points: y, y0, z0, z1 + # side normal to x-axis + if i < 2: + y = (lpnt[1] + all[i] * lvec[1] + size[1]) / dy + y0 = wp.max(0.0, wp.min(float(nrow - 2), wp.floor(y))) + if i == 1: + z0 = hfield_data[adr + int(wp.round(y0)) * nrow + ncol - 1] + z1 = hfield_data[adr + int(wp.round(y0 + 1.0)) * nrow + ncol - 1] + else: + z0 = hfield_data[adr + int(wp.round(y0)) * nrow] + z1 = hfield_data[adr + int(wp.round(y0 + 1.0)) * nrow] + # side normal to y-axis + else: + y = (lpnt[0] + all[i] * lvec[0] + size[0]) / dx + y0 = wp.max(0.0, wp.min(float(ncol - 2), wp.floor(y))) + if i == 3: + z0 = hfield_data[adr + int(wp.round(y0)) + (nrow - 1) * ncol] + z1 = hfield_data[adr + int(wp.round(y0 + 1.0)) + (nrow - 1) * ncol] + else: + z0 = hfield_data[adr + int(wp.round(y0))] + z1 = hfield_data[adr + int(wp.round(y0 + 1.0))] + + # check if point is below line segments + if z < z0 * (y0 + 1.0 - y) + z1 * (y - y0): + x = all[i] + + return x + + +@wp.func +def _ray_mesh( + # Model: + nmeshface: int, + mesh_vertadr: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_faceadr: wp.array(dtype=int), + mesh_face: wp.array(dtype=wp.vec3i), + # In: + data_id: int, + pos: wp.vec3, + mat: wp.mat33, + pnt: wp.vec3, + vec: wp.vec3, +) -> float: + """Returns the distance and geomid for ray mesh intersections.""" + pnt, vec = _ray_map(pos, mat, pnt, vec) + + # compute orthogonal basis vectors + if wp.abs(vec[0]) < wp.abs(vec[1]): + if wp.abs(vec[0]) < wp.abs(vec[2]): + b0 = wp.vec3(0.0, vec[2], -vec[1]) + else: + b0 = wp.vec3(vec[1], -vec[0], 0.0) + else: + if wp.abs(vec[1]) < wp.abs(vec[2]): + b0 = wp.vec3(-vec[2], 0.0, vec[0]) + else: + b0 = wp.vec3(vec[1], -vec[0], 0.0) + + # normalize first vector + b0 = wp.normalize(b0) + + # compute second vector as cross product + b1 = wp.cross(vec, b0) + b1 = wp.normalize(b1) + + min_dist = float(wp.inf) + + # get mesh vertex data range + vert_start = mesh_vertadr[data_id] + + # get mesh face and vertex data + face_start = mesh_faceadr[data_id] + + if data_id + 1 < mesh_faceadr.shape[0]: + face_end = mesh_faceadr[data_id + 1] + else: + face_end = nmeshface + + # iterate through all faces + for i in range(face_start, face_end): + # get vertices for this face + v_idx = mesh_face[i] + + # create triangle struct + v0 = mesh_vert[vert_start + v_idx.x] + v1 = mesh_vert[vert_start + v_idx.y] + v2 = mesh_vert[vert_start + v_idx.z] + + # calculate intersection + dist = _ray_triangle(v0, v1, v2, pnt, vec, b0, b1) + if dist < min_dist: + min_dist = dist + + return min_dist + + +@wp.func +def ray_geom(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3, geomtype: int) -> float: + """Returns distance along ray to intersection with geom, or infinity if none.""" + + # TODO(team): static loop unrolling to remove unnecessary branching + if geomtype == int(GeomType.PLANE.value): + return _ray_plane(pos, mat, size, pnt, vec) + elif geomtype == int(GeomType.SPHERE.value): + return _ray_sphere(pos, size[0] * size[0], pnt, vec) + elif geomtype == int(GeomType.CAPSULE.value): + return _ray_capsule(pos, mat, size, pnt, vec) + elif geomtype == int(GeomType.ELLIPSOID.value): + return _ray_ellipsoid(pos, mat, size, pnt, vec) + elif geomtype == int(GeomType.CYLINDER.value): + return _ray_cylinder(pos, mat, size, pnt, vec) + elif geomtype == int(GeomType.BOX.value): + dist, _ = _ray_box(pos, mat, size, pnt, vec) + return dist + else: + return wp.inf + + +@wp.func +def _ray_geom_mesh( + # Model: + nmeshface: int, + body_weldid: wp.array(dtype=int), + geom_type: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_group: wp.array(dtype=int), + geom_matid: wp.array2d(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + geom_rgba: wp.array2d(dtype=wp.vec4), + hfield_adr: wp.array(dtype=int), + hfield_nrow: wp.array(dtype=int), + hfield_ncol: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + hfield_data: wp.array(dtype=float), + mesh_vertadr: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_faceadr: wp.array(dtype=int), + mesh_face: wp.array(dtype=wp.vec3i), + mat_rgba: wp.array2d(dtype=wp.vec4), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + worldid: int, + pnt: wp.vec3, + vec: wp.vec3, + geomgroup: vec6, + flg_static: bool, + bodyexclude: int, + geomid: int, +) -> float: + if not _ray_eliminate( + body_weldid, + geom_bodyid, + geom_group, + geom_matid[worldid], + geom_rgba[worldid], + mat_rgba[worldid], + geomid, + geomgroup, + flg_static, + bodyexclude, + ): + pos = geom_xpos_in[worldid, geomid] + mat = geom_xmat_in[worldid, geomid] + type = geom_type[geomid] + + if type == int(GeomType.MESH.value): + return _ray_mesh( + nmeshface, + mesh_vertadr, + mesh_vert, + mesh_faceadr, + mesh_face, + geom_dataid[geomid], + pos, + mat, + pnt, + vec, + ) + elif type == int(GeomType.HFIELD.value): + return _ray_hfield( + geom_type, + geom_dataid, + hfield_adr, + hfield_nrow, + hfield_ncol, + hfield_size, + hfield_data, + pos, + mat, + pnt, + vec, + geomid, + ) + else: + return ray_geom(pos, mat, geom_size[worldid, geomid], pnt, vec, type) + else: + return wp.inf + + +@wp.kernel +def _ray( + # Model: + ngeom: int, + nmeshface: int, + body_weldid: wp.array(dtype=int), + geom_type: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_group: wp.array(dtype=int), + geom_matid: wp.array2d(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + geom_rgba: wp.array2d(dtype=wp.vec4), + hfield_adr: wp.array(dtype=int), + hfield_nrow: wp.array(dtype=int), + hfield_ncol: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + hfield_data: wp.array(dtype=float), + mesh_vertadr: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_faceadr: wp.array(dtype=int), + mesh_face: wp.array(dtype=wp.vec3i), + mat_rgba: wp.array2d(dtype=wp.vec4), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + pnt: wp.array2d(dtype=wp.vec3), + vec: wp.array2d(dtype=wp.vec3), + geomgroup: vec6, + flg_static: bool, + bodyexclude: wp.array(dtype=int), + # Out: + dist_out: wp.array(dtype=float, ndim=2), + geomid_out: wp.array(dtype=int, ndim=2), +): + worldid, rayid, tid = wp.tid() + + num_threads = wp.block_dim() + + min_dist = float(wp.inf) + min_geomid = int(-1) + + upper = ((ngeom + num_threads - 1) // num_threads) * num_threads + for geomid in range(tid, upper, num_threads): + if geomid < ngeom: + dist = _ray_geom_mesh( + nmeshface, + body_weldid, + geom_type, + geom_bodyid, + geom_dataid, + geom_group, + geom_matid, + geom_size, + geom_rgba, + hfield_adr, + hfield_nrow, + hfield_ncol, + hfield_size, + hfield_data, + mesh_vertadr, + mesh_vert, + mesh_faceadr, + mesh_face, + mat_rgba, + geom_xpos_in, + geom_xmat_in, + worldid, + pnt[worldid, rayid], + vec[worldid, rayid], + geomgroup, + flg_static, + bodyexclude[rayid], + geomid, + ) + else: + dist = wp.inf + + tile_dist = wp.tile(dist) + local_min_geomid = wp.tile_argmin(tile_dist) + local_min_dist = tile_dist[local_min_geomid[0]] + + tile_geomid = wp.tile(geomid) + + if local_min_dist < min_dist: + min_dist = local_min_dist + min_geomid = tile_geomid[local_min_geomid[0]] + + if wp.isinf(min_dist): + dist_out[worldid, rayid] = -1.0 + else: + dist_out[worldid, rayid] = min_dist + geomid_out[worldid, rayid] = min_geomid + + +def ray( + m: Model, + d: Data, + pnt: wp.array2d(dtype=wp.vec3), + vec: wp.array2d(dtype=wp.vec3), + geomgroup: vec6 = None, + flg_static: bool = True, + bodyexclude: int = -1, +) -> tuple[wp.array2d(dtype=float), wp.array2d(dtype=int)]: + """Returns the distance at which rays intersect with primitive geoms. + + Args: + m (Model): The model containing kinematic and dynamic information (device). + d (Data): The data object containing the current state and output arrays (device). + pnt (wp.array2d(dtype=wp.vec3)): Ray origin points. + vec (wp.array2d(dtype=wp.vec3)): Ray directions. + geomgroup (vec6, optional): Group inclusion/exclusion mask. + If all are wp.inf, ignore. + flg_static (bool, optional): If True, allows rays to intersect with static geoms. + Defaults to True. + bodyexclude (int, optional): Ignore geoms on specified body id (-1 to disable). + Defaults to -1. + + Returns: + wp.array2d(dtype=float): Distances from ray origins to geom surfaces. + wp.array2d(dtype=int): IDs of intersected geoms (-1 if none). + """ + + assert pnt.shape[0] == vec.shape[0] + assert d.ray_dist.shape[1] == d.ray_geomid.shape[1] + assert pnt.shape[0] == d.ray_dist.shape[1] + + if geomgroup is None: + geomgroup = vec6(-1, -1, -1, -1, -1, -1) + + d.ray_bodyexclude.fill_(bodyexclude) + + rays(m, d, pnt, vec, geomgroup, flg_static, d.ray_bodyexclude, d.ray_dist, d.ray_geomid) + + return d.ray_dist, d.ray_geomid + + +def rays( + m: Model, + d: Data, + pnt: wp.array2d(dtype=wp.vec3), + vec: wp.array2d(dtype=wp.vec3), + geomgroup: vec6, + flg_static: bool, + bodyexclude: wp.array(dtype=int), + dist: wp.array2d(dtype=wp.vec3), + geomid: wp.array2d(dtype=int), +): + wp.launch_tiled( + _ray, + dim=(d.nworld, pnt.shape[1]), + inputs=[ + m.ngeom, + m.nmeshface, + m.body_weldid, + m.geom_type, + m.geom_bodyid, + m.geom_dataid, + m.geom_group, + m.geom_matid, + m.geom_size, + m.geom_rgba, + m.hfield_adr, + m.hfield_nrow, + m.hfield_ncol, + m.hfield_size, + m.hfield_data, + m.mesh_vertadr, + m.mesh_vert, + m.mesh_faceadr, + m.mesh_face, + m.mat_rgba, + d.geom_xpos, + d.geom_xmat, + pnt, + vec, + geomgroup, + flg_static, + bodyexclude, + dist, + geomid, + ], + block_dim=m.block_dim.ray, + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray_test.py new file mode 100644 index 00000000..a81dda7e --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray_test.py @@ -0,0 +1,312 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Tests for ray functions.""" + +import mujoco +import numpy as np +import warp as wp +from absl.testing import absltest + +import mujoco_warp as mjwarp + +from mujoco.mjx.third_party.mujoco_warp._src import test_util +from mujoco.mjx.third_party.mujoco_warp._src.types import vec6 + +# tolerance for difference between MuJoCo and MJX ray calculations - mostly +# due to float precision +_TOLERANCE = 5e-5 + + +def _assert_eq(a, b, name): + tol = _TOLERANCE * 10 # avoid test noise + err_msg = f"mismatch: {name}" + np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol) + + +class RayTest(absltest.TestCase): + def test_ray_nothing(self): + """Tests that ray returns -1 when nothing is hit.""" + mjm, mjd, m, d = test_util.fixture("ray.xml") + + pnt = wp.array([wp.vec3(12.146, 1.865, 3.895)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.vec3(0.0, 0.0, -1.0)], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] # Extract from [[-1]] + dist_np = dist.numpy()[0, 0] # Extract from [[-1.]] + _assert_eq(geomid_np, -1, "geom_id") + _assert_eq(dist_np, -1, "dist") + + def test_ray_plane(self): + """Tests ray<>plane matches MuJoCo.""" + mjm, mjd, m, d = test_util.fixture("ray.xml") + + # looking down at a slight angle + pnt = wp.array([wp.vec3(2.0, 1.0, 3.0)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, 0, "geom_id") + pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0] + unused = np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused) + _assert_eq(dist_np, mj_dist, "dist") + + # looking on wrong side of plane + pnt = wp.array([wp.vec3(0.0, 0.0, -0.5)], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, -1, "geom_id") + _assert_eq(dist_np, -1, "dist") + + def test_ray_sphere(self): + """Tests ray<>sphere matches MuJoCo.""" + mjm, mjd, m, d = test_util.fixture("ray.xml") + + # looking down at sphere at a slight angle + pnt = wp.array([wp.vec3(0.0, 0.0, 1.6)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, 1, "geom_id") + pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0] + unused = np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused) + _assert_eq(dist_np, mj_dist, "dist") + + def test_ray_capsule(self): + """Tests ray<>capsule matches MuJoCo.""" + mjm, mjd, m, d = test_util.fixture("ray.xml") + + # looking down at capsule at a slight angle + pnt = wp.array([wp.vec3(0.5, 1.0, 1.6)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(0.0, 0.05, -1.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, 2, "geom_id") + pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0] + unused = np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused) + _assert_eq(dist_np, mj_dist, "dist") + + # looking up at capsule from below + pnt = wp.array([wp.vec3(-0.5, 1.0, 0.05)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(0.0, 0.05, 1.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, 2, "geom_id") + pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0] + unused = np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused) + _assert_eq(dist_np, mj_dist, "dist") + + # looking at cylinder of capsule from the side + pnt = wp.array([wp.vec3(0.0, 1.0, 0.75)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(1.0, 0.0, 0.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, 2, "geom_id") + pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0] + unused = np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused) + _assert_eq(dist_np, mj_dist, "dist") + + def test_ray_cylinder(self): + """Tests ray<>cylinder matches MuJoCo.""" + mjm, mjd, m, d = test_util.fixture("ray.xml") + + pnt = wp.array([wp.vec3(2.0, 0.0, 0.05)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(0.0, 0.05, 1.0))], dtype=wp.vec3).reshape((1, 1)) + + mj_geomid = np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(mjm, mjd, pnt.numpy()[0, 0], vec.numpy()[0, 0], None, 1, -1, mj_geomid) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + + _assert_eq(geomid.numpy()[0, 0], mj_geomid[0], "geomid") + _assert_eq(dist.numpy()[0, 0], mj_dist, "dist") + + def test_ray_box(self): + """Tests ray<>box matches MuJoCo.""" + mjm, mjd, m, d = test_util.fixture("ray.xml") + + # looking down at box at a slight angle + pnt = wp.array([wp.vec3(1.0, 0.0, 1.6)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(0.0, 0.05, -1.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, 3, "geom_id") + pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0] + unused = np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused) + _assert_eq(dist_np, mj_dist, "dist") + + # looking up at box from below + pnt = wp.array([wp.vec3(1.0, 0.0, 0.05)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(0.0, 0.05, 1.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, 3, "geom_id") + pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0] + unused = np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused) + _assert_eq(dist_np, mj_dist, "dist") + + def test_ray_mesh(self): + """Tests ray<>mesh matches MuJoCo.""" + mjm, mjd, m, d = test_util.fixture("ray.xml") + + # look at the tetrahedron + pnt = wp.array([wp.vec3(2.0, 2.0, 2.0)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(-1.0, -1.0, -1.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, 4, "geom_id") + + pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0] + unused = np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused) + _assert_eq(dist_np, mj_dist, "dist-tetrahedron") + + # look away from the dodecahedron + pnt = wp.array([wp.vec3(4.0, 2.0, 2.0)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(2.0, 1.0, 1.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + _assert_eq(geomid_np, -1, "geom_id") + + # look at the dodecahedron + pnt = wp.array([wp.vec3(4.0, 2.0, 2.0)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(-2.0, -1.0, -1.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, 5, "geom_id") + + pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0] + unused = np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused) + _assert_eq(dist_np, mj_dist, "dist-dodecahedron") + + def test_ray_hfield(self): + mjm, mjd, m, d = test_util.fixture("ray.xml") + + pnt = wp.array([wp.vec3(0.0, 2.0, 2.0)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.vec3(0.0, 0.0, -1.0)], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + + mj_geomid = np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(mjm, mjd, pnt.numpy()[0, 0], vec.numpy()[0, 0], None, 1, -1, mj_geomid) + + _assert_eq(dist.numpy()[0, 0], mj_dist, "dist") + _assert_eq(geomid.numpy()[0, 0], mj_geomid[0], "geomid") + + def test_ray_geomgroup(self): + """Tests ray geomgroup filter.""" + mjm, mjd, m, d = test_util.fixture("ray.xml") + + # hits plane with geom_group[0] = 1 + pnt = wp.array([wp.vec3(2.0, 1.0, 3.0)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1)) + geomgroup = vec6(1, 0, 0, 0, 0, 0) + dist, geomid = mjwarp.ray(m, d, pnt, vec, geomgroup=geomgroup) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, 0, "geom_id") + + pnt_np, vec_np = pnt.numpy()[0, 0], vec.numpy()[0, 0] + unused = np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(mjm, mjd, pnt_np, vec_np, None, 1, -1, unused) + _assert_eq(dist_np, mj_dist, "dist") + + # nothing hit with geom_group[0] = 0 + pnt = wp.array([wp.vec3(2.0, 1.0, 3.0)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1)) + geomgroup = vec6(0, 0, 0, 0, 0, 0) + dist, geomid = mjwarp.ray(m, d, pnt, vec, geomgroup=geomgroup) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, -1, "geom_id") + _assert_eq(dist_np, -1, "dist") + + def test_ray_flg_static(self): + """Tests ray flg_static filter.""" + mjm, mjd, m, d = test_util.fixture("ray.xml") + + # nothing hit with flg_static = False + pnt = wp.array([wp.vec3(2.0, 1.0, 3.0)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec, flg_static=False) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, -1, "geom_id") + _assert_eq(dist_np, -1, "dist") + + def test_ray_bodyexclude(self): + """Tests ray bodyexclude filter.""" + mjm, mjd, m, d = test_util.fixture("ray.xml") + + # nothing hit with bodyexclude = 0 (world body) + pnt = wp.array([wp.vec3(2.0, 1.0, 3.0)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec, bodyexclude=0) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, -1, "geom_id") + _assert_eq(dist_np, -1, "dist") + + def test_ray_invisible(self): + """Tests ray doesn't hit transparent geoms.""" + mjm, mjd, m, d = test_util.fixture("ray.xml") + + # nothing hit with transparent geoms + m.geom_rgba = wp.array2d([[wp.vec4(0.0, 0.0, 0.0, 0.0)] * 8], dtype=wp.vec4) + mujoco.mj_forward(mjm, mjd) + + pnt = wp.array([wp.vec3(2.0, 1.0, 3.0)], dtype=wp.vec3).reshape((1, 1)) + vec = wp.array([wp.normalize(wp.vec3(0.1, 0.2, -1.0))], dtype=wp.vec3).reshape((1, 1)) + dist, geomid = mjwarp.ray(m, d, pnt, vec) + wp.synchronize() + geomid_np = geomid.numpy()[0, 0] + dist_np = dist.numpy()[0, 0] + _assert_eq(geomid_np, -1, "geom_id") + _assert_eq(dist_np, -1, "dist") + + +if __name__ == "__main__": + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py new file mode 100644 index 00000000..ecf5f322 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py @@ -0,0 +1,2113 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +from typing import Any, Tuple + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src import math +from mujoco.mjx.third_party.mujoco_warp._src import ray +from mujoco.mjx.third_party.mujoco_warp._src import smooth +from mujoco.mjx.third_party.mujoco_warp._src import support +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType +from mujoco.mjx.third_party.mujoco_warp._src.types import ConstraintType +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import DataType +from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import JointType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import ObjType +from mujoco.mjx.third_party.mujoco_warp._src.types import SensorType +from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType +from mujoco.mjx.third_party.mujoco_warp._src.types import vec6 +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel + + +@wp.func +def _write_scalar( + # Model: + sensor_datatype: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_cutoff: wp.array(dtype=float), + # In: + sensorid: int, + sensor: Any, + # Out: + out: wp.array(dtype=float), +): + adr = sensor_adr[sensorid] + cutoff = sensor_cutoff[sensorid] + + if cutoff > 0.0: + datatype = sensor_datatype[sensorid] + if datatype == int(DataType.REAL.value): + out[adr] = wp.clamp(sensor, -cutoff, cutoff) + elif datatype == int(DataType.POSITIVE.value): + out[adr] = wp.min(sensor, cutoff) + else: + out[adr] = sensor + + +@wp.func +def _write_vector( + # Model: + sensor_datatype: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_cutoff: wp.array(dtype=float), + # In: + sensorid: int, + sensordim: int, + sensor: Any, + # Out: + out: wp.array(dtype=float), +): + adr = sensor_adr[sensorid] + cutoff = sensor_cutoff[sensorid] + + if cutoff > 0.0: + datatype = sensor_datatype[sensorid] + if datatype == int(DataType.REAL.value): + for i in range(sensordim): + out[adr + i] = wp.clamp(sensor[i], -cutoff, cutoff) + elif datatype == int(DataType.POSITIVE.value): + for i in range(sensordim): + out[adr + i] = wp.min(sensor[i], cutoff) + else: + for i in range(sensordim): + out[adr + i] = sensor[i] + + +@wp.func +def _magnetometer( + # Model: + opt_magnetic: wp.array(dtype=wp.vec3), + # Data in: + site_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + worldid: int, + objid: int, +) -> wp.vec3: + magnetic = opt_magnetic[worldid] + return wp.transpose(site_xmat_in[worldid, objid]) @ magnetic + + +@wp.func +def _cam_projection( + # Model: + cam_fovy: wp.array(dtype=float), + cam_resolution: wp.array(dtype=wp.vec2i), + cam_sensorsize: wp.array(dtype=wp.vec2), + cam_intrinsic: wp.array(dtype=wp.vec4), + # Data in: + site_xpos_in: wp.array2d(dtype=wp.vec3), + cam_xpos_in: wp.array2d(dtype=wp.vec3), + cam_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + worldid: int, + objid: int, + refid: int, +) -> wp.vec2: + sensorsize = cam_sensorsize[refid] + intrinsic = cam_intrinsic[refid] + fovy = cam_fovy[refid] + res = cam_resolution[refid] + + target_xpos = site_xpos_in[worldid, objid] + xpos = cam_xpos_in[worldid, refid] + xmat = cam_xmat_in[worldid, refid] + + translation = wp.mat44f(1.0, 0.0, 0.0, -xpos[0], 0.0, 1.0, 0.0, -xpos[1], 0.0, 0.0, 1.0, -xpos[2], 0.0, 0.0, 0.0, 1.0) + rotation = wp.mat44f( + xmat[0, 0], xmat[1, 0], xmat[2, 0], 0.0, + xmat[0, 1], xmat[1, 1], xmat[2, 1], 0.0, + xmat[0, 2], xmat[1, 2], xmat[2, 2], 0.0, + 0.0, 0.0, 0.0, 1.0, + ) # fmt: skip + + # focal transformation matrix (3 x 4) + if sensorsize[0] != 0.0 and sensorsize[1] != 0.0: + fx = intrinsic[0] / (sensorsize[0] + MJ_MINVAL) * float(res[0]) + fy = intrinsic[1] / (sensorsize[1] + MJ_MINVAL) * float(res[1]) + focal = wp.mat44f(-fx, 0.0, 0.0, 0.0, 0.0, fy, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0) + else: + f = 0.5 / wp.tan(fovy * wp.static(wp.pi / 360.0)) * float(res[1]) + focal = wp.mat44f(-f, 0.0, 0.0, 0.0, 0.0, f, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0) + + # image matrix (3 x 3) + image = wp.mat44f( + 1.0, 0.0, 0.5 * float(res[0]), 0.0, 0.0, 1.0, 0.5 * float(res[1]), 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0 + ) + + # projection matrix (3 x 4): product of all 4 matrices + # TODO(team): compute proj directly + proj = image @ focal @ rotation @ translation + + # projection matrix multiples homogeneous [x, y, z, 1] vectors + pos_hom = wp.vec4(target_xpos[0], target_xpos[1], target_xpos[2], 1.0) + + # project world coordinates into pixel space, see: + # https://en.wikipedia.org/wiki/3D_projection#Mathematical_formula + pixel_coord_hom = proj @ pos_hom + + # avoid dividing by tiny numbers + denom = pixel_coord_hom[2] + if wp.abs(denom) < MJ_MINVAL: + denom = wp.clamp(denom, -MJ_MINVAL, MJ_MINVAL) + + # compute projection + return wp.vec2f(pixel_coord_hom[0], pixel_coord_hom[1]) / denom + + +@wp.kernel +def _sensor_rangefinder_init( + # Model: + sensor_objid: wp.array(dtype=int), + sensor_rangefinder_adr: wp.array(dtype=int), + # Data in: + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + # Data out: + sensor_rangefinder_pnt_out: wp.array2d(dtype=wp.vec3), + sensor_rangefinder_vec_out: wp.array2d(dtype=wp.vec3), +): + worldid, rfid = wp.tid() + sensorid = sensor_rangefinder_adr[rfid] + objid = sensor_objid[sensorid] + site_xpos = site_xpos_in[worldid, objid] + site_xmat = site_xmat_in[worldid, objid] + + sensor_rangefinder_pnt_out[worldid, rfid] = site_xpos + sensor_rangefinder_vec_out[worldid, rfid] = wp.vec3(site_xmat[0, 2], site_xmat[1, 2], site_xmat[2, 2]) + + +@wp.func +def _joint_pos(jnt_qposadr: wp.array(dtype=int), qpos_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float: + return qpos_in[worldid, jnt_qposadr[objid]] + + +@wp.func +def _tendon_pos(ten_length_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float: + return ten_length_in[worldid, objid] + + +@wp.func +def _actuator_pos(actuator_length_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float: + return actuator_length_in[worldid, objid] + + +@wp.func +def _ball_quat(jnt_qposadr: wp.array(dtype=int), qpos_in: wp.array2d(dtype=float), worldid: int, objid: int) -> wp.quat: + adr = jnt_qposadr[objid] + quat = wp.quat( + qpos_in[worldid, adr + 0], + qpos_in[worldid, adr + 1], + qpos_in[worldid, adr + 2], + qpos_in[worldid, adr + 3], + ) + return wp.normalize(quat) + + +@wp.kernel +def _limit_pos_zero( + # Model: + sensor_adr: wp.array(dtype=int), + sensor_limitpos_adr: wp.array(dtype=int), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, limitposid = wp.tid() + sensordata_out[worldid, sensor_adr[sensor_limitpos_adr[limitposid]]] = 0.0 + + +@wp.kernel +def _limit_pos( + # Model: + sensor_datatype: wp.array(dtype=int), + sensor_objid: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_cutoff: wp.array(dtype=float), + sensor_limitpos_adr: wp.array(dtype=int), + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nl_in: wp.array(dtype=int), + efc_type_in: wp.array2d(dtype=int), + efc_id_in: wp.array2d(dtype=int), + efc_pos_in: wp.array2d(dtype=float), + efc_margin_in: wp.array2d(dtype=float), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, efcid, limitposid = wp.tid() + + ne = ne_in[worldid] + nf = nf_in[worldid] + nl = nl_in[worldid] + + # skip if not limit + if efcid < ne + nf or efcid >= ne + nf + nl: + return + + sensorid = sensor_limitpos_adr[limitposid] + if efc_id_in[worldid, efcid] == sensor_objid[sensorid]: + efc_type = efc_type_in[worldid, efcid] + if efc_type == int(ConstraintType.LIMIT_JOINT.value) or efc_type == int(ConstraintType.LIMIT_TENDON.value): + val = efc_pos_in[worldid, efcid] - efc_margin_in[worldid, efcid] + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, sensordata_out[worldid]) + + +@wp.func +def _frame_pos( + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + xmat_in: wp.array2d(dtype=wp.mat33), + xipos_in: wp.array2d(dtype=wp.vec3), + ximat_in: wp.array2d(dtype=wp.mat33), + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + cam_xpos_in: wp.array2d(dtype=wp.vec3), + cam_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + worldid: int, + objid: int, + objtype: int, + refid: int, + reftype: int, +) -> wp.vec3: + if objtype == int(ObjType.BODY.value): + xpos = xipos_in[worldid, objid] + elif objtype == int(ObjType.XBODY.value): + xpos = xpos_in[worldid, objid] + elif objtype == int(ObjType.GEOM.value): + xpos = geom_xpos_in[worldid, objid] + elif objtype == int(ObjType.SITE.value): + xpos = site_xpos_in[worldid, objid] + elif objtype == int(ObjType.CAMERA.value): + xpos = cam_xpos_in[worldid, objid] + else: # UNKNOWN + xpos = wp.vec3(0.0) + + if refid == -1: + return xpos + + if reftype == int(ObjType.BODY.value): + xpos_ref = xipos_in[worldid, refid] + xmat_ref = ximat_in[worldid, refid] + elif objtype == int(ObjType.XBODY.value): + xpos_ref = xpos_in[worldid, refid] + xmat_ref = xmat_in[worldid, refid] + elif reftype == int(ObjType.GEOM.value): + xpos_ref = geom_xpos_in[worldid, refid] + xmat_ref = geom_xmat_in[worldid, refid] + elif reftype == int(ObjType.SITE.value): + xpos_ref = site_xpos_in[worldid, refid] + xmat_ref = site_xmat_in[worldid, refid] + elif reftype == int(ObjType.CAMERA.value): + xpos_ref = cam_xpos_in[worldid, refid] + xmat_ref = cam_xmat_in[worldid, refid] + + else: # UNKNOWN + xpos_ref = wp.vec3(0.0) + xmat_ref = wp.identity(3, wp.float32) + + return wp.transpose(xmat_ref) @ (xpos - xpos_ref) + + +@wp.func +def _frame_axis( + # Data in: + xmat_in: wp.array2d(dtype=wp.mat33), + ximat_in: wp.array2d(dtype=wp.mat33), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + site_xmat_in: wp.array2d(dtype=wp.mat33), + cam_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + worldid: int, + objid: int, + objtype: int, + refid: int, + reftype: int, + frame_axis: int, +) -> wp.vec3: + if objtype == int(ObjType.BODY.value): + xmat = ximat_in[worldid, objid] + axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis]) + elif objtype == int(ObjType.XBODY.value): + xmat = xmat_in[worldid, objid] + axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis]) + elif objtype == int(ObjType.GEOM.value): + xmat = geom_xmat_in[worldid, objid] + axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis]) + elif objtype == int(ObjType.SITE.value): + xmat = site_xmat_in[worldid, objid] + axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis]) + elif objtype == int(ObjType.CAMERA.value): + xmat = cam_xmat_in[worldid, objid] + axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis]) + else: # UNKNOWN + axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis]) + + if refid == -1: + return axis + + if reftype == int(ObjType.BODY.value): + xmat_ref = ximat_in[worldid, refid] + elif reftype == int(ObjType.XBODY.value): + xmat_ref = xmat_in[worldid, refid] + elif reftype == int(ObjType.GEOM.value): + xmat_ref = geom_xmat_in[worldid, refid] + elif reftype == int(ObjType.SITE.value): + xmat_ref = site_xmat_in[worldid, refid] + elif reftype == int(ObjType.CAMERA.value): + xmat_ref = cam_xmat_in[worldid, refid] + else: # UNKNOWN + xmat_ref = wp.identity(3, dtype=wp.float32) + + return wp.transpose(xmat_ref) @ axis + + +@wp.func +def _frame_quat( + # Model: + body_iquat: wp.array2d(dtype=wp.quat), + geom_bodyid: wp.array(dtype=int), + geom_quat: wp.array2d(dtype=wp.quat), + site_bodyid: wp.array(dtype=int), + site_quat: wp.array2d(dtype=wp.quat), + cam_bodyid: wp.array(dtype=int), + cam_quat: wp.array2d(dtype=wp.quat), + # Data in: + xquat_in: wp.array2d(dtype=wp.quat), + # In: + worldid: int, + objid: int, + objtype: int, + refid: int, + reftype: int, +) -> wp.quat: + if objtype == int(ObjType.BODY.value): + quat = math.mul_quat(xquat_in[worldid, objid], body_iquat[worldid, objid]) + elif objtype == int(ObjType.XBODY.value): + quat = xquat_in[worldid, objid] + elif objtype == int(ObjType.GEOM.value): + quat = math.mul_quat(xquat_in[worldid, geom_bodyid[objid]], geom_quat[worldid, objid]) + elif objtype == int(ObjType.SITE.value): + quat = math.mul_quat(xquat_in[worldid, site_bodyid[objid]], site_quat[worldid, objid]) + elif objtype == int(ObjType.CAMERA.value): + quat = math.mul_quat(xquat_in[worldid, cam_bodyid[objid]], cam_quat[worldid, objid]) + else: # UNKNOWN + quat = wp.quat(1.0, 0.0, 0.0, 0.0) + + if refid == -1: + return quat + + if reftype == int(ObjType.BODY.value): + refquat = math.mul_quat(xquat_in[worldid, refid], body_iquat[worldid, refid]) + elif reftype == int(ObjType.XBODY.value): + refquat = xquat_in[worldid, refid] + elif reftype == int(ObjType.GEOM.value): + refquat = math.mul_quat(xquat_in[worldid, geom_bodyid[refid]], geom_quat[worldid, refid]) + elif reftype == int(ObjType.SITE.value): + refquat = math.mul_quat(xquat_in[worldid, site_bodyid[refid]], site_quat[worldid, refid]) + elif reftype == int(ObjType.CAMERA.value): + refquat = math.mul_quat(xquat_in[worldid, cam_bodyid[refid]], cam_quat[worldid, refid]) + else: # UNKNOWN + refquat = wp.quat(1.0, 0.0, 0.0, 0.0) + + return math.mul_quat(math.quat_inv(refquat), quat) + + +@wp.func +def _subtree_com(subtree_com_in: wp.array2d(dtype=wp.vec3), worldid: int, objid: int) -> wp.vec3: + return subtree_com_in[worldid, objid] + + +@wp.func +def _clock(time_in: wp.array(dtype=float), worldid: int) -> float: + return time_in[worldid] + + +@wp.kernel +def _sensor_pos( + # Model: + opt_magnetic: wp.array(dtype=wp.vec3), + body_iquat: wp.array2d(dtype=wp.quat), + jnt_qposadr: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + geom_quat: wp.array2d(dtype=wp.quat), + site_bodyid: wp.array(dtype=int), + site_quat: wp.array2d(dtype=wp.quat), + cam_bodyid: wp.array(dtype=int), + cam_quat: wp.array2d(dtype=wp.quat), + cam_fovy: wp.array(dtype=float), + cam_resolution: wp.array(dtype=wp.vec2i), + cam_sensorsize: wp.array(dtype=wp.vec2), + cam_intrinsic: wp.array(dtype=wp.vec4), + sensor_type: wp.array(dtype=int), + sensor_datatype: wp.array(dtype=int), + sensor_objtype: wp.array(dtype=int), + sensor_objid: wp.array(dtype=int), + sensor_reftype: wp.array(dtype=int), + sensor_refid: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_cutoff: wp.array(dtype=float), + sensor_pos_adr: wp.array(dtype=int), + rangefinder_sensor_adr: wp.array(dtype=int), + # Data in: + time_in: wp.array(dtype=float), + energy_in: wp.array(dtype=wp.vec2), + qpos_in: wp.array2d(dtype=float), + xpos_in: wp.array2d(dtype=wp.vec3), + xquat_in: wp.array2d(dtype=wp.quat), + xmat_in: wp.array2d(dtype=wp.mat33), + xipos_in: wp.array2d(dtype=wp.vec3), + ximat_in: wp.array2d(dtype=wp.mat33), + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + cam_xpos_in: wp.array2d(dtype=wp.vec3), + cam_xmat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + actuator_length_in: wp.array2d(dtype=float), + ten_length_in: wp.array2d(dtype=float), + sensor_rangefinder_dist_in: wp.array2d(dtype=float), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, posid = wp.tid() + sensorid = sensor_pos_adr[posid] + sensortype = sensor_type[sensorid] + objid = sensor_objid[sensorid] + out = sensordata_out[worldid] + + if sensortype == int(SensorType.MAGNETOMETER.value): + vec3 = _magnetometer(opt_magnetic, site_xmat_in, worldid, objid) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + elif sensortype == int(SensorType.CAMPROJECTION.value): + refid = sensor_refid[sensorid] + vec2 = _cam_projection( + cam_fovy, cam_resolution, cam_sensorsize, cam_intrinsic, site_xpos_in, cam_xpos_in, cam_xmat_in, worldid, objid, refid + ) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 2, vec2, out) + elif sensortype == int(SensorType.RANGEFINDER.value): + val = sensor_rangefinder_dist_in[worldid, rangefinder_sensor_adr[sensorid]] + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) + elif sensortype == int(SensorType.JOINTPOS.value): + val = _joint_pos(jnt_qposadr, qpos_in, worldid, objid) + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) + elif sensortype == int(SensorType.TENDONPOS.value): + val = _tendon_pos(ten_length_in, worldid, objid) + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) + elif sensortype == int(SensorType.ACTUATORPOS.value): + val = _actuator_pos(actuator_length_in, worldid, objid) + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) + elif sensortype == int(SensorType.BALLQUAT.value): + quat = _ball_quat(jnt_qposadr, qpos_in, worldid, objid) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 4, quat, out) + elif sensortype == int(SensorType.FRAMEPOS.value): + objtype = sensor_objtype[sensorid] + refid = sensor_refid[sensorid] + reftype = sensor_reftype[sensorid] + vec3 = _frame_pos( + xpos_in, + xmat_in, + xipos_in, + ximat_in, + geom_xpos_in, + geom_xmat_in, + site_xpos_in, + site_xmat_in, + cam_xpos_in, + cam_xmat_in, + worldid, + objid, + objtype, + refid, + reftype, + ) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + elif ( + sensortype == int(SensorType.FRAMEXAXIS.value) + or sensortype == int(SensorType.FRAMEYAXIS.value) + or sensortype == int(SensorType.FRAMEZAXIS.value) + ): + objtype = sensor_objtype[sensorid] + refid = sensor_refid[sensorid] + reftype = sensor_reftype[sensorid] + if sensortype == int(SensorType.FRAMEXAXIS.value): + axis = 0 + elif sensortype == int(SensorType.FRAMEYAXIS.value): + axis = 1 + elif sensortype == int(SensorType.FRAMEZAXIS.value): + axis = 2 + vec3 = _frame_axis( + ximat_in, xmat_in, geom_xmat_in, site_xmat_in, cam_xmat_in, worldid, objid, objtype, refid, reftype, axis + ) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + elif sensortype == int(SensorType.FRAMEQUAT.value): + objtype = sensor_objtype[sensorid] + refid = sensor_refid[sensorid] + reftype = sensor_reftype[sensorid] + quat = _frame_quat( + body_iquat, + geom_bodyid, + geom_quat, + site_bodyid, + site_quat, + cam_bodyid, + cam_quat, + xquat_in, + worldid, + objid, + objtype, + refid, + reftype, + ) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 4, quat, out) + elif sensortype == int(SensorType.SUBTREECOM.value): + vec3 = _subtree_com(subtree_com_in, worldid, objid) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + elif sensortype == int(SensorType.E_POTENTIAL.value): + val = energy_in[worldid][0] + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) + elif sensortype == int(SensorType.E_KINETIC.value): + val = energy_in[worldid][1] + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) + elif sensortype == int(SensorType.CLOCK.value): + val = _clock(time_in, worldid) + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) + + +@event_scope +def sensor_pos(m: Model, d: Data): + """Compute position-dependent sensor values.""" + + if m.opt.disableflags & DisableBit.SENSOR: + return + + # rangefinder + if m.sensor_rangefinder_adr.size > 0: + # get position and direction + wp.launch( + _sensor_rangefinder_init, + dim=(d.nworld, m.sensor_rangefinder_adr.size), + inputs=[ + m.sensor_objid, + m.sensor_rangefinder_adr, + d.site_xpos, + d.site_xmat, + ], + outputs=[ + d.sensor_rangefinder_pnt, + d.sensor_rangefinder_vec, + ], + ) + + # get distances + ray.rays( + m, + d, + d.sensor_rangefinder_pnt, + d.sensor_rangefinder_vec, + vec6(wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf), + True, + m.sensor_rangefinder_bodyid, + d.sensor_rangefinder_dist, + d.sensor_rangefinder_geomid, + ) + + if m.sensor_e_potential: + energy_pos(m, d) + + if m.sensor_e_kinetic: + energy_vel(m, d) + + wp.launch( + _sensor_pos, + dim=(d.nworld, m.sensor_pos_adr.size), + inputs=[ + m.opt.magnetic, + m.body_iquat, + m.jnt_qposadr, + m.geom_bodyid, + m.geom_quat, + m.site_bodyid, + m.site_quat, + m.cam_bodyid, + m.cam_quat, + m.cam_fovy, + m.cam_resolution, + m.cam_sensorsize, + m.cam_intrinsic, + m.sensor_type, + m.sensor_datatype, + m.sensor_objtype, + m.sensor_objid, + m.sensor_reftype, + m.sensor_refid, + m.sensor_adr, + m.sensor_cutoff, + m.sensor_pos_adr, + m.rangefinder_sensor_adr, + d.time, + d.energy, + d.qpos, + d.xpos, + d.xquat, + d.xmat, + d.xipos, + d.ximat, + d.geom_xpos, + d.geom_xmat, + d.site_xpos, + d.site_xmat, + d.cam_xpos, + d.cam_xmat, + d.subtree_com, + d.actuator_length, + d.ten_length, + d.sensor_rangefinder_dist, + ], + outputs=[d.sensordata], + ) + + # jointlimitpos and tendonlimitpos + wp.launch( + _limit_pos_zero, + dim=(d.nworld, m.sensor_limitpos_adr.size), + inputs=[m.sensor_adr, m.sensor_limitpos_adr], + outputs=[d.sensordata], + ) + + wp.launch( + _limit_pos, + dim=(d.nworld, d.njmax, m.sensor_limitpos_adr.size), + inputs=[ + m.sensor_datatype, + m.sensor_objid, + m.sensor_adr, + m.sensor_cutoff, + m.sensor_limitpos_adr, + d.ne, + d.nf, + d.nl, + d.efc.type, + d.efc.id, + d.efc.pos, + d.efc.margin, + ], + outputs=[ + d.sensordata, + ], + ) + + +@wp.func +def _velocimeter( + # Model: + body_rootid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + # Data in: + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + # In: + worldid: int, + objid: int, +) -> wp.vec3: + bodyid = site_bodyid[objid] + pos = site_xpos_in[worldid, objid] + rot = site_xmat_in[worldid, objid] + cvel = cvel_in[worldid, bodyid] + ang = wp.spatial_top(cvel) + lin = wp.spatial_bottom(cvel) + subtree_com = subtree_com_in[worldid, body_rootid[bodyid]] + dif = pos - subtree_com + return wp.transpose(rot) @ (lin - wp.cross(dif, ang)) + + +@wp.func +def _gyro( + # Model: + site_bodyid: wp.array(dtype=int), + # Data in: + site_xmat_in: wp.array2d(dtype=wp.mat33), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + # In: + worldid: int, + objid: int, +) -> wp.vec3: + bodyid = site_bodyid[objid] + rot = site_xmat_in[worldid, objid] + cvel = cvel_in[worldid, bodyid] + ang = wp.spatial_top(cvel) + return wp.transpose(rot) @ ang + + +@wp.func +def _joint_vel(jnt_dofadr: wp.array(dtype=int), qvel_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float: + return qvel_in[worldid, jnt_dofadr[objid]] + + +@wp.func +def _tendon_vel(ten_velocity_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float: + return ten_velocity_in[worldid, objid] + + +@wp.func +def _actuator_vel(actuator_velocity_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float: + return actuator_velocity_in[worldid, objid] + + +@wp.func +def _ball_ang_vel(jnt_dofadr: wp.array(dtype=int), qvel_in: wp.array2d(dtype=float), worldid: int, objid: int) -> wp.vec3: + adr = jnt_dofadr[objid] + return wp.vec3(qvel_in[worldid, adr + 0], qvel_in[worldid, adr + 1], qvel_in[worldid, adr + 2]) + + +@wp.kernel +def _limit_vel_zero( + # Model: + sensor_adr: wp.array(dtype=int), + sensor_limitvel_adr: wp.array(dtype=int), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, limitvelid = wp.tid() + sensordata_out[worldid, sensor_adr[sensor_limitvel_adr[limitvelid]]] = 0.0 + + +@wp.kernel +def _limit_vel( + # Model: + sensor_datatype: wp.array(dtype=int), + sensor_objid: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_cutoff: wp.array(dtype=float), + sensor_limitvel_adr: wp.array(dtype=int), + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nl_in: wp.array(dtype=int), + efc_type_in: wp.array2d(dtype=int), + efc_id_in: wp.array2d(dtype=int), + efc_vel_in: wp.array2d(dtype=float), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, efcid, limitvelid = wp.tid() + + ne = ne_in[worldid] + nf = nf_in[worldid] + nl = nl_in[worldid] + + # skip if not limit + if efcid < ne + nf or efcid >= ne + nf + nl: + return + + sensorid = sensor_limitvel_adr[limitvelid] + if efc_id_in[worldid, efcid] == sensor_objid[sensorid]: + efc_type = efc_type_in[worldid, efcid] + if efc_type == int(ConstraintType.LIMIT_JOINT.value) or efc_type == int(ConstraintType.LIMIT_TENDON.value): + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, efc_vel_in[worldid, efcid], sensordata_out[worldid]) + + +@wp.func +def _cvel_offset( + # Model: + body_rootid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + cam_bodyid: wp.array(dtype=int), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + xipos_in: wp.array2d(dtype=wp.vec3), + geom_xpos_in: wp.array2d(dtype=wp.vec3), + site_xpos_in: wp.array2d(dtype=wp.vec3), + cam_xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + # In: + worldid: int, + objtype: int, + objid: int, +) -> Tuple[wp.spatial_vector, wp.vec3]: + if objtype == int(ObjType.BODY.value): + pos = xipos_in[worldid, objid] + bodyid = objid + elif objtype == int(ObjType.XBODY.value): + pos = xpos_in[worldid, objid] + bodyid = objid + elif objtype == int(ObjType.GEOM.value): + pos = geom_xpos_in[worldid, objid] + bodyid = geom_bodyid[objid] + elif objtype == int(ObjType.SITE.value): + pos = site_xpos_in[worldid, objid] + bodyid = site_bodyid[objid] + elif objtype == int(ObjType.CAMERA.value): + pos = cam_xpos_in[worldid, objid] + bodyid = cam_bodyid[objid] + else: # UNKNOWN + pos = wp.vec3(0.0) + bodyid = 0 + + return cvel_in[worldid, bodyid], pos - subtree_com_in[worldid, body_rootid[bodyid]] + + +@wp.func +def _frame_linvel( + # Model: + body_rootid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + cam_bodyid: wp.array(dtype=int), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + xmat_in: wp.array2d(dtype=wp.mat33), + xipos_in: wp.array2d(dtype=wp.vec3), + ximat_in: wp.array2d(dtype=wp.mat33), + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + cam_xpos_in: wp.array2d(dtype=wp.vec3), + cam_xmat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + # In: + worldid: int, + objid: int, + objtype: int, + refid: int, + reftype: int, +) -> wp.vec3: + if objtype == int(ObjType.BODY.value): + xpos = xipos_in[worldid, objid] + elif objtype == int(ObjType.XBODY.value): + xpos = xpos_in[worldid, objid] + elif objtype == int(ObjType.GEOM.value): + xpos = geom_xpos_in[worldid, objid] + elif objtype == int(ObjType.SITE.value): + xpos = site_xpos_in[worldid, objid] + elif objtype == int(ObjType.CAMERA.value): + xpos = cam_xpos_in[worldid, objid] + else: # UNKNOWN + xpos = wp.vec3(0.0) + + if reftype == int(ObjType.BODY.value): + xposref = xipos_in[worldid, refid] + xmatref = ximat_in[worldid, refid] + elif reftype == int(ObjType.XBODY.value): + xposref = xpos_in[worldid, refid] + xmatref = xmat_in[worldid, refid] + elif reftype == int(ObjType.GEOM.value): + xposref = geom_xpos_in[worldid, refid] + xmatref = geom_xmat_in[worldid, refid] + elif reftype == int(ObjType.SITE.value): + xposref = site_xpos_in[worldid, refid] + xmatref = site_xmat_in[worldid, refid] + elif reftype == int(ObjType.CAMERA.value): + xposref = cam_xpos_in[worldid, refid] + xmatref = cam_xmat_in[worldid, refid] + else: # UNKNOWN + xposref = wp.vec3(0.0) + xmatref = wp.identity(3, dtype=float) + + cvel, offset = _cvel_offset( + body_rootid, + geom_bodyid, + site_bodyid, + cam_bodyid, + xpos_in, + xipos_in, + geom_xpos_in, + site_xpos_in, + cam_xpos_in, + subtree_com_in, + cvel_in, + worldid, + objtype, + objid, + ) + cvelref, offsetref = _cvel_offset( + body_rootid, + geom_bodyid, + site_bodyid, + cam_bodyid, + xpos_in, + xipos_in, + geom_xpos_in, + site_xpos_in, + cam_xpos_in, + subtree_com_in, + cvel_in, + worldid, + reftype, + refid, + ) + clinvel = wp.spatial_bottom(cvel) + cangvel = wp.spatial_top(cvel) + cangvelref = wp.spatial_top(cvelref) + xlinvel = clinvel - wp.cross(offset, cangvel) + + if refid > -1: + clinvelref = wp.spatial_bottom(cvelref) + xlinvelref = clinvelref - wp.cross(offsetref, cangvelref) + rvec = xpos - xposref + rel_vel = xlinvel - xlinvelref + wp.cross(rvec, cangvelref) + return wp.transpose(xmatref) @ rel_vel + else: + return xlinvel + + +@wp.func +def _frame_angvel( + # Model: + body_rootid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + cam_bodyid: wp.array(dtype=int), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + xmat_in: wp.array2d(dtype=wp.mat33), + xipos_in: wp.array2d(dtype=wp.vec3), + ximat_in: wp.array2d(dtype=wp.mat33), + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + cam_xpos_in: wp.array2d(dtype=wp.vec3), + cam_xmat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + # In: + worldid: int, + objid: int, + objtype: int, + refid: int, + reftype: int, +) -> wp.vec3: + cvel, _ = _cvel_offset( + body_rootid, + geom_bodyid, + site_bodyid, + cam_bodyid, + xpos_in, + xipos_in, + geom_xpos_in, + site_xpos_in, + cam_xpos_in, + subtree_com_in, + cvel_in, + worldid, + objtype, + objid, + ) + cangvel = wp.spatial_top(cvel) + + if refid > -1: + if reftype == int(ObjType.BODY.value): + xmatref = ximat_in[worldid, refid] + elif reftype == int(ObjType.XBODY.value): + xmatref = xmat_in[worldid, refid] + elif reftype == int(ObjType.GEOM.value): + xmatref = geom_xmat_in[worldid, refid] + elif reftype == int(ObjType.SITE.value): + xmatref = site_xmat_in[worldid, refid] + elif reftype == int(ObjType.CAMERA.value): + xmatref = cam_xmat_in[worldid, refid] + else: # UNKNOWN + xmatref = wp.identity(3, dtype=float) + + cvelref, _ = _cvel_offset( + body_rootid, + geom_bodyid, + site_bodyid, + cam_bodyid, + xpos_in, + xipos_in, + geom_xpos_in, + site_xpos_in, + cam_xpos_in, + subtree_com_in, + cvel_in, + worldid, + reftype, + refid, + ) + cangvelref = wp.spatial_top(cvelref) + + return wp.transpose(xmatref) @ (cangvel - cangvelref) + else: + return cangvel + + +@wp.func +def _subtree_linvel(subtree_linvel_in: wp.array2d(dtype=wp.vec3), worldid: int, objid: int) -> wp.vec3: + return subtree_linvel_in[worldid, objid] + + +@wp.func +def _subtree_angmom(subtree_angmom_in: wp.array2d(dtype=wp.vec3), worldid: int, objid: int) -> wp.vec3: + return subtree_angmom_in[worldid, objid] + + +@wp.kernel +def _sensor_vel( + # Model: + body_rootid: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + cam_bodyid: wp.array(dtype=int), + sensor_type: wp.array(dtype=int), + sensor_datatype: wp.array(dtype=int), + sensor_objtype: wp.array(dtype=int), + sensor_objid: wp.array(dtype=int), + sensor_reftype: wp.array(dtype=int), + sensor_refid: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_cutoff: wp.array(dtype=float), + sensor_vel_adr: wp.array(dtype=int), + # Data in: + qvel_in: wp.array2d(dtype=float), + xpos_in: wp.array2d(dtype=wp.vec3), + xmat_in: wp.array2d(dtype=wp.mat33), + xipos_in: wp.array2d(dtype=wp.vec3), + ximat_in: wp.array2d(dtype=wp.mat33), + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + cam_xpos_in: wp.array2d(dtype=wp.vec3), + cam_xmat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + ten_velocity_in: wp.array2d(dtype=float), + actuator_velocity_in: wp.array2d(dtype=float), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + subtree_linvel_in: wp.array2d(dtype=wp.vec3), + subtree_angmom_in: wp.array2d(dtype=wp.vec3), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, velid = wp.tid() + sensorid = sensor_vel_adr[velid] + sensortype = sensor_type[sensorid] + objid = sensor_objid[sensorid] + out = sensordata_out[worldid] + + if sensortype == int(SensorType.VELOCIMETER.value): + vec3 = _velocimeter(body_rootid, site_bodyid, site_xpos_in, site_xmat_in, subtree_com_in, cvel_in, worldid, objid) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + elif sensortype == int(SensorType.GYRO.value): + vec3 = _gyro(site_bodyid, site_xmat_in, cvel_in, worldid, objid) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + elif sensortype == int(SensorType.JOINTVEL.value): + val = _joint_vel(jnt_dofadr, qvel_in, worldid, objid) + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) + elif sensortype == int(SensorType.TENDONVEL.value): + val = _tendon_vel(ten_velocity_in, worldid, objid) + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) + elif sensortype == int(SensorType.ACTUATORVEL.value): + val = _actuator_vel(actuator_velocity_in, worldid, objid) + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) + elif sensortype == int(SensorType.BALLANGVEL.value): + vec3 = _ball_ang_vel(jnt_dofadr, qvel_in, worldid, objid) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + elif sensortype == int(SensorType.FRAMELINVEL.value): + objtype = sensor_objtype[sensorid] + refid = sensor_refid[sensorid] + reftype = sensor_reftype[sensorid] + frame_linvel = _frame_linvel( + body_rootid, + geom_bodyid, + site_bodyid, + cam_bodyid, + xpos_in, + xmat_in, + xipos_in, + ximat_in, + geom_xpos_in, + geom_xmat_in, + site_xpos_in, + site_xmat_in, + cam_xpos_in, + cam_xmat_in, + subtree_com_in, + cvel_in, + worldid, + objid, + objtype, + refid, + reftype, + ) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, frame_linvel, out) + elif sensortype == int(SensorType.FRAMEANGVEL.value): + objtype = sensor_objtype[sensorid] + refid = sensor_refid[sensorid] + reftype = sensor_reftype[sensorid] + frame_angvel = _frame_angvel( + body_rootid, + geom_bodyid, + site_bodyid, + cam_bodyid, + xpos_in, + xmat_in, + xipos_in, + ximat_in, + geom_xpos_in, + geom_xmat_in, + site_xpos_in, + site_xmat_in, + cam_xpos_in, + cam_xmat_in, + subtree_com_in, + cvel_in, + worldid, + objid, + objtype, + refid, + reftype, + ) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, frame_angvel, out) + elif sensortype == int(SensorType.SUBTREELINVEL.value): + vec3 = _subtree_linvel(subtree_linvel_in, worldid, objid) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + elif sensortype == int(SensorType.SUBTREEANGMOM.value): + vec3 = _subtree_angmom(subtree_angmom_in, worldid, objid) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + + +@event_scope +def sensor_vel(m: Model, d: Data): + """Compute velocity-dependent sensor values.""" + + if m.opt.disableflags & DisableBit.SENSOR: + return + + if m.sensor_subtree_vel: + smooth.subtree_vel(m, d) + + wp.launch( + _sensor_vel, + dim=(d.nworld, m.sensor_vel_adr.size), + inputs=[ + m.body_rootid, + m.jnt_dofadr, + m.geom_bodyid, + m.site_bodyid, + m.cam_bodyid, + m.sensor_type, + m.sensor_datatype, + m.sensor_objtype, + m.sensor_objid, + m.sensor_reftype, + m.sensor_refid, + m.sensor_adr, + m.sensor_cutoff, + m.sensor_vel_adr, + d.qvel, + d.xpos, + d.xmat, + d.xipos, + d.ximat, + d.geom_xpos, + d.geom_xmat, + d.site_xpos, + d.site_xmat, + d.cam_xpos, + d.cam_xmat, + d.subtree_com, + d.ten_velocity, + d.actuator_velocity, + d.cvel, + d.subtree_linvel, + d.subtree_angmom, + ], + outputs=[d.sensordata], + ) + + wp.launch( + _limit_vel_zero, + dim=(d.nworld, m.sensor_limitvel_adr.size), + inputs=[m.sensor_adr, m.sensor_limitvel_adr], + outputs=[d.sensordata], + ) + + wp.launch( + _limit_vel, + dim=(d.nworld, d.njmax, m.sensor_limitvel_adr.size), + inputs=[ + m.sensor_datatype, + m.sensor_objid, + m.sensor_adr, + m.sensor_cutoff, + m.sensor_limitvel_adr, + d.ne, + d.nf, + d.nl, + d.efc.type, + d.efc.id, + d.efc.vel, + ], + outputs=[ + d.sensordata, + ], + ) + + +@wp.func +def _accelerometer( + # Model: + body_rootid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + # Data in: + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + cacc_in: wp.array2d(dtype=wp.spatial_vector), + # In: + worldid: int, + objid: int, +) -> wp.vec3: + bodyid = site_bodyid[objid] + rot = site_xmat_in[worldid, objid] + rotT = wp.transpose(rot) + cvel = cvel_in[worldid, bodyid] + cvel_top = wp.spatial_top(cvel) + cvel_bottom = wp.spatial_bottom(cvel) + cacc = cacc_in[worldid, bodyid] + cacc_top = wp.spatial_top(cacc) + cacc_bottom = wp.spatial_bottom(cacc) + dif = site_xpos_in[worldid, objid] - subtree_com_in[worldid, body_rootid[bodyid]] + ang = rotT @ cvel_top + lin = rotT @ (cvel_bottom - wp.cross(dif, cvel_top)) + acc = rotT @ (cacc_bottom - wp.cross(dif, cacc_top)) + correction = wp.cross(ang, lin) + return acc + correction + + +@wp.func +def _force( + # Model: + site_bodyid: wp.array(dtype=int), + # Data in: + site_xmat_in: wp.array2d(dtype=wp.mat33), + cfrc_int_in: wp.array2d(dtype=wp.spatial_vector), + # In: + worldid: int, + objid: int, +) -> wp.vec3: + bodyid = site_bodyid[objid] + cfrc_int = cfrc_int_in[worldid, bodyid] + site_xmat = site_xmat_in[worldid, objid] + return wp.transpose(site_xmat) @ wp.spatial_bottom(cfrc_int) + + +@wp.func +def _torque( + # Model: + body_rootid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + # Data in: + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cfrc_int_in: wp.array2d(dtype=wp.spatial_vector), + # In: + worldid: int, + objid: int, +) -> wp.vec3: + bodyid = site_bodyid[objid] + cfrc_int = cfrc_int_in[worldid, bodyid] + site_xmat = site_xmat_in[worldid, objid] + dif = site_xpos_in[worldid, objid] - subtree_com_in[worldid, body_rootid[bodyid]] + return wp.transpose(site_xmat) @ (wp.spatial_top(cfrc_int) - wp.cross(dif, wp.spatial_bottom(cfrc_int))) + + +@wp.func +def _actuator_force(actuator_force_in: wp.array2d(dtype=float), worldid: int, objid: int) -> float: + return actuator_force_in[worldid, objid] + + +@wp.func +def _joint_actuator_force( + # Model: + jnt_dofadr: wp.array(dtype=int), + # Data in: + qfrc_actuator_in: wp.array2d(dtype=float), + # In: + worldid: int, + objid: int, +) -> float: + return qfrc_actuator_in[worldid, jnt_dofadr[objid]] + + +@wp.kernel +def _tendon_actuator_force_zero( + # Model: + sensor_adr: wp.array(dtype=int), + sensor_tendonactfrc_adr: wp.array(dtype=int), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, tenactfrcid = wp.tid() + sensorid = sensor_tendonactfrc_adr[tenactfrcid] + adr = sensor_adr[sensorid] + sensordata_out[worldid, adr] = 0.0 + + +@wp.kernel +def _tendon_actuator_force( + # Model: + actuator_trntype: wp.array(dtype=int), + actuator_trnid: wp.array(dtype=wp.vec2i), + sensor_objid: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_tendonactfrc_adr: wp.array(dtype=int), + # Data in: + actuator_force_in: wp.array2d(dtype=float), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, tenactfrcid, actid = wp.tid() + sensorid = sensor_tendonactfrc_adr[tenactfrcid] + + if actuator_trntype[actid] == int(TrnType.TENDON.value) and actuator_trnid[actid][0] == sensor_objid[sensorid]: + adr = sensor_adr[sensorid] + sensordata_out[worldid, adr] += actuator_force_in[worldid, actid] + + +@wp.kernel +def _tendon_actuator_force_cutoff( + # Model: + sensor_datatype: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_cutoff: wp.array(dtype=float), + sensor_tendonactfrc_adr: wp.array(dtype=int), + # Data in: + sensordata_in: wp.array2d(dtype=float), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, tenactfrcid = wp.tid() + sensorid = sensor_tendonactfrc_adr[tenactfrcid] + adr = sensor_adr[sensorid] + val = sensordata_in[worldid, adr] + + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, sensordata_out[worldid]) + + +@wp.kernel +def _limit_frc_zero( + # Model: + sensor_adr: wp.array(dtype=int), + sensor_limitfrc_adr: wp.array(dtype=int), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, limitfrcid = wp.tid() + sensordata_out[worldid, sensor_adr[sensor_limitfrc_adr[limitfrcid]]] = 0.0 + + +@wp.kernel +def _limit_frc( + # Model: + sensor_datatype: wp.array(dtype=int), + sensor_objid: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_cutoff: wp.array(dtype=float), + sensor_limitfrc_adr: wp.array(dtype=int), + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nl_in: wp.array(dtype=int), + efc_type_in: wp.array2d(dtype=int), + efc_id_in: wp.array2d(dtype=int), + efc_force_in: wp.array2d(dtype=float), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, efcid, limitfrcid = wp.tid() + + ne = ne_in[worldid] + nf = nf_in[worldid] + nl = nl_in[worldid] + + # skip if not limit + if efcid < ne + nf or efcid >= ne + nf + nl: + return + + sensorid = sensor_limitfrc_adr[limitfrcid] + if efc_id_in[worldid, efcid] == sensor_objid[sensorid]: + efc_type = efc_type_in[worldid, efcid] + if efc_type == int(ConstraintType.LIMIT_JOINT.value) or efc_type == int(ConstraintType.LIMIT_TENDON.value): + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, efc_force_in[worldid, efcid], sensordata_out[worldid]) + + +@wp.func +def _framelinacc( + # Model: + body_rootid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + cam_bodyid: wp.array(dtype=int), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + xipos_in: wp.array2d(dtype=wp.vec3), + geom_xpos_in: wp.array2d(dtype=wp.vec3), + site_xpos_in: wp.array2d(dtype=wp.vec3), + cam_xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + cacc_in: wp.array2d(dtype=wp.spatial_vector), + # In: + worldid: int, + objid: int, + objtype: int, +) -> wp.vec3: + if objtype == int(ObjType.BODY.value): + bodyid = objid + pos = xipos_in[worldid, objid] + elif objtype == int(ObjType.XBODY.value): + bodyid = objid + pos = xpos_in[worldid, objid] + elif objtype == int(ObjType.GEOM.value): + bodyid = geom_bodyid[objid] + pos = geom_xpos_in[worldid, objid] + elif objtype == int(ObjType.SITE.value): + bodyid = site_bodyid[objid] + pos = site_xpos_in[worldid, objid] + elif objtype == int(ObjType.CAMERA.value): + bodyid = cam_bodyid[objid] + pos = cam_xpos_in[worldid, objid] + else: # UNKNOWN + bodyid = 0 + pos = wp.vec3(0.0) + + cacc = cacc_in[worldid, bodyid] + cvel = cvel_in[worldid, bodyid] + offset = pos - subtree_com_in[worldid, body_rootid[bodyid]] + ang = wp.spatial_top(cvel) + lin = wp.spatial_bottom(cvel) - wp.cross(offset, ang) + acc = wp.spatial_bottom(cacc) - wp.cross(offset, wp.spatial_top(cacc)) + correction = wp.cross(ang, lin) + + return acc + correction + + +@wp.func +def _frameangacc( + # Model: + geom_bodyid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + cam_bodyid: wp.array(dtype=int), + # Data in: + cacc_in: wp.array2d(dtype=wp.spatial_vector), + # In: + worldid: int, + objid: int, + objtype: int, +) -> wp.vec3: + if objtype == int(ObjType.BODY.value) or objtype == int(ObjType.XBODY.value): + bodyid = objid + elif objtype == int(ObjType.GEOM.value): + bodyid = geom_bodyid[objid] + elif objtype == int(ObjType.SITE.value): + bodyid = site_bodyid[objid] + elif objtype == int(ObjType.CAMERA.value): + bodyid = cam_bodyid[objid] + else: # UNKNOWN + bodyid = 0 + + return wp.spatial_top(cacc_in[worldid, bodyid]) + + +@wp.kernel +def _sensor_acc( + # Model: + body_rootid: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + cam_bodyid: wp.array(dtype=int), + sensor_type: wp.array(dtype=int), + sensor_datatype: wp.array(dtype=int), + sensor_objtype: wp.array(dtype=int), + sensor_objid: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_cutoff: wp.array(dtype=float), + sensor_acc_adr: wp.array(dtype=int), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + xipos_in: wp.array2d(dtype=wp.vec3), + geom_xpos_in: wp.array2d(dtype=wp.vec3), + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + cam_xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + actuator_force_in: wp.array2d(dtype=float), + qfrc_actuator_in: wp.array2d(dtype=float), + cacc_in: wp.array2d(dtype=wp.spatial_vector), + cfrc_int_in: wp.array2d(dtype=wp.spatial_vector), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, accid = wp.tid() + sensorid = sensor_acc_adr[accid] + sensortype = sensor_type[sensorid] + objid = sensor_objid[sensorid] + out = sensordata_out[worldid] + + if sensortype == int(SensorType.ACCELEROMETER.value): + vec3 = _accelerometer( + body_rootid, site_bodyid, site_xpos_in, site_xmat_in, subtree_com_in, cvel_in, cacc_in, worldid, objid + ) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + elif sensortype == int(SensorType.FORCE.value): + vec3 = _force(site_bodyid, site_xmat_in, cfrc_int_in, worldid, objid) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + elif sensortype == int(SensorType.TORQUE.value): + vec3 = _torque(body_rootid, site_bodyid, site_xpos_in, site_xmat_in, subtree_com_in, cfrc_int_in, worldid, objid) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + elif sensortype == int(SensorType.ACTUATORFRC.value): + val = _actuator_force(actuator_force_in, worldid, objid) + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) + elif sensortype == int(SensorType.JOINTACTFRC.value): + val = _joint_actuator_force(jnt_dofadr, qfrc_actuator_in, worldid, objid) + _write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) + elif sensortype == int(SensorType.FRAMELINACC.value): + objtype = sensor_objtype[sensorid] + vec3 = _framelinacc( + body_rootid, + geom_bodyid, + site_bodyid, + cam_bodyid, + xpos_in, + xipos_in, + geom_xpos_in, + site_xpos_in, + cam_xpos_in, + subtree_com_in, + cvel_in, + cacc_in, + worldid, + objid, + objtype, + ) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + elif sensortype == int(SensorType.FRAMEANGACC.value): + objtype = sensor_objtype[sensorid] + vec3 = _frameangacc( + geom_bodyid, + site_bodyid, + cam_bodyid, + cacc_in, + worldid, + objid, + objtype, + ) + _write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) + + +@wp.kernel +def _sensor_touch_zero( + # Model: + sensor_adr: wp.array(dtype=int), + sensor_touch_adr: wp.array(dtype=int), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + worldid, sensortouchadrid = wp.tid() + sensorid = sensor_touch_adr[sensortouchadrid] + adr = sensor_adr[sensorid] + sensordata_out[worldid, adr] = 0.0 + + +@wp.kernel +def _sensor_touch( + # Model: + opt_cone: int, + geom_bodyid: wp.array(dtype=int), + site_type: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + site_size: wp.array(dtype=wp.vec3), + sensor_objid: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_touch_adr: wp.array(dtype=int), + # Data in: + ncon_in: wp.array(dtype=int), + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + contact_pos_in: wp.array(dtype=wp.vec3), + contact_frame_in: wp.array(dtype=wp.mat33), + contact_dim_in: wp.array(dtype=int), + contact_geom_in: wp.array(dtype=wp.vec2i), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_force_in: wp.array2d(dtype=float), + # Data out: + sensordata_out: wp.array2d(dtype=float), +): + conid, sensortouchadrid = wp.tid() + + if conid > ncon_in[0]: + return + + sensorid = sensor_touch_adr[sensortouchadrid] + + objid = sensor_objid[sensorid] + bodyid = site_bodyid[objid] + + # find contact in sensor zone, add normal force + + # contacting bodies + geom = contact_geom_in[conid] + conbody = wp.vec2i(geom_bodyid[geom[0]], geom_bodyid[geom[1]]) + + # select contacts involving sensorized body + worldid = contact_worldid_in[conid] + efc_address0 = contact_efc_address_in[conid, 0] + if efc_address0 >= 0 and (bodyid == conbody[0] or bodyid == conbody[1]): + # get contact normal force + normalforce = efc_force_in[worldid, efc_address0] + + if opt_cone == int(ConeType.PYRAMIDAL.value): + dim = contact_dim_in[conid] + for i in range(1, 2 * (dim - 1)): + normalforce += efc_force_in[worldid, contact_efc_address_in[conid, i]] + + if normalforce <= 0.0: + return + + # convert contact normal force to global frame, normalize + frame = contact_frame_in[conid] + conray = wp.vec3(frame[0, 0], frame[0, 1], frame[0, 2]) * normalforce + conray, _ = math.normalize_with_norm(conray) + + # flip ray direction if sensor is on body2 + if bodyid == conbody[1]: + conray = -conray + + # add if ray-zone intersection (always true when contact.pos inside zone) + if ( + ray.ray_geom( + site_xpos_in[worldid, objid], + site_xmat_in[worldid, objid], + site_size[objid], + contact_pos_in[conid], + conray, + site_type[objid], + ) + >= 0.0 + ): + adr = sensor_adr[sensorid] + wp.atomic_add(sensordata_out[worldid], adr, normalforce) + + +@event_scope +def sensor_acc(m: Model, d: Data): + """Compute acceleration-dependent sensor values.""" + if m.opt.disableflags & DisableBit.SENSOR: + return + + wp.launch( + _sensor_touch_zero, + dim=(d.nworld, m.sensor_touch_adr.size), + inputs=[ + m.sensor_adr, + m.sensor_touch_adr, + ], + outputs=[ + d.sensordata, + ], + ) + + wp.launch( + _sensor_touch, + dim=(d.nconmax, m.sensor_touch_adr.size), + inputs=[ + m.opt.cone, + m.geom_bodyid, + m.site_type, + m.site_bodyid, + m.site_size, + m.sensor_objid, + m.sensor_adr, + m.sensor_touch_adr, + d.ncon, + d.site_xpos, + d.site_xmat, + d.contact.pos, + d.contact.frame, + d.contact.dim, + d.contact.geom, + d.contact.efc_address, + d.contact.worldid, + d.efc.force, + ], + outputs=[ + d.sensordata, + ], + ) + + if m.sensor_rne_postconstraint: + smooth.rne_postconstraint(m, d) + + wp.launch( + _sensor_acc, + dim=(d.nworld, m.sensor_acc_adr.size), + inputs=[ + m.body_rootid, + m.jnt_dofadr, + m.geom_bodyid, + m.site_bodyid, + m.cam_bodyid, + m.sensor_type, + m.sensor_datatype, + m.sensor_objtype, + m.sensor_objid, + m.sensor_adr, + m.sensor_cutoff, + m.sensor_acc_adr, + d.xpos, + d.xipos, + d.geom_xpos, + d.site_xpos, + d.site_xmat, + d.cam_xpos, + d.subtree_com, + d.cvel, + d.actuator_force, + d.qfrc_actuator, + d.cacc, + d.cfrc_int, + ], + outputs=[d.sensordata], + ) + + wp.launch( + _tendon_actuator_force_zero, + dim=(d.nworld, m.sensor_tendonactfrc_adr.size), + inputs=[ + m.sensor_adr, + m.sensor_tendonactfrc_adr, + ], + outputs=[ + d.sensordata, + ], + ) + + wp.launch( + _tendon_actuator_force, + dim=(d.nworld, m.sensor_tendonactfrc_adr.size, m.nu), + inputs=[ + m.actuator_trntype, + m.actuator_trnid, + m.sensor_objid, + m.sensor_adr, + m.sensor_tendonactfrc_adr, + d.actuator_force, + ], + outputs=[ + d.sensordata, + ], + ) + + wp.launch( + _tendon_actuator_force_cutoff, + dim=(d.nworld, m.sensor_tendonactfrc_adr.size), + inputs=[ + m.sensor_datatype, + m.sensor_adr, + m.sensor_cutoff, + m.sensor_tendonactfrc_adr, + d.sensordata, + ], + outputs=[d.sensordata], + ) + + wp.launch( + _limit_frc_zero, + dim=(d.nworld, m.sensor_limitfrc_adr.size), + inputs=[m.sensor_adr, m.sensor_limitfrc_adr], + outputs=[d.sensordata], + ) + + wp.launch( + _limit_frc, + dim=(d.nworld, d.njmax, m.sensor_limitfrc_adr.size), + inputs=[ + m.sensor_datatype, + m.sensor_objid, + m.sensor_adr, + m.sensor_cutoff, + m.sensor_limitfrc_adr, + d.ne, + d.nf, + d.nl, + d.efc.type, + d.efc.id, + d.efc.force, + ], + outputs=[ + d.sensordata, + ], + ) + + +@wp.kernel +def _energy_pos_zero( + # Data out: + energy_out: wp.array(dtype=wp.vec2), +): + worldid = wp.tid() + energy_out[worldid][0] = 0.0 + + +@wp.kernel +def _energy_pos_gravity( + # Model: + opt_gravity: wp.array(dtype=wp.vec3), + body_mass: wp.array2d(dtype=float), + # Data in: + xipos_in: wp.array2d(dtype=wp.vec3), + # Data out: + energy_out: wp.array(dtype=wp.vec2), +): + worldid, bodyid = wp.tid() + gravity = opt_gravity[worldid] + bodyid += 1 # skip world body + + energy = wp.vec2( + body_mass[worldid, bodyid] * wp.dot(gravity, xipos_in[worldid, bodyid]), + 0.0, + ) + + wp.atomic_sub(energy_out, worldid, energy) + + +@wp.kernel +def _energy_pos_passive_joint( + # Model: + qpos_spring: wp.array2d(dtype=float), + jnt_type: wp.array(dtype=int), + jnt_qposadr: wp.array(dtype=int), + jnt_stiffness: wp.array2d(dtype=float), + # Data in: + qpos_in: wp.array2d(dtype=float), + # Data out: + energy_out: wp.array(dtype=wp.vec2), +): + worldid, jntid = wp.tid() + stiffness = jnt_stiffness[worldid, jntid] + + if stiffness == 0.0: + return + + padr = jnt_qposadr[jntid] + jnttype = jnt_type[jntid] + + if jnttype == int(JointType.FREE.value): + dif0 = wp.vec3( + qpos_in[worldid, padr + 0] - qpos_spring[worldid, padr + 0], + qpos_in[worldid, padr + 1] - qpos_spring[worldid, padr + 1], + qpos_in[worldid, padr + 2] - qpos_spring[worldid, padr + 2], + ) + + # convert quaternion difference into angular "velocity" + quat1 = wp.quat( + qpos_in[worldid, padr + 3], + qpos_in[worldid, padr + 4], + qpos_in[worldid, padr + 5], + qpos_in[worldid, padr + 6], + ) + quat1 = wp.normalize(quat1) + + quat_spring = wp.quat( + qpos_spring[worldid, padr + 3], + qpos_spring[worldid, padr + 4], + qpos_spring[worldid, padr + 5], + qpos_spring[worldid, padr + 6], + ) + + dif1 = math.quat_sub(quat1, quat_spring) + + energy = wp.vec2( + 0.5 * stiffness * (wp.dot(dif0, dif0) + wp.dot(dif1, dif1)), + 0.0, + ) + + wp.atomic_add(energy_out, worldid, energy) + + elif jnttype == int(JointType.BALL.value): + quat = wp.quat( + qpos_in[worldid, padr + 0], + qpos_in[worldid, padr + 1], + qpos_in[worldid, padr + 2], + qpos_in[worldid, padr + 3], + ) + quat = wp.normalize(quat) + + quat_spring = wp.quat( + qpos_spring[worldid, padr + 0], + qpos_spring[worldid, padr + 1], + qpos_spring[worldid, padr + 2], + qpos_spring[worldid, padr + 3], + ) + + dif = math.quat_sub(quat, quat_spring) + energy = wp.vec2( + 0.5 * stiffness * wp.dot(dif, dif), + 0.0, + ) + wp.atomic_add(energy_out, worldid, energy) + elif jnttype == int(JointType.SLIDE.value) or jnttype == int(JointType.HINGE.value): + dif_ = qpos_in[worldid, padr] - qpos_spring[worldid, padr] + energy = wp.vec2( + 0.5 * stiffness * dif_ * dif_, + 0.0, + ) + wp.atomic_add(energy_out, worldid, energy) + + +@wp.kernel +def _energy_pos_passive_tendon( + # Model: + tendon_stiffness: wp.array2d(dtype=float), + tendon_lengthspring: wp.array2d(dtype=wp.vec2), + # Data in: + ten_length_in: wp.array2d(dtype=float), + # Data out: + energy_out: wp.array(dtype=wp.vec2), +): + worldid, tenid = wp.tid() + + stiffness = tendon_stiffness[worldid, tenid] + + if stiffness == 0.0: + return + + length = ten_length_in[worldid, tenid] + + # compute spring displacement + lengthspring = tendon_lengthspring[worldid, tenid] + lower = lengthspring[0] + upper = lengthspring[1] + + if length > upper: + displacement = upper - length + elif length < lower: + displacement = lower - length + else: + displacement = 0.0 + + energy = wp.vec2(0.5 * stiffness * displacement * displacement, 0.0) + wp.atomic_add(energy_out, worldid, energy) + + +def energy_pos(m: Model, d: Data): + """Position-dependent energy (potential).""" + wp.launch(_energy_pos_zero, dim=(d.nworld,), outputs=[d.energy]) + + # init potential energy: -sum_i(body_i.mass * dot(gravity, body_i.pos)) + if not m.opt.disableflags & DisableBit.GRAVITY: + wp.launch( + _energy_pos_gravity, dim=(d.nworld, m.nbody - 1), inputs=[m.opt.gravity, m.body_mass, d.xipos], outputs=[d.energy] + ) + + if not m.opt.disableflags & DisableBit.PASSIVE: + # add joint-level springs + wp.launch( + _energy_pos_passive_joint, + dim=(d.nworld, m.njnt), + inputs=[ + m.qpos_spring, + m.jnt_type, + m.jnt_qposadr, + m.jnt_stiffness, + d.qpos, + ], + outputs=[d.energy], + ) + + # add tendon-level springs + if m.ntendon: + wp.launch( + _energy_pos_passive_tendon, + dim=(d.nworld, m.ntendon), + inputs=[ + m.tendon_stiffness, + m.tendon_lengthspring, + d.ten_length, + ], + outputs=[d.energy], + ) + + # TODO(team): flex + + +@cache_kernel +def _energy_vel_kinetic(nv: int): + @nested_kernel + def energy_vel_kinetic( + # Data in: + qvel_in: wp.array2d(dtype=float), + # In: + Mqvel: wp.array2d(dtype=float), + # Out: + energy_out: wp.array(dtype=wp.vec2), + ): + worldid = wp.tid() + + qvel_tile = wp.tile_load(qvel_in[worldid], shape=wp.static(nv)) + Mqvel_tile = wp.tile_load(Mqvel[worldid], shape=wp.static(nv)) + + # qvel * (M @ qvel) + qvelMqvel_tile = wp.tile_map(wp.mul, qvel_tile, Mqvel_tile) + + # sum(qvel * (M @ qvel)) + quadratic_tile = wp.tile_reduce(wp.add, qvelMqvel_tile) + + energy_out[worldid][1] = 0.5 * quadratic_tile[0] + + return energy_vel_kinetic + + +def energy_vel(m: Model, d: Data): + """Velocity-dependent energy (kinetic).""" + + # kinetic energy: 0.5 * qvel.T @ M @ qvel + + # M @ qvel + skip = wp.zeros(d.nworld, dtype=bool) + support.mul_m(m, d, d.efc.mv, d.qvel, skip) + + wp.launch_tiled( + _energy_vel_kinetic(m.nv), + dim=(d.nworld,), + inputs=[d.qvel, d.efc.mv], + outputs=[d.energy], + block_dim=m.block_dim.energy_vel_kinetic, + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py new file mode 100644 index 00000000..458dea1f --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py @@ -0,0 +1,411 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Tests for sensor functions.""" + +import mujoco +import numpy as np +import warp as wp +from absl.testing import absltest +from absl.testing import parameterized + +import mujoco_warp as mjwarp + +from mujoco.mjx.third_party.mujoco_warp._src import test_util + +# tolerance for difference between MuJoCo and MJWarp calculations - mostly +# due to float precision +_TOLERANCE = 5e-5 + + +def _assert_eq(a, b, name): + tol = _TOLERANCE * 10 # avoid test noise + err_msg = f"mismatch: {name}" + np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol) + + +class SensorTest(parameterized.TestCase): + def test_sensor(self): + """Test sensors.""" + mjm, mjd, m, d = test_util.fixture( + xml=""" + + + """, + keyframe=0, + kick=True, + ) + + d.sensordata.zero_() + + mjwarp.sensor_pos(m, d) + mjwarp.sensor_vel(m, d) + mjwarp.sensor_acc(m, d) + + _assert_eq(d.sensordata.numpy()[0], mjd.sensordata, "sensordata") + + def test_rangefinder(self): + """Test rangefinder.""" + for keyframe in range(2): + _, mjd, m, d = test_util.fixture( + xml=""" + + + + + + + + + + + + + + + + + + + + + + + + """, + ) + + d.sensordata.zero_() + + mjwarp.sensor_pos(m, d) + + _assert_eq(d.sensordata.numpy()[0], mjd.sensordata, "sensordata") + + def test_touch_sensor(self): + """Test touch sensor.""" + for keyframe in range(2): + _, mjd, m, d = test_util.fixture( + xml=""" + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + """, + keyframe=keyframe, + ) + + d.sensordata.zero_() + + mjwarp.sensor_acc(m, d) + + _assert_eq(d.sensordata.numpy()[0], mjd.sensordata, "sensordata") + + def test_tendon_sensor(self): + """Test tendon sensors.""" + _, mjd, m, d = test_util.fixture("tendon/fixed.xml", keyframe=0, sparse=False) + + d.sensordata.zero_() + + mjwarp.sensor_pos(m, d) + mjwarp.sensor_vel(m, d) + + _assert_eq(d.sensordata.numpy()[0], mjd.sensordata, "sensordata") + + @parameterized.parameters("humanoid/humanoid.xml", "constraints.xml") + def test_energy(self, xml): + mjm, mjd, m, d = test_util.fixture(xml, constraint=False, kick=True) + + d.energy.zero_() + + mujoco.mj_energyPos(mjm, mjd) + mjwarp.energy_pos(m, d) + + _assert_eq(d.energy.numpy()[0][0], mjd.energy[0], "potential energy") + + mujoco.mj_energyVel(mjm, mjd) + mjwarp.energy_vel(m, d) + + _assert_eq(d.energy.numpy()[0][1], mjd.energy[1], "kinetic energy") + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py new file mode 100644 index 00000000..b22ee540 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py @@ -0,0 +1,3128 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src import math +from mujoco.mjx.third_party.mujoco_warp._src import support +from mujoco.mjx.third_party.mujoco_warp._src import util_misc +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import CamLightType +from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import JointType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import ObjType +from mujoco.mjx.third_party.mujoco_warp._src.types import TileSet +from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType +from mujoco.mjx.third_party.mujoco_warp._src.types import WrapType +from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 +from mujoco.mjx.third_party.mujoco_warp._src.types import vec10 +from mujoco.mjx.third_party.mujoco_warp._src.types import vec11 +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel + +wp.set_module_options({"enable_backward": False}) + + +@wp.kernel +def _kinematics_root( + # Data out: + xpos_out: wp.array2d(dtype=wp.vec3), + xquat_out: wp.array2d(dtype=wp.quat), + xmat_out: wp.array2d(dtype=wp.mat33), + xipos_out: wp.array2d(dtype=wp.vec3), + ximat_out: wp.array2d(dtype=wp.mat33), +): + worldid = wp.tid() + xpos_out[worldid, 0] = wp.vec3(0.0) + xquat_out[worldid, 0] = wp.quat(1.0, 0.0, 0.0, 0.0) + xipos_out[worldid, 0] = wp.vec3(0.0) + xmat_out[worldid, 0] = wp.identity(n=3, dtype=wp.float32) + ximat_out[worldid, 0] = wp.identity(n=3, dtype=wp.float32) + + +@wp.kernel +def _kinematics_level( + # Model: + qpos0: wp.array2d(dtype=float), + body_parentid: wp.array(dtype=int), + body_jntnum: wp.array(dtype=int), + body_jntadr: wp.array(dtype=int), + body_pos: wp.array2d(dtype=wp.vec3), + body_quat: wp.array2d(dtype=wp.quat), + body_ipos: wp.array2d(dtype=wp.vec3), + body_iquat: wp.array2d(dtype=wp.quat), + jnt_type: wp.array(dtype=int), + jnt_qposadr: wp.array(dtype=int), + jnt_pos: wp.array2d(dtype=wp.vec3), + jnt_axis: wp.array2d(dtype=wp.vec3), + # Data in: + qpos_in: wp.array2d(dtype=float), + xpos_in: wp.array2d(dtype=wp.vec3), + xquat_in: wp.array2d(dtype=wp.quat), + xmat_in: wp.array2d(dtype=wp.mat33), + # In: + body_tree_: wp.array(dtype=int), + # Data out: + xpos_out: wp.array2d(dtype=wp.vec3), + xquat_out: wp.array2d(dtype=wp.quat), + xmat_out: wp.array2d(dtype=wp.mat33), + xipos_out: wp.array2d(dtype=wp.vec3), + ximat_out: wp.array2d(dtype=wp.mat33), + xanchor_out: wp.array2d(dtype=wp.vec3), + xaxis_out: wp.array2d(dtype=wp.vec3), +): + worldid, nodeid = wp.tid() + bodyid = body_tree_[nodeid] + jntadr = body_jntadr[bodyid] + jntnum = body_jntnum[bodyid] + qpos = qpos_in[worldid] + + if jntnum == 0: + # no joints - apply fixed translation and rotation relative to parent + pid = body_parentid[bodyid] + xpos = (xmat_in[worldid, pid] * body_pos[worldid, bodyid]) + xpos_in[worldid, pid] + xquat = math.mul_quat(xquat_in[worldid, pid], body_quat[worldid, bodyid]) + elif jntnum == 1 and jnt_type[jntadr] == wp.static(JointType.FREE.value): + # free joint + qadr = jnt_qposadr[jntadr] + xpos = wp.vec3(qpos[qadr], qpos[qadr + 1], qpos[qadr + 2]) + xquat = wp.quat(qpos[qadr + 3], qpos[qadr + 4], qpos[qadr + 5], qpos[qadr + 6]) + xquat = wp.normalize(xquat) + xanchor_out[worldid, jntadr] = xpos + xaxis_out[worldid, jntadr] = jnt_axis[worldid, jntadr] + else: + # regular or no joints + # apply fixed translation and rotation relative to parent + pid = body_parentid[bodyid] + xpos = (xmat_in[worldid, pid] * body_pos[worldid, bodyid]) + xpos_in[worldid, pid] + xquat = math.mul_quat(xquat_in[worldid, pid], body_quat[worldid, bodyid]) + + for _ in range(jntnum): + qadr = jnt_qposadr[jntadr] + jnt_type_ = jnt_type[jntadr] + jnt_axis_ = jnt_axis[worldid, jntadr] + xanchor = math.rot_vec_quat(jnt_pos[worldid, jntadr], xquat) + xpos + xaxis = math.rot_vec_quat(jnt_axis_, xquat) + + if jnt_type_ == wp.static(JointType.BALL.value): + qloc = wp.quat( + qpos[qadr + 0], + qpos[qadr + 1], + qpos[qadr + 2], + qpos[qadr + 3], + ) + qloc = wp.normalize(qloc) + xquat = math.mul_quat(xquat, qloc) + # correct for off-center rotation + xpos = xanchor - math.rot_vec_quat(jnt_pos[worldid, jntadr], xquat) + elif jnt_type_ == wp.static(JointType.SLIDE.value): + xpos += xaxis * (qpos[qadr] - qpos0[worldid, qadr]) + elif jnt_type_ == wp.static(JointType.HINGE.value): + qpos0_ = qpos0[worldid, qadr] + qloc_ = math.axis_angle_to_quat(jnt_axis_, qpos[qadr] - qpos0_) + xquat = math.mul_quat(xquat, qloc_) + # correct for off-center rotation + xpos = xanchor - math.rot_vec_quat(jnt_pos[worldid, jntadr], xquat) + + xanchor_out[worldid, jntadr] = xanchor + xaxis_out[worldid, jntadr] = xaxis + jntadr += 1 + + xpos_out[worldid, bodyid] = xpos + xquat_out[worldid, bodyid] = wp.normalize(xquat) + xmat_out[worldid, bodyid] = math.quat_to_mat(xquat) + xipos_out[worldid, bodyid] = xpos + math.rot_vec_quat(body_ipos[worldid, bodyid], xquat) + ximat_out[worldid, bodyid] = math.quat_to_mat(math.mul_quat(xquat, body_iquat[worldid, bodyid])) + + +@wp.kernel +def _geom_local_to_global( + # Model: + geom_bodyid: wp.array(dtype=int), + geom_pos: wp.array2d(dtype=wp.vec3), + geom_quat: wp.array2d(dtype=wp.quat), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + xquat_in: wp.array2d(dtype=wp.quat), + geom_skip_in: wp.array(dtype=bool), + # Data out: + geom_skip_out: wp.array(dtype=bool), + geom_xpos_out: wp.array2d(dtype=wp.vec3), + geom_xmat_out: wp.array2d(dtype=wp.mat33), +): + worldid, geomid = wp.tid() + bodyid = geom_bodyid[geomid] + if not geom_skip_in[geomid]: + # Calculate only if necessary + xpos = xpos_in[worldid, bodyid] + xquat = xquat_in[worldid, bodyid] + geom_xpos_out[worldid, geomid] = xpos + math.rot_vec_quat(geom_pos[worldid, geomid], xquat) + geom_xmat_out[worldid, geomid] = math.quat_to_mat(math.mul_quat(xquat, geom_quat[worldid, geomid])) + + if bodyid == 0: + # static geom pose are calculated only once + geom_skip_out[geomid] = True + + +@wp.kernel +def _site_local_to_global( + # Model: + site_bodyid: wp.array(dtype=int), + site_pos: wp.array2d(dtype=wp.vec3), + site_quat: wp.array2d(dtype=wp.quat), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + xquat_in: wp.array2d(dtype=wp.quat), + # Data out: + site_xpos_out: wp.array2d(dtype=wp.vec3), + site_xmat_out: wp.array2d(dtype=wp.mat33), +): + worldid, siteid = wp.tid() + bodyid = site_bodyid[siteid] + xpos = xpos_in[worldid, bodyid] + xquat = xquat_in[worldid, bodyid] + site_xpos_out[worldid, siteid] = xpos + math.rot_vec_quat(site_pos[worldid, siteid], xquat) + site_xmat_out[worldid, siteid] = math.quat_to_mat(math.mul_quat(xquat, site_quat[worldid, siteid])) + + +@wp.kernel +def _flex_vertices( + # Model: + flex_vertbodyid: wp.array(dtype=int), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + # Data out: + flexvert_xpos_out: wp.array2d(dtype=wp.vec3), +): + worldid, vertid = wp.tid() + flexvert_xpos_out[worldid, vertid] = xpos_in[worldid, flex_vertbodyid[vertid]] + + +@wp.kernel +def _flex_edges( + # Model: + body_dofadr: wp.array(dtype=int), + flex_vertadr: wp.array(dtype=int), + flex_vertbodyid: wp.array(dtype=int), + flex_edge: wp.array(dtype=wp.vec2i), + # Data in: + qvel_in: wp.array2d(dtype=float), + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), + # Data out: + flexedge_length_out: wp.array2d(dtype=float), + flexedge_velocity_out: wp.array2d(dtype=float), +): + worldid, edgeid = wp.tid() + f = 0 # TODO(quaglino): get f from edgeid + vbase = flex_vertadr[f] + v = flex_edge[edgeid] + pos1 = flexvert_xpos_in[worldid, vbase + v[0]] + pos2 = flexvert_xpos_in[worldid, vbase + v[1]] + vec = pos2 - pos1 + vecnorm = wp.length(vec) + flexedge_length_out[worldid, edgeid] = vecnorm + # TODO(quaglino): use Jacobian + i = body_dofadr[flex_vertbodyid[vbase + v[0]]] + j = body_dofadr[flex_vertbodyid[vbase + v[1]]] + vel1 = wp.vec3(qvel_in[worldid, i], qvel_in[worldid, i + 1], qvel_in[worldid, i + 2]) + vel2 = wp.vec3(qvel_in[worldid, j], qvel_in[worldid, j + 1], qvel_in[worldid, j + 2]) + flexedge_velocity_out[worldid, edgeid] = wp.dot(vel2 - vel1, vec) / vecnorm + + +@wp.kernel +def _mocap( + # Model: + body_ipos: wp.array2d(dtype=wp.vec3), + body_iquat: wp.array2d(dtype=wp.quat), + mocap_bodyid: wp.array(dtype=int), + # Data in: + mocap_pos_in: wp.array2d(dtype=wp.vec3), + mocap_quat_in: wp.array2d(dtype=wp.quat), + # Data out: + xpos_out: wp.array2d(dtype=wp.vec3), + xquat_out: wp.array2d(dtype=wp.quat), + xmat_out: wp.array2d(dtype=wp.mat33), + xipos_out: wp.array2d(dtype=wp.vec3), + ximat_out: wp.array2d(dtype=wp.mat33), +): + worldid, mocapid = wp.tid() + bodyid = mocap_bodyid[mocapid] + mocap_quat = wp.normalize(mocap_quat_in[worldid, mocapid]) + xpos = mocap_pos_in[worldid, mocapid] + xpos_out[worldid, bodyid] = xpos + xquat_out[worldid, bodyid] = mocap_quat + xmat_out[worldid, bodyid] = math.quat_to_mat(mocap_quat) + xipos_out[worldid, bodyid] = xpos + math.rot_vec_quat(body_ipos[worldid, bodyid], mocap_quat) + ximat_out[worldid, bodyid] = math.quat_to_mat(math.mul_quat(mocap_quat, body_iquat[worldid, bodyid])) + + +@event_scope +def kinematics(m: Model, d: Data): + """ + Computes forward kinematics for all bodies, sites, geoms, and flexible elements. + + This function updates the global positions and orientations of all bodies, as well as the + derived positions and orientations of geoms, sites, and flexible elements, based on the + current joint positions and any attached mocap bodies. + """ + wp.launch(_kinematics_root, dim=(d.nworld), inputs=[], outputs=[d.xpos, d.xquat, d.xmat, d.xipos, d.ximat]) + + for i in range(1, len(m.body_tree)): + body_tree = m.body_tree[i] + wp.launch( + _kinematics_level, + dim=(d.nworld, body_tree.size), + inputs=[ + m.qpos0, + m.body_parentid, + m.body_jntnum, + m.body_jntadr, + m.body_pos, + m.body_quat, + m.body_ipos, + m.body_iquat, + m.jnt_type, + m.jnt_qposadr, + m.jnt_pos, + m.jnt_axis, + d.qpos, + d.xpos, + d.xquat, + d.xmat, + body_tree, + ], + outputs=[d.xpos, d.xquat, d.xmat, d.xipos, d.ximat, d.xanchor, d.xaxis], + ) + + wp.launch( + _mocap, + dim=(d.nworld, m.nmocap), + inputs=[m.body_ipos, m.body_iquat, m.mocap_bodyid, d.mocap_pos, d.mocap_quat], + outputs=[d.xpos, d.xquat, d.xmat, d.xipos, d.ximat], + ) + + wp.launch( + _geom_local_to_global, + dim=(d.nworld, m.ngeom), + inputs=[m.geom_bodyid, m.geom_pos, m.geom_quat, d.xpos, d.xquat, d.geom_skip], + outputs=[d.geom_skip, d.geom_xpos, d.geom_xmat], + ) + + wp.launch( + _site_local_to_global, + dim=(d.nworld, m.nsite), + inputs=[m.site_bodyid, m.site_pos, m.site_quat, d.xpos, d.xquat], + outputs=[d.site_xpos, d.site_xmat], + ) + + wp.launch(_flex_vertices, dim=(d.nworld, m.nflexvert), inputs=[m.flex_vertbodyid, d.xpos], outputs=[d.flexvert_xpos]) + wp.launch( + _flex_edges, + dim=(d.nworld, m.nflexedge), + inputs=[m.body_dofadr, m.flex_vertadr, m.flex_vertbodyid, m.flex_edge, d.qvel, d.flexvert_xpos], + outputs=[d.flexedge_length, d.flexedge_velocity], + ) + + +@wp.kernel +def _subtree_com_init( + # Model: + body_mass: wp.array2d(dtype=float), + # Data in: + xipos_in: wp.array2d(dtype=wp.vec3), + # Data out: + xipos_out: wp.array2d(dtype=wp.vec3), +): + worldid, bodyid = wp.tid() + xipos_out[worldid, bodyid] = xipos_in[worldid, bodyid] * body_mass[worldid, bodyid] + + +@wp.kernel +def _subtree_com_acc( + # Model: + body_parentid: wp.array(dtype=int), + # Data in: + subtree_com_in: wp.array2d(dtype=wp.vec3), + # In: + body_tree_: wp.array(dtype=int), + # Data out: + subtree_com_out: wp.array2d(dtype=wp.vec3), +): + worldid, nodeid = wp.tid() + bodyid = body_tree_[nodeid] + pid = body_parentid[bodyid] + wp.atomic_add(subtree_com_out, worldid, pid, subtree_com_in[worldid, bodyid]) + + +@wp.kernel +def _subtree_div( + # Model: + subtree_mass: wp.array2d(dtype=float), + # Data out: + subtree_com_out: wp.array2d(dtype=wp.vec3), +): + worldid, bodyid = wp.tid() + subtree_com_out[worldid, bodyid] /= subtree_mass[worldid, bodyid] + + +@wp.kernel +def _cinert( + # Model: + body_rootid: wp.array(dtype=int), + body_mass: wp.array2d(dtype=float), + body_inertia: wp.array2d(dtype=wp.vec3), + # Data in: + xipos_in: wp.array2d(dtype=wp.vec3), + ximat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + # Data out: + cinert_out: wp.array2d(dtype=vec10), +): + worldid, bodyid = wp.tid() + mat = ximat_in[worldid, bodyid] + inert = body_inertia[worldid, bodyid] + mass = body_mass[worldid, bodyid] + dif = xipos_in[worldid, bodyid] - subtree_com_in[worldid, body_rootid[bodyid]] + # express inertia in com-based frame (mju_inertCom) + + res = vec10() + # res_rot = mat * diag(inert) * mat' + tmp = mat @ wp.diag(inert) @ wp.transpose(mat) + res[0] = tmp[0, 0] + res[1] = tmp[1, 1] + res[2] = tmp[2, 2] + res[3] = tmp[0, 1] + res[4] = tmp[0, 2] + res[5] = tmp[1, 2] + # res_rot -= mass * dif_cross * dif_cross + res[0] += mass * (dif[1] * dif[1] + dif[2] * dif[2]) + res[1] += mass * (dif[0] * dif[0] + dif[2] * dif[2]) + res[2] += mass * (dif[0] * dif[0] + dif[1] * dif[1]) + res[3] -= mass * dif[0] * dif[1] + res[4] -= mass * dif[0] * dif[2] + res[5] -= mass * dif[1] * dif[2] + # res_tran = mass * dif + res[6] = mass * dif[0] + res[7] = mass * dif[1] + res[8] = mass * dif[2] + # res_mass = mass + res[9] = mass + + cinert_out[worldid, bodyid] = res + + +@wp.kernel +def _cdof( + # Model: + body_rootid: wp.array(dtype=int), + jnt_type: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + jnt_bodyid: wp.array(dtype=int), + # Data in: + xmat_in: wp.array2d(dtype=wp.mat33), + xanchor_in: wp.array2d(dtype=wp.vec3), + xaxis_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + # Data out: + cdof_out: wp.array2d(dtype=wp.spatial_vector), +): + worldid, jntid = wp.tid() + bodyid = jnt_bodyid[jntid] + dofid = jnt_dofadr[jntid] + jnt_type_ = jnt_type[jntid] + xaxis = xaxis_in[worldid, jntid] + xmat = wp.transpose(xmat_in[worldid, bodyid]) + + # compute com-anchor vector + offset = subtree_com_in[worldid, body_rootid[bodyid]] - xanchor_in[worldid, jntid] + + res = cdof_out[worldid] + if jnt_type_ == wp.static(JointType.FREE.value): + res[dofid + 0] = wp.spatial_vector(0.0, 0.0, 0.0, 1.0, 0.0, 0.0) + res[dofid + 1] = wp.spatial_vector(0.0, 0.0, 0.0, 0.0, 1.0, 0.0) + res[dofid + 2] = wp.spatial_vector(0.0, 0.0, 0.0, 0.0, 0.0, 1.0) + # I_3 rotation in child frame (assume no subsequent rotations) + res[dofid + 3] = wp.spatial_vector(xmat[0], wp.cross(xmat[0], offset)) + res[dofid + 4] = wp.spatial_vector(xmat[1], wp.cross(xmat[1], offset)) + res[dofid + 5] = wp.spatial_vector(xmat[2], wp.cross(xmat[2], offset)) + elif jnt_type_ == wp.static(JointType.BALL.value): # ball + # I_3 rotation in child frame (assume no subsequent rotations) + res[dofid + 0] = wp.spatial_vector(xmat[0], wp.cross(xmat[0], offset)) + res[dofid + 1] = wp.spatial_vector(xmat[1], wp.cross(xmat[1], offset)) + res[dofid + 2] = wp.spatial_vector(xmat[2], wp.cross(xmat[2], offset)) + elif jnt_type_ == wp.static(JointType.SLIDE.value): + res[dofid] = wp.spatial_vector(wp.vec3(0.0), xaxis) + elif jnt_type_ == wp.static(JointType.HINGE.value): # hinge + res[dofid] = wp.spatial_vector(xaxis, wp.cross(xaxis, offset)) + + +@event_scope +def com_pos(m: Model, d: Data): + """ + Computes subtree center of mass positions. Transforms inertia and motion to global frame + centered at subtree CoM. + + Accumulates the mass-weighted positions up the kinematic tree, divides by total mass, and + computes composite inertias and motion degrees of freedom in the subtree CoM frame. + """ + wp.launch(_subtree_com_init, dim=(d.nworld, m.nbody), inputs=[m.body_mass, d.xipos, d.subtree_com]) + + for i in reversed(range(len(m.body_tree))): + body_tree = m.body_tree[i] + wp.launch( + _subtree_com_acc, + dim=(d.nworld, body_tree.size), + inputs=[m.body_parentid, d.subtree_com, body_tree], + outputs=[d.subtree_com], + ) + + wp.launch(_subtree_div, dim=(d.nworld, m.nbody), inputs=[m.subtree_mass], outputs=[d.subtree_com]) + wp.launch( + _cinert, + dim=(d.nworld, m.nbody), + inputs=[m.body_rootid, m.body_mass, m.body_inertia, d.xipos, d.ximat, d.subtree_com], + outputs=[d.cinert], + ) + wp.launch( + _cdof, + dim=(d.nworld, m.njnt), + inputs=[m.body_rootid, m.jnt_type, m.jnt_dofadr, m.jnt_bodyid, d.xmat, d.xanchor, d.xaxis, d.subtree_com], + outputs=[d.cdof], + ) + + +@wp.kernel +def _cam_local_to_global( + # Model: + cam_bodyid: wp.array(dtype=int), + cam_pos: wp.array2d(dtype=wp.vec3), + cam_quat: wp.array2d(dtype=wp.quat), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + xquat_in: wp.array2d(dtype=wp.quat), + # Data out: + cam_xpos_out: wp.array2d(dtype=wp.vec3), + cam_xmat_out: wp.array2d(dtype=wp.mat33), +): + """Fixed cameras.""" + worldid, camid = wp.tid() + bodyid = cam_bodyid[camid] + xpos = xpos_in[worldid, bodyid] + xquat = xquat_in[worldid, bodyid] + cam_xpos_out[worldid, camid] = xpos + math.rot_vec_quat(cam_pos[worldid, camid], xquat) + cam_xmat_out[worldid, camid] = math.quat_to_mat(math.mul_quat(xquat, cam_quat[worldid, camid])) + + +@wp.kernel +def _cam_fn( + # Model: + cam_mode: wp.array(dtype=int), + cam_bodyid: wp.array(dtype=int), + cam_targetbodyid: wp.array(dtype=int), + cam_poscom0: wp.array2d(dtype=wp.vec3), + cam_pos0: wp.array2d(dtype=wp.vec3), + cam_mat0: wp.array2d(dtype=wp.mat33), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + # Data out: + cam_xpos_out: wp.array2d(dtype=wp.vec3), + cam_xmat_out: wp.array2d(dtype=wp.mat33), +): + worldid, camid = wp.tid() + is_target_cam = (cam_mode[camid] == wp.static(CamLightType.TARGETBODY.value)) or ( + cam_mode[camid] == wp.static(CamLightType.TARGETBODYCOM.value) + ) + invalid_target = is_target_cam and (cam_targetbodyid[camid] < 0) + if invalid_target: + return + elif cam_mode[camid] == wp.static(CamLightType.TRACK.value): + cam_xmat_out[worldid, camid] = cam_mat0[worldid, camid] + body_xpos = xpos_in[worldid, cam_bodyid[camid]] + cam_xpos_out[worldid, camid] = body_xpos + cam_pos0[worldid, camid] + elif cam_mode[camid] == wp.static(CamLightType.TRACKCOM.value): + cam_xmat_out[worldid, camid] = cam_mat0[worldid, camid] + cam_xpos_out[worldid, camid] = subtree_com_in[worldid, cam_bodyid[camid]] + cam_poscom0[worldid, camid] + elif cam_mode[camid] == wp.static(CamLightType.TARGETBODY.value) or cam_mode[camid] == wp.static( + CamLightType.TARGETBODYCOM.value + ): + pos = xpos_in[worldid, cam_targetbodyid[camid]] + if cam_mode[camid] == wp.static(CamLightType.TARGETBODYCOM.value): + pos = subtree_com_in[worldid, cam_targetbodyid[camid]] + # zaxis = -desired camera direction, in global frame + mat_3 = wp.normalize(cam_xpos_out[worldid, camid] - pos) + # xaxis: orthogonal to zaxis and to (0,0,1) + mat_1 = wp.normalize(wp.cross(wp.vec3(0.0, 0.0, 1.0), mat_3)) + mat_2 = wp.normalize(wp.cross(mat_3, mat_1)) + # fmt: off + cam_xmat_out[worldid, camid] = wp.mat33( + mat_1[0], mat_2[0], mat_3[0], + mat_1[1], mat_2[1], mat_3[1], + mat_1[2], mat_2[2], mat_3[2] + ) + # fmt: on + + +@wp.kernel +def _light_local_to_global( + # Model: + light_bodyid: wp.array(dtype=int), + light_pos: wp.array2d(dtype=wp.vec3), + light_dir: wp.array2d(dtype=wp.vec3), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + xquat_in: wp.array2d(dtype=wp.quat), + # Data out: + light_xpos_out: wp.array2d(dtype=wp.vec3), + light_xdir_out: wp.array2d(dtype=wp.vec3), +): + """Fixed lights.""" + worldid, lightid = wp.tid() + bodyid = light_bodyid[lightid] + xpos = xpos_in[worldid, bodyid] + xquat = xquat_in[worldid, bodyid] + light_xpos_out[worldid, lightid] = xpos + math.rot_vec_quat(light_pos[worldid, lightid], xquat) + light_xdir_out[worldid, lightid] = math.rot_vec_quat(light_dir[worldid, lightid], xquat) + + +@wp.kernel +def _light_fn( + # Model: + light_mode: wp.array(dtype=int), + light_bodyid: wp.array(dtype=int), + light_targetbodyid: wp.array(dtype=int), + light_poscom0: wp.array2d(dtype=wp.vec3), + light_pos0: wp.array2d(dtype=wp.vec3), + light_dir0: wp.array2d(dtype=wp.vec3), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + light_xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + # Data out: + light_xpos_out: wp.array2d(dtype=wp.vec3), + light_xdir_out: wp.array2d(dtype=wp.vec3), +): + worldid, lightid = wp.tid() + is_target_light = (light_mode[lightid] == wp.static(CamLightType.TARGETBODY.value)) or ( + light_mode[lightid] == wp.static(CamLightType.TARGETBODYCOM.value) + ) + invalid_target = is_target_light and (light_targetbodyid[lightid] < 0) + if invalid_target: + return + elif light_mode[lightid] == wp.static(CamLightType.TRACK.value): + light_xdir_out[worldid, lightid] = light_dir0[worldid, lightid] + body_xpos = xpos_in[worldid, light_bodyid[lightid]] + light_xpos_out[worldid, lightid] = body_xpos + light_pos0[worldid, lightid] + elif light_mode[lightid] == wp.static(CamLightType.TRACKCOM.value): + light_xdir_out[worldid, lightid] = light_dir0[worldid, lightid] + light_xpos_out[worldid, lightid] = subtree_com_in[worldid, light_bodyid[lightid]] + light_poscom0[worldid, lightid] + elif light_mode[lightid] == wp.static(CamLightType.TARGETBODY.value) or light_mode[lightid] == wp.static( + CamLightType.TARGETBODYCOM.value + ): + pos = xpos_in[worldid, light_targetbodyid[lightid]] + if light_mode[lightid] == wp.static(CamLightType.TARGETBODYCOM.value): + pos = subtree_com_in[worldid, light_targetbodyid[lightid]] + light_xdir_out[worldid, lightid] = pos - light_xpos_in[worldid, lightid] + light_xdir_out[worldid, lightid] = wp.normalize(light_xdir_out[worldid, lightid]) + + +@event_scope +def camlight(m: Model, d: Data): + """ + Computes camera and light positions and orientations. + + Updates the global positions and orientations for all cameras and lights in the model, + including special handling for tracking and target modes. + """ + wp.launch( + _cam_local_to_global, + dim=(d.nworld, m.ncam), + inputs=[m.cam_bodyid, m.cam_pos, m.cam_quat, d.xpos, d.xquat], + outputs=[d.cam_xpos, d.cam_xmat], + ) + wp.launch( + _cam_fn, + dim=(d.nworld, m.ncam), + inputs=[m.cam_mode, m.cam_bodyid, m.cam_targetbodyid, m.cam_poscom0, m.cam_pos0, m.cam_mat0, d.xpos, d.subtree_com], + outputs=[d.cam_xpos, d.cam_xmat], + ) + wp.launch( + _light_local_to_global, + dim=(d.nworld, m.nlight), + inputs=[m.light_bodyid, m.light_pos, m.light_dir, d.xpos, d.xquat], + outputs=[d.light_xpos, d.light_xdir], + ) + wp.launch( + _light_fn, + dim=(d.nworld, m.nlight), + inputs=[ + m.light_mode, + m.light_bodyid, + m.light_targetbodyid, + m.light_poscom0, + m.light_pos0, + m.light_dir0, + d.xpos, + d.light_xpos, + d.subtree_com, + ], + outputs=[d.light_xpos, d.light_xdir], + ) + + +@wp.kernel +def _crb_accumulate( + # Model: + body_parentid: wp.array(dtype=int), + # Data in: + crb_in: wp.array2d(dtype=vec10), + # In: + body_tree_: wp.array(dtype=int), + # Data out: + crb_out: wp.array2d(dtype=vec10), +): + worldid, nodeid = wp.tid() + bodyid = body_tree_[nodeid] + pid = body_parentid[bodyid] + if pid == 0: + return + wp.atomic_add(crb_out, worldid, pid, crb_in[worldid, bodyid]) + + +@wp.kernel +def _qM_sparse( + # Model: + dof_bodyid: wp.array(dtype=int), + dof_parentid: wp.array(dtype=int), + dof_Madr: wp.array(dtype=int), + dof_armature: wp.array2d(dtype=float), + # Data in: + cdof_in: wp.array2d(dtype=wp.spatial_vector), + crb_in: wp.array2d(dtype=vec10), + # Data out: + qM_out: wp.array3d(dtype=float), +): + worldid, dofid = wp.tid() + madr_ij = dof_Madr[dofid] + bodyid = dof_bodyid[dofid] + + # init M(i,i) with armature inertia + qM_out[worldid, 0, madr_ij] = dof_armature[worldid, dofid] + + # precompute buf = crb_body_i * cdof_i + buf = math.inert_vec(crb_in[worldid, bodyid], cdof_in[worldid, dofid]) + + # sparse backward pass over ancestors + while dofid >= 0: + qM_out[worldid, 0, madr_ij] += wp.dot(cdof_in[worldid, dofid], buf) + madr_ij += 1 + dofid = dof_parentid[dofid] + + +@wp.kernel +def _qM_dense( + # Model: + dof_bodyid: wp.array(dtype=int), + dof_parentid: wp.array(dtype=int), + dof_armature: wp.array2d(dtype=float), + # Data in: + cdof_in: wp.array2d(dtype=wp.spatial_vector), + crb_in: wp.array2d(dtype=vec10), + # Data out: + qM_out: wp.array3d(dtype=float), +): + worldid, dofid = wp.tid() + bodyid = dof_bodyid[dofid] + # init M(i,i) with armature inertia + M = dof_armature[worldid, dofid] + + # precompute buf = crb_body_i * cdof_i + buf = math.inert_vec(crb_in[worldid, bodyid], cdof_in[worldid, dofid]) + M += wp.dot(cdof_in[worldid, dofid], buf) + + qM_out[worldid, dofid, dofid] = M + + # sparse backward pass over ancestors + dofidi = dofid + dofid = dof_parentid[dofid] + while dofid >= 0: + qMij = wp.dot(cdof_in[worldid, dofid], buf) + qM_out[worldid, dofidi, dofid] += qMij + qM_out[worldid, dofid, dofidi] += qMij + dofid = dof_parentid[dofid] + + +@event_scope +def crb(m: Model, d: Data): + """ + Computes composite rigid body inertias for each body and the joint-space inertia matrix. + + Accumulates composite rigid body inertias up the kinematic tree and computes the + joint-space inertia matrix in either sparse or dense format, depending on model options. + """ + wp.copy(d.crb, d.cinert) + + for i in reversed(range(len(m.body_tree))): + body_tree = m.body_tree[i] + wp.launch(_crb_accumulate, dim=(d.nworld, body_tree.size), inputs=[m.body_parentid, d.crb, body_tree], outputs=[d.crb]) + + d.qM.zero_() + if m.opt.is_sparse: + wp.launch( + _qM_sparse, + dim=(d.nworld, m.nv), + inputs=[m.dof_bodyid, m.dof_parentid, m.dof_Madr, m.dof_armature, d.cdof, d.crb], + outputs=[d.qM], + ) + else: + wp.launch( + _qM_dense, dim=(d.nworld, m.nv), inputs=[m.dof_bodyid, m.dof_parentid, m.dof_armature, d.cdof, d.crb], outputs=[d.qM] + ) + + +@wp.kernel +def _tendon_armature( + # Model: + opt_is_sparse: bool, + dof_parentid: wp.array(dtype=int), + dof_Madr: wp.array(dtype=int), + tendon_armature: wp.array2d(dtype=float), + # Data in: + ten_J_in: wp.array3d(dtype=float), + # Data out: + qM_out: wp.array3d(dtype=float), +): + worldid, tenid, dofid = wp.tid() + + if opt_is_sparse: + madr_ij = dof_Madr[dofid] + + armature = tendon_armature[worldid, tenid] + + if armature == 0.0: + return + + ten_Ji = ten_J_in[worldid, tenid, dofid] + + if ten_Ji == 0.0: + return + + # sparse backward pass over ancestors + dofidi = dofid + while dofid >= 0: + if dofid != dofidi: + ten_Jj = ten_J_in[worldid, tenid, dofid] + else: + ten_Jj = ten_Ji + + qMij = armature * ten_Jj * ten_Ji + + if opt_is_sparse: + wp.atomic_add(qM_out[worldid, 0], madr_ij, qMij) + madr_ij += 1 + else: + wp.atomic_add(qM_out[worldid, dofidi], dofid, qMij) + if dofidi != dofid: + wp.atomic_add(qM_out[worldid, dofid], dofidi, qMij) + + dofid = dof_parentid[dofid] + + +@event_scope +def tendon_armature(m: Model, d: Data): + """Add tendon armature to qM.""" + wp.launch( + _tendon_armature, + dim=(d.nworld, m.ntendon, m.nv), + inputs=[m.opt.is_sparse, m.dof_parentid, m.dof_Madr, m.tendon_armature, d.ten_J], + outputs=[d.qM], + ) + + +@wp.kernel +def _copy_CSR( + # Model: + mapM2M: wp.array(dtype=int), + # In: + M_in: wp.array3d(dtype=float), + # Out: + L_out: wp.array3d(dtype=float), +): + worldid, ind = wp.tid() + L_out[worldid, 0, ind] = M_in[worldid, 0, mapM2M[ind]] + + +@wp.kernel +def _qLD_acc( + # Model: + M_rownnz: wp.array(dtype=int), + M_rowadr: wp.array(dtype=int), + # In: + qLD_updates_: wp.array(dtype=wp.vec3i), + L_in: wp.array3d(dtype=float), + # Out: + L_out: wp.array3d(dtype=float), +): + worldid, nodeid = wp.tid() + update = qLD_updates_[nodeid] + i, k, Madr_ki = update[0], update[1], update[2] + Madr_i = M_rowadr[i] # Address of row being updated + diag_k = M_rowadr[k] + M_rownnz[k] - 1 # Address of diagonal element of k + # tmp = M(k,i) / M(k,k) + tmp = L_out[worldid, 0, Madr_ki] / L_out[worldid, 0, diag_k] + for j in range(M_rownnz[i]): + # M(i,j) -= M(k,j) * tmp + wp.atomic_sub(L_out[worldid, 0], Madr_i + j, L_in[worldid, 0, M_rowadr[k] + j] * tmp) + # M(k,i) = tmp + L_out[worldid, 0, Madr_ki] = tmp + + +@wp.kernel +def _qLDiag_div( + # Model: + M_rownnz: wp.array(dtype=int), + M_rowadr: wp.array(dtype=int), + # In: + L_in: wp.array3d(dtype=float), + # Out: + D_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + diag_i = M_rowadr[dofid] + M_rownnz[dofid] - 1 # Address of diagonal element of i + D_out[worldid, dofid] = 1.0 / L_in[worldid, 0, diag_i] + + +def _factor_i_sparse(m: Model, d: Data, M: wp.array3d(dtype=float), L: wp.array3d(dtype=float), D: wp.array2d(dtype=float)): + """Sparse L'*D*L factorization of inertia-like matrix M, assumed spd.""" + wp.launch(_copy_CSR, dim=(d.nworld, m.nC), inputs=[m.mapM2M, M], outputs=[L]) + + for i in reversed(range(len(m.qLD_updates))): + qLD_updates = m.qLD_updates[i] + wp.launch(_qLD_acc, dim=(d.nworld, qLD_updates.size), inputs=[m.M_rownnz, m.M_rowadr, qLD_updates, L], outputs=[L]) + + wp.launch(_qLDiag_div, dim=(d.nworld, m.nv), inputs=[m.M_rownnz, m.M_rowadr, L], outputs=[D]) + + +@cache_kernel +def _tile_cholesky_factorize(tile: TileSet): + """Returns a kernel for dense Cholesky factorization of a tile.""" + + @nested_kernel + def cholesky_factorize( + # Data In: + qM_in: wp.array3d(dtype=float), + # In: + adr: wp.array(dtype=int), + # Out: + L_out: wp.array3d(dtype=float), + ): + worldid, nodeid = wp.tid() + TILE_SIZE = wp.static(tile.size) + + dofid = adr[nodeid] + M_tile = wp.tile_load(qM_in[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid)) + L_tile = wp.tile_cholesky(M_tile) + wp.tile_store(L_out[worldid], L_tile, offset=(dofid, dofid)) + + return cholesky_factorize + + +def _factor_i_dense(m: Model, d: Data, M: wp.array, L: wp.array): + """Dense Cholesky factorization of inertia-like matrix M, assumed spd.""" + for tile in m.qM_tiles: + wp.launch_tiled( + _tile_cholesky_factorize(tile), + dim=(d.nworld, tile.adr.size), + inputs=[M, tile.adr], + outputs=[L], + block_dim=m.block_dim.cholesky_factorize, + ) + + +@event_scope +def factor_m(m: Model, d: Data): + """Factorization of inertia-like matrix M, assumed spd.""" + if m.opt.is_sparse: + _factor_i_sparse(m, d, d.qM, d.qLD, d.qLDiagInv) + else: + _factor_i_dense(m, d, d.qM, d.qLD) + + +@wp.kernel +def _cacc_world( + # In: + gravity: wp.array(dtype=wp.vec3), + # Data out: + cacc_out: wp.array2d(dtype=wp.spatial_vector), +): + worldid = wp.tid() + cacc_out[worldid, 0] = wp.spatial_vector(wp.vec3(0.0), -gravity[worldid]) + + +def _rne_cacc_world(m: Model, d: Data): + if m.opt.disableflags & DisableBit.GRAVITY: + d.cacc.zero_() + else: + wp.launch(_cacc_world, dim=[d.nworld], inputs=[m.opt.gravity], outputs=[d.cacc]) + + +@wp.kernel +def _cacc( + # Model: + body_parentid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), + # Data in: + qvel_in: wp.array2d(dtype=float), + qacc_in: wp.array2d(dtype=float), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + cdof_dot_in: wp.array2d(dtype=wp.spatial_vector), + cacc_in: wp.array2d(dtype=wp.spatial_vector), + # In: + body_tree_: wp.array(dtype=int), + flg_acc: bool, + # Data out: + cacc_out: wp.array2d(dtype=wp.spatial_vector), +): + worldid, nodeid = wp.tid() + bodyid = body_tree_[nodeid] + dofnum = body_dofnum[bodyid] + pid = body_parentid[bodyid] + dofadr = body_dofadr[bodyid] + local_cacc = cacc_in[worldid, pid] + for i in range(dofnum): + local_cacc += cdof_dot_in[worldid, dofadr + i] * qvel_in[worldid, dofadr + i] + if flg_acc: + local_cacc += cdof_in[worldid, dofadr + i] * qacc_in[worldid, dofadr + i] + cacc_out[worldid, bodyid] = local_cacc + + +def _rne_cacc_forward(m: Model, d: Data, flg_acc: bool = False): + for body_tree in m.body_tree: + wp.launch( + _cacc, + dim=(d.nworld, body_tree.size), + inputs=[m.body_parentid, m.body_dofnum, m.body_dofadr, d.qvel, d.qacc, d.cdof, d.cdof_dot, d.cacc, body_tree, flg_acc], + outputs=[d.cacc], + ) + + +@wp.kernel +def _cfrc( + # Data in: + cinert_in: wp.array2d(dtype=vec10), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + cacc_in: wp.array2d(dtype=wp.spatial_vector), + cfrc_ext_in: wp.array2d(dtype=wp.spatial_vector), + # In: + flg_cfrc_ext: bool, + # Data out: + cfrc_int_out: wp.array2d(dtype=wp.spatial_vector), +): + worldid, bodyid = wp.tid() + bodyid += 1 # skip world body + cacc = cacc_in[worldid, bodyid] + cinert = cinert_in[worldid, bodyid] + cvel = cvel_in[worldid, bodyid] + frc = math.inert_vec(cinert, cacc) + frc += math.motion_cross_force(cvel, math.inert_vec(cinert, cvel)) + if flg_cfrc_ext: + frc -= cfrc_ext_in[worldid, bodyid] + + cfrc_int_out[worldid, bodyid] = frc + + +def _rne_cfrc(m: Model, d: Data, flg_cfrc_ext: bool = False): + wp.launch( + _cfrc, dim=[d.nworld, m.nbody - 1], inputs=[d.cinert, d.cvel, d.cacc, d.cfrc_ext, flg_cfrc_ext], outputs=[d.cfrc_int] + ) + + +@wp.kernel +def _cfrc_backward( + # Model: + body_parentid: wp.array(dtype=int), + # Data in: + cfrc_int_in: wp.array2d(dtype=wp.spatial_vector), + # In: + body_tree_: wp.array(dtype=int), + # Data out: + cfrc_int_out: wp.array2d(dtype=wp.spatial_vector), +): + worldid, nodeid = wp.tid() + bodyid = body_tree_[nodeid] + pid = body_parentid[bodyid] + if bodyid != 0: + wp.atomic_add(cfrc_int_out[worldid], pid, cfrc_int_in[worldid, bodyid]) + + +def _rne_cfrc_backward(m: Model, d: Data): + for body_tree in reversed(m.body_tree): + wp.launch( + _cfrc_backward, dim=[d.nworld, body_tree.size], inputs=[m.body_parentid, d.cfrc_int, body_tree], outputs=[d.cfrc_int] + ) + + +@wp.kernel +def _qfrc_bias( + # Model: + dof_bodyid: wp.array(dtype=int), + # Data in: + cdof_in: wp.array2d(dtype=wp.spatial_vector), + cfrc_int_in: wp.array2d(dtype=wp.spatial_vector), + # Data out: + qfrc_bias_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + bodyid = dof_bodyid[dofid] + qfrc_bias_out[worldid, dofid] = wp.dot(cdof_in[worldid, dofid], cfrc_int_in[worldid, bodyid]) + + +@event_scope +def rne(m: Model, d: Data, flg_acc: bool = False): + """ + Computes inverse dynamics using the recursive Newton-Euler algorithm. + + Computes the bias forces (qfrc_bias) and internal forces (cfrc_int) for the current state, + including the effects of gravity and optionally joint accelerations. + + Args: + m (Model): The model containing kinematic and dynamic information. + d (Data): The data object containing the current state and output arrays. + flg_acc (bool, optional): If True, includes joint accelerations in the computation. + Defaults to False. + """ + _rne_cacc_world(m, d) + _rne_cacc_forward(m, d, flg_acc=flg_acc) + _rne_cfrc(m, d) + _rne_cfrc_backward(m, d) + wp.launch(_qfrc_bias, dim=[d.nworld, m.nv], inputs=[m.dof_bodyid, d.cdof, d.cfrc_int], outputs=[d.qfrc_bias]) + + +@wp.kernel +def _cfrc_ext( + # Model: + body_rootid: wp.array(dtype=int), + # Data in: + xfrc_applied_in: wp.array2d(dtype=wp.spatial_vector), + xipos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + # Data out: + cfrc_ext_out: wp.array2d(dtype=wp.spatial_vector), +): + worldid, bodyid = wp.tid() + if bodyid == 0: + cfrc_ext_out[worldid, 0] = wp.spatial_vector(0.0, 0.0, 0.0, 0.0, 0.0, 0.0) + else: + xfrc_applied = xfrc_applied_in[worldid, bodyid] + subtree_com = subtree_com_in[worldid, body_rootid[bodyid]] + xipos = xipos_in[worldid, bodyid] + cfrc_ext_out[worldid, bodyid] = support.transform_force(xfrc_applied, subtree_com - xipos) + + +@wp.kernel +def _cfrc_ext_equality( + # Model: + body_rootid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + site_pos: wp.array2d(dtype=wp.vec3), + eq_obj1id: wp.array(dtype=int), + eq_obj2id: wp.array(dtype=int), + eq_objtype: wp.array(dtype=int), + eq_data: wp.array2d(dtype=vec11), + # Data in: + ne_connect_in: wp.array(dtype=int), + ne_weld_in: wp.array(dtype=int), + xpos_in: wp.array2d(dtype=wp.vec3), + xmat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + efc_id_in: wp.array2d(dtype=int), + efc_force_in: wp.array2d(dtype=float), + # Data out: + cfrc_ext_out: wp.array2d(dtype=wp.spatial_vector), +): + worldid, eqid = wp.tid() + + ne_connect = ne_connect_in[worldid] + ne_weld = ne_weld_in[worldid] + num_connect = ne_connect // 3 + + if eqid >= num_connect + ne_weld // 6: + return + + is_connect = eqid < num_connect + if is_connect: + efcid = 3 * eqid + cfrc_torque = wp.vec3(0.0, 0.0, 0.0) # no torque from connect + else: + efcid = 6 * eqid - ne_connect + cfrc_torque = wp.vec3(efc_force_in[worldid, efcid + 3], efc_force_in[worldid, efcid + 4], efc_force_in[worldid, efcid + 5]) + + cfrc_force = wp.vec3( + efc_force_in[worldid, efcid + 0], + efc_force_in[worldid, efcid + 1], + efc_force_in[worldid, efcid + 2], + ) + + id = efc_id_in[worldid, efcid] + eq_data_ = eq_data[worldid, id] + body_semantic = eq_objtype[id] == wp.static(ObjType.BODY.value) + + obj1 = eq_obj1id[id] + obj2 = eq_obj2id[id] + + if body_semantic: + bodyid1 = obj1 + bodyid2 = obj2 + else: + bodyid1 = site_bodyid[obj1] + bodyid2 = site_bodyid[obj2] + + # body 1 + if bodyid1: + if body_semantic: + if is_connect: + offset = wp.vec3(eq_data_[0], eq_data_[1], eq_data_[2]) + else: + offset = wp.vec3(eq_data_[3], eq_data_[4], eq_data_[5]) + else: + offset = site_pos[worldid, obj1] + + # transform point on body1: local -> global + pos = xmat_in[worldid, bodyid1] @ offset + xpos_in[worldid, bodyid1] + + # subtree CoM-based torque_force vector + newpos = subtree_com_in[worldid, body_rootid[bodyid1]] + + dif = newpos - pos + cfrc_com = wp.spatial_vector(cfrc_torque - wp.cross(dif, cfrc_force), cfrc_force) + + # apply (opposite for body 1) + wp.atomic_add(cfrc_ext_out[worldid], bodyid1, cfrc_com) + + # body 2 + if bodyid2: + if body_semantic: + if is_connect: + offset = wp.vec3(eq_data_[3], eq_data_[4], eq_data_[5]) + else: + offset = wp.vec3(eq_data_[0], eq_data_[1], eq_data_[2]) + else: + offset = site_pos[worldid, obj2] + + # transform point on body2: local -> global + pos = xmat_in[worldid, bodyid2] @ offset + xpos_in[worldid, bodyid2] + + # subtree CoM-based torque_force vector + newpos = subtree_com_in[worldid, body_rootid[bodyid2]] + + dif = newpos - pos + cfrc_com = wp.spatial_vector(cfrc_torque - wp.cross(dif, cfrc_force), cfrc_force) + + # apply + wp.atomic_sub(cfrc_ext_out[worldid], bodyid2, cfrc_com) + + +@wp.func +def transform_force(force: wp.vec3, torque: wp.vec3, offset: wp.vec3) -> wp.spatial_vector: + torque -= wp.cross(offset, force) + return wp.spatial_vector(torque, force) + + +@wp.kernel +def _cfrc_ext_contact( + # Model: + opt_cone: int, + body_rootid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + # Data in: + ncon_in: wp.array(dtype=int), + subtree_com_in: wp.array2d(dtype=wp.vec3), + contact_pos_in: wp.array(dtype=wp.vec3), + contact_frame_in: wp.array(dtype=wp.mat33), + contact_friction_in: wp.array(dtype=vec5), + contact_dim_in: wp.array(dtype=int), + contact_geom_in: wp.array(dtype=wp.vec2i), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_force_in: wp.array2d(dtype=float), + # Data out: + cfrc_ext_out: wp.array2d(dtype=wp.spatial_vector), +): + contactid = wp.tid() + + if contactid >= ncon_in[0]: + return + + geom = contact_geom_in[contactid] + id1 = geom_bodyid[geom[0]] + id2 = geom_bodyid[geom[1]] + + if id1 == 0 and id2 == 0: + return + + worldid = contact_worldid_in[contactid] + + # contact force in world frame + force = support.contact_force_fn( + opt_cone, + ncon_in, + contact_frame_in, + contact_friction_in, + contact_dim_in, + contact_efc_address_in, + efc_force_in, + worldid, + contactid, + to_world_frame=True, + ) + + pos = contact_pos_in[contactid] + + # contact force on bodies + if id1: + com1 = subtree_com_in[worldid, body_rootid[id1]] + wp.atomic_sub(cfrc_ext_out[worldid], id1, support.transform_force(force, com1 - pos)) + + if id2: + com2 = subtree_com_in[worldid, body_rootid[id2]] + wp.atomic_add(cfrc_ext_out[worldid], id2, support.transform_force(force, com2 - pos)) + + +@event_scope +def rne_postconstraint(m: Model, d: Data): + """ + Computes the recursive Newton-Euler algorithm after constraints are applied. + + Computes cacc, cfrc_ext, and cfrc_int, including the effects of applied forces, equality + constraints, and contacts. + """ + # cfrc_ext = perturb + wp.launch( + _cfrc_ext, + dim=(d.nworld, m.nbody), + inputs=[m.body_rootid, d.xfrc_applied, d.xipos, d.subtree_com], + outputs=[d.cfrc_ext], + ) + + wp.launch( + _cfrc_ext_equality, + dim=(d.nworld, m.neq), + inputs=[ + m.body_rootid, + m.site_bodyid, + m.site_pos, + m.eq_obj1id, + m.eq_obj2id, + m.eq_objtype, + m.eq_data, + d.ne_connect, + d.ne_weld, + d.xpos, + d.xmat, + d.subtree_com, + d.efc.id, + d.efc.force, + ], + outputs=[d.cfrc_ext], + ) + + # cfrc_ext += contacts + wp.launch( + _cfrc_ext_contact, + dim=(d.nconmax,), + inputs=[ + m.opt.cone, + m.body_rootid, + m.geom_bodyid, + d.ncon, + d.subtree_com, + d.contact.pos, + d.contact.frame, + d.contact.friction, + d.contact.dim, + d.contact.geom, + d.contact.efc_address, + d.contact.worldid, + d.efc.force, + ], + outputs=[d.cfrc_ext], + ) + + # forward pass over bodies: compute cacc, cfrc_int + _rne_cacc_world(m, d) + _rne_cacc_forward(m, d, flg_acc=True) + + # cfrc_body = cinert * cacc + cvel x (cinert * cvel) + _rne_cfrc(m, d, flg_cfrc_ext=True) + + # backward pass over bodies: accumulate cfrc_int from children + _rne_cfrc_backward(m, d) + + +@wp.kernel +def _tendon_dot( + # Model: + nv: int, + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + jnt_type: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + dof_bodyid: wp.array(dtype=int), + dof_jntid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + tendon_adr: wp.array(dtype=int), + tendon_num: wp.array(dtype=int), + tendon_armature: wp.array2d(dtype=float), + wrap_objid: wp.array(dtype=int), + wrap_prm: wp.array(dtype=float), + wrap_type: wp.array(dtype=int), + # Data in: + site_xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + cdof_dot_in: wp.array2d(dtype=wp.spatial_vector), + # Data out: + ten_Jdot_out: wp.array3d(dtype=float), +): + worldid, tenid = wp.tid() + + armature = tendon_armature[worldid, tenid] + if armature == 0.0: + return + + # fixed tendon has zero Jdot + adr = tendon_adr[tenid] + if wrap_type[adr] == int(WrapType.JOINT.value): + return + + # process spatial tendon + divisor = float(1.0) + num = tendon_num[tenid] + j = int(0) + while j < num - 1: + # get 1st and 2nd object + type0 = wrap_type[adr + j + 0] + type1 = wrap_type[adr + j + 1] + id0 = wrap_objid[adr + j + 0] + id1 = wrap_objid[adr + j + 1] + + # pulley + pulley = int(WrapType.PULLEY.value) + if (type0 == pulley) or (type1 == pulley): + # get divisor + if type0 == pulley: + divisor = wrap_prm[adr + j] + + j += 1 + continue + + # init sequence; assume it start with site + wpnt0 = site_xpos_in[worldid, id0] + + bodyid0 = site_bodyid[id0] + pos0 = site_xpos_in[worldid, id0] + cvel0 = cvel_in[worldid, bodyid0] + subtree_com0 = subtree_com_in[worldid, body_rootid[bodyid0]] + dif0 = pos0 - subtree_com0 + wvel0 = wp.spatial_bottom(cvel0) - wp.cross(dif0, wp.spatial_top(cvel0)) + wbody0 = site_bodyid[id0] + + # second object is geom: process site-geom-site + if (type1 == int(WrapType.SPHERE.value)) or (type1 == int(WrapType.CYLINDER.value)): + # TODO(team): derivatives of util_misc.wrap + return + + # complete sequence + wbody1 = site_bodyid[id1] + wpnt1 = site_xpos_in[worldid, id1] + + bodyid1 = site_bodyid[id1] + pos1 = site_xpos_in[worldid, id1] + cvel1 = cvel_in[worldid, bodyid1] + subtree_com1 = subtree_com_in[worldid, body_rootid[bodyid1]] + dif1 = pos1 - subtree_com1 + wvel1 = wp.spatial_bottom(cvel1) - wp.cross(dif1, wp.spatial_top(cvel1)) + + # accumulate moments if consecutive points are in different bodies + if wbody0 != wbody1: + # dpnt = 3D position difference, normalize + dpnt, norm = math.normalize_with_norm(wpnt1 - wpnt0) + + # dvel = d / dt (dpnt) + dvel = wvel1 - wvel0 + dot = wp.dot(dpnt, dvel) + dvel += dpnt * (-dot) + if norm > MJ_MINVAL: + dvel /= norm + else: + dvel = wp.vec3(0.0) + + # get endpoint Jacobian time derivatives, subtract + # TODO(team): parallelize? + for i in range(nv): + jac1, _ = support.jac_dot( + body_parentid, + body_rootid, + jnt_type, + jnt_dofadr, + dof_bodyid, + dof_jntid, + subtree_com_in, + cdof_in, + cvel_in, + cdof_dot_in, + wpnt0, + wbody0, + i, + worldid, + ) + jac2, _ = support.jac_dot( + body_parentid, + body_rootid, + jnt_type, + jnt_dofadr, + dof_bodyid, + dof_jntid, + subtree_com_in, + cdof_in, + cvel_in, + cdof_dot_in, + wpnt1, + wbody1, + i, + worldid, + ) + jacdif = jac2 - jac1 + + # chain rule, first term: Jdot += d / dt (jac2 - jac1) * dpnt + Jdot = wp.dot(jacdif, dpnt) + + # get endpoint Jacobians, subtract + jac1, _ = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + wpnt0, + wbody0, + i, + worldid, + ) + jac2, _ = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + wpnt1, + wbody1, + i, + worldid, + ) + jacdif = jac2 - jac1 + + # chain rule, second term: Jdot += (jac2 - jac1) * d / dt (dpnt) + Jdot += wp.dot(jacdif, dvel) + + ten_Jdot_out[worldid, tenid, i] += Jdot / divisor + + # TODO(team): j += 2 if geom wrapping + j += 1 + + +@wp.kernel +def _tendon_bias_coef( + # Model: + tendon_armature: wp.array2d(dtype=float), + # Data in: + qvel_in: wp.array2d(dtype=float), + ten_Jdot_in: wp.array3d(dtype=float), + # Data out: + ten_bias_coef_out: wp.array2d(dtype=float), +): + worldid, tenid, dofid = wp.tid() + + armature = tendon_armature[worldid, tenid] + if armature == 0.0: + return + + ten_Jdot = ten_Jdot_in[worldid, tenid, dofid] + if ten_Jdot == 0.0: + return + + wp.atomic_add(ten_bias_coef_out[worldid], tenid, ten_Jdot * qvel_in[worldid, dofid]) + + +@wp.kernel +def _tendon_bias_qfrc( + # Model: + tendon_armature: wp.array2d(dtype=float), + # Data in: + ten_J_in: wp.array3d(dtype=float), + ten_bias_coef_in: wp.array2d(dtype=float), + # Out: + qfrc_out: wp.array2d(dtype=float), +): + worldid, tenid, dofid = wp.tid() + + armature = tendon_armature[worldid, tenid] + if armature == 0.0: + return + + ten_J = ten_J_in[worldid, tenid, dofid] + if ten_J == 0.0: + return + + wp.atomic_add(qfrc_out[worldid], dofid, ten_J * armature * ten_bias_coef_in[worldid, tenid]) + + +@event_scope +def tendon_bias(m: Model, d: Data, qfrc: wp.array2d(dtype=float)): + """Add bias force due to tendon armature.""" + d.ten_Jdot.zero_() + wp.launch( + _tendon_dot, + dim=(d.nworld, m.ntendon), + inputs=[ + m.nv, + m.body_parentid, + m.body_rootid, + m.jnt_type, + m.jnt_dofadr, + m.dof_bodyid, + m.dof_jntid, + m.site_bodyid, + m.tendon_adr, + m.tendon_num, + m.tendon_armature, + m.wrap_objid, + m.wrap_prm, + m.wrap_type, + d.site_xpos, + d.subtree_com, + d.cdof, + d.cvel, + d.cdof_dot, + ], + outputs=[ + d.ten_Jdot, + ], + ) + + d.ten_bias_coef.zero_() + wp.launch( + _tendon_bias_coef, + dim=(d.nworld, m.ntendon, m.nv), + inputs=[ + m.tendon_armature, + d.qvel, + d.ten_Jdot, + ], + outputs=[ + d.ten_bias_coef, + ], + ) + + wp.launch( + _tendon_bias_qfrc, + dim=(d.nworld, m.ntendon, m.nv), + inputs=[ + m.tendon_armature, + d.ten_J, + d.ten_bias_coef, + ], + outputs=[ + qfrc, + ], + ) + + +@wp.kernel +def _comvel_root(cvel_out: wp.array2d(dtype=wp.spatial_vector)): + worldid, elementid = wp.tid() + cvel_out[worldid, 0][elementid] = 0.0 + + +@wp.kernel +def _comvel_level( + # Model: + body_parentid: wp.array(dtype=int), + body_jntnum: wp.array(dtype=int), + body_jntadr: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), + jnt_type: wp.array(dtype=int), + # Data in: + qvel_in: wp.array2d(dtype=float), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + # In: + body_tree_: wp.array(dtype=int), + # Data out: + cvel_out: wp.array2d(dtype=wp.spatial_vector), + cdof_dot_out: wp.array2d(dtype=wp.spatial_vector), +): + worldid, nodeid = wp.tid() + bodyid = body_tree_[nodeid] + dofid = body_dofadr[bodyid] + jntid = body_jntadr[bodyid] + jntnum = body_jntnum[bodyid] + pid = body_parentid[bodyid] + + if jntnum == 0: + cvel_out[worldid, bodyid] = cvel_in[worldid, pid] + return + + cvel = cvel_in[worldid, pid] + qvel = qvel_in[worldid] + cdof = cdof_in[worldid] + + for j in range(jntid, jntid + jntnum): + jnttype = jnt_type[j] + + if jnttype == wp.static(JointType.FREE.value): + cvel += cdof[dofid + 0] * qvel[dofid + 0] + cvel += cdof[dofid + 1] * qvel[dofid + 1] + cvel += cdof[dofid + 2] * qvel[dofid + 2] + + cdof_dot_out[worldid, dofid + 3] = math.motion_cross(cvel, cdof[dofid + 3]) + cdof_dot_out[worldid, dofid + 4] = math.motion_cross(cvel, cdof[dofid + 4]) + cdof_dot_out[worldid, dofid + 5] = math.motion_cross(cvel, cdof[dofid + 5]) + + cvel += cdof[dofid + 3] * qvel[dofid + 3] + cvel += cdof[dofid + 4] * qvel[dofid + 4] + cvel += cdof[dofid + 5] * qvel[dofid + 5] + + dofid += 6 + elif jnttype == wp.static(JointType.BALL.value): + cdof_dot_out[worldid, dofid + 0] = math.motion_cross(cvel, cdof[dofid + 0]) + cdof_dot_out[worldid, dofid + 1] = math.motion_cross(cvel, cdof[dofid + 1]) + cdof_dot_out[worldid, dofid + 2] = math.motion_cross(cvel, cdof[dofid + 2]) + + cvel += cdof[dofid + 0] * qvel[dofid + 0] + cvel += cdof[dofid + 1] * qvel[dofid + 1] + cvel += cdof[dofid + 2] * qvel[dofid + 2] + + dofid += 3 + else: + cdof_dot_out[worldid, dofid] = math.motion_cross(cvel, cdof[dofid]) + cvel += cdof[dofid] * qvel[dofid] + + dofid += 1 + + cvel_out[worldid, bodyid] = cvel + + +@event_scope +def com_vel(m: Model, d: Data): + """ + Computes the spatial velocities (cvel) and the derivative cdof_dot for all bodies. + + Propagates velocities down the kinematic tree, updating the spatial velocity and + derivative for each body. + """ + wp.launch(_comvel_root, dim=(d.nworld, 6), inputs=[], outputs=[d.cvel]) + + for body_tree in m.body_tree: + wp.launch( + _comvel_level, + dim=(d.nworld, body_tree.size), + inputs=[m.body_parentid, m.body_jntnum, m.body_jntadr, m.body_dofadr, m.jnt_type, d.qvel, d.cdof, d.cvel, body_tree], + outputs=[d.cvel, d.cdof_dot], + ) + + +@wp.kernel +def _transmission( + # Model: + nv: int, + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + body_weldid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), + jnt_type: wp.array(dtype=int), + jnt_qposadr: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + dof_bodyid: wp.array(dtype=int), + dof_parentid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + site_quat: wp.array2d(dtype=wp.quat), + actuator_trntype: wp.array(dtype=int), + actuator_trnid: wp.array(dtype=wp.vec2i), + actuator_gear: wp.array2d(dtype=wp.spatial_vector), + actuator_cranklength: wp.array(dtype=float), + tendon_adr: wp.array(dtype=int), + tendon_num: wp.array(dtype=int), + wrap_objid: wp.array(dtype=int), + wrap_type: wp.array(dtype=int), + # Data in: + qpos_in: wp.array2d(dtype=float), + xquat_in: wp.array2d(dtype=wp.quat), + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + ten_length_in: wp.array2d(dtype=float), + ten_J_in: wp.array3d(dtype=float), + # Data out: + actuator_length_out: wp.array2d(dtype=float), + actuator_moment_out: wp.array3d(dtype=float), +): + worldid, actid = wp.tid() + trntype = actuator_trntype[actid] + gear = actuator_gear[worldid, actid] + if trntype == wp.static(TrnType.JOINT.value) or trntype == wp.static(TrnType.JOINTINPARENT.value): + qpos = qpos_in[worldid] + jntid = actuator_trnid[actid][0] + jnt_typ = jnt_type[jntid] + qadr = jnt_qposadr[jntid] + vadr = jnt_dofadr[jntid] + if jnt_typ == wp.static(JointType.FREE.value): + actuator_length_out[worldid, actid] = 0.0 + if trntype == wp.static(TrnType.JOINTINPARENT.value): + quat = wp.normalize( + wp.quat( + qpos[qadr + 3], + qpos[qadr + 4], + qpos[qadr + 5], + qpos[qadr + 6], + ) + ) + quat_neg = math.quat_inv(quat) + gearaxis = math.rot_vec_quat(wp.spatial_bottom(gear), quat_neg) + actuator_moment_out[worldid, actid, vadr + 0] = gear[0] + actuator_moment_out[worldid, actid, vadr + 1] = gear[1] + actuator_moment_out[worldid, actid, vadr + 2] = gear[2] + actuator_moment_out[worldid, actid, vadr + 3] = gearaxis[0] + actuator_moment_out[worldid, actid, vadr + 4] = gearaxis[1] + actuator_moment_out[worldid, actid, vadr + 5] = gearaxis[2] + else: + for i in range(6): + actuator_moment_out[worldid, actid, vadr + i] = gear[i] + elif jnt_typ == wp.static(JointType.BALL.value): + q = wp.quat(qpos[qadr + 0], qpos[qadr + 1], qpos[qadr + 2], qpos[qadr + 3]) + q = wp.normalize(q) + axis_angle = math.quat_to_vel(q) + gearaxis = wp.spatial_top(gear) # [:3] + if trntype == wp.static(TrnType.JOINTINPARENT.value): + quat_neg = math.quat_inv(q) + gearaxis = math.rot_vec_quat(gearaxis, quat_neg) + actuator_length_out[worldid, actid] = wp.dot(axis_angle, gearaxis) + for i in range(3): + actuator_moment_out[worldid, actid, vadr + i] = gearaxis[i] + elif jnt_typ == wp.static(JointType.SLIDE.value) or jnt_typ == wp.static(JointType.HINGE.value): + actuator_length_out[worldid, actid] = qpos[qadr] * gear[0] + actuator_moment_out[worldid, actid, vadr] = gear[0] + else: + wp.printf("unrecognized joint type") + elif trntype == wp.static(TrnType.SLIDERCRANK.value): + # get data + trnid = actuator_trnid[actid] + id = trnid[0] + idslider = trnid[1] + gear0 = gear[0] + rod = actuator_cranklength[actid] + site_xmat = site_xmat_in[worldid, idslider] + axis = wp.vec3(site_xmat[0, 2], site_xmat[1, 2], site_xmat[2, 2]) + site_xpos_id = site_xpos_in[worldid, id] + site_xpos_idslider = site_xpos_in[worldid, idslider] + vec = site_xpos_id - site_xpos_idslider + + # compute length and determinant + # length = a' * v - sqrt(det); det = (a' * v)^2 + r^2 - v' * v + av = wp.dot(vec, axis) + det = av * av + rod * rod - wp.dot(vec, vec) + ok = 1 + if det <= 0.0: + ok = 0 + sdet = 0.0 + length = av + else: + sdet = wp.sqrt(det) + length = av - sdet + + actuator_length_out[worldid, actid] = length * gear0 + + # compute derivatives of length w.r.t. vec and axis + if ok == 1: + scale = 1.0 - av / sdet + dldv = axis * scale + vec / sdet + dlda = vec * scale + else: + dldv = axis + dlda = vec + + # apply chain rule + # TODO(team): parallelize? + for i in range(nv): + # get Jacobians of axis(jacA) and vec(jac) + # mj_jacPointAxis + jacp, jacr = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + site_xpos_idslider, + site_bodyid[idslider], + i, + worldid, + ) + jacS = jacp + jacA = wp.cross(jacr, axis) + + # mj_jacSite + jac, _ = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + site_xpos_id, + site_bodyid[id], + i, + worldid, + ) + jac -= jacS + + # apply the chain rule + moment = wp.dot(dlda, jacA) + wp.dot(dldv, jac) + actuator_moment_out[worldid, actid, i] = moment * gear0 + elif trntype == wp.static(TrnType.TENDON.value): + tenid = actuator_trnid[actid][0] + + gear0 = gear[0] + actuator_length_out[worldid, actid] = ten_length_in[worldid, tenid] * gear0 + + # fixed + adr = tendon_adr[tenid] + if wrap_type[adr] == wp.static(WrapType.JOINT.value): + ten_num = tendon_num[tenid] + for i in range(ten_num): + dofadr = jnt_dofadr[wrap_objid[adr + i]] + actuator_moment_out[worldid, actid, dofadr] = ten_J_in[worldid, tenid, dofadr] * gear0 + else: # spatial + for dofadr in range(nv): + actuator_moment_out[worldid, actid, dofadr] = ten_J_in[worldid, tenid, dofadr] * gear0 + elif trntype == wp.static(TrnType.BODY.value): + # cannot compute meaningful length, set to zero + actuator_length_out[worldid, actid] = 0.0 + + # initialize moment + for i in range(nv): + actuator_moment_out[worldid, actid, i] = 0.0 + + # moment computed by _transmission_body_moment and _transmission_body_moment_scale + elif trntype == int(TrnType.SITE.value): + trnid = actuator_trnid[actid] + siteid = trnid[0] + refid = trnid[1] + + gear = actuator_gear[worldid, actid] + gear_translation = wp.spatial_top(gear) + gear_rotational = wp.spatial_bottom(gear) + + # reference site undefined + if refid == -1: + # wrench: gear expressed in global frame + site_xmat = site_xmat_in[worldid, siteid] + wrench_translation = site_xmat @ gear_translation + wrench_rotation = site_xmat @ gear_rotational + + # moment: global Jacobian projected on wrench + # TODO(team): parallelize + for i in range(nv): + jacp, jacr = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + site_xpos_in[worldid, siteid], + site_bodyid[siteid], + i, + worldid, + ) + actuator_length_out[worldid, actid] = 0.0 + actuator_moment_out[worldid, actid, i] = wp.dot(jacp, wrench_translation) + wp.dot(jacr, wrench_rotation) + # reference site defined + else: + # initialize last dof address for each body + bodyid = site_bodyid[siteid] + bodyrefid = site_bodyid[refid] + b0 = body_weldid[bodyid] + b1 = body_weldid[bodyrefid] + dofadr0 = body_dofadr[b0] + body_dofnum[b0] - 1 + dofadr1 = body_dofadr[b1] + body_dofnum[b1] - 1 + + # find common ancestral dof, if any + dofadr_common = -1 + if dofadr0 >= 0 and dofadr1 >= 0: + # traverse up the tree until common ancestral dof is found + while dofadr0 != dofadr1: + if dofadr0 < dofadr1: + dofadr1 = dof_parentid[dofadr1] + else: + dofadr0 = dof_parentid[dofadr0] + + if dofadr0 == -1 or dofadr1 == -1: + # reached tree root, no common ancestral dof + break + + # found common ancestral dof + if dofadr0 == dofadr1: + dofadr_common = dofadr0 + + translational_transmission = not (gear[0] == 0.0 and gear[1] == 0.0 and gear[2] == 0.0) + rotational_transmission = not (gear[3] == 0.0 and gear[4] == 0.0 and gear[5] == 0.0) + + site_xpos = site_xpos_in[worldid, siteid] + ref_xpos = site_xpos_in[worldid, refid] + ref_xmat = site_xmat_in[worldid, refid] + + length = float(0.0) + + if translational_transmission: + # vec: site position in reference site frame + vec = wp.transpose(ref_xmat) @ (site_xpos - ref_xpos) + length += wp.dot(vec, gear_translation) + + wrench_translation = ref_xmat @ gear_translation + + if rotational_transmission: + # get site and refsite quats from parent bodies (avoid converting matrix to quat) + quat = math.mul_quat(site_quat[worldid, siteid], xquat_in[worldid, bodyid]) + refquat = math.mul_quat(site_quat[worldid, refid], xquat_in[worldid, bodyrefid]) + + # convert difference to expmap (axis-angle) + vec = math.quat_sub(quat, refquat) + length += wp.dot(vec, gear_rotational) + + wrench_rotation = ref_xmat @ gear_rotational + + actuator_length_out[worldid, actid] = length + + # TODO(team): parallelize + for i in range(nv): + jacp, jacr = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + site_xpos, + site_bodyid[siteid], + i, + worldid, + ) + + # jacref: global Jacobian of reference site + jacpref, jacrref = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + ref_xpos, + site_bodyid[refid], + i, + worldid, + ) + + jacpdif = jacp - jacpref + jacrdif = jacr - jacrref + + # if common ancestral dof was found, clear the columns of its parental chain + da = dofadr_common + while da >= 0: + if da == i: + jacpdif = wp.vec3(0.0) + jacrdif = wp.vec3(0.0) + break + da = dof_parentid[da] + + # moment: global Jacobian projected on wrench + moment = float(0.0) + + if translational_transmission: + moment += wp.dot(jacpdif, wrench_translation) + if rotational_transmission: + moment += wp.dot(jacrdif, wrench_rotation) + + actuator_moment_out[worldid, actid, i] = moment + else: + wp.printf("unhandled transmission type %d\n", trntype) + + +@wp.kernel +def _transmission_body_moment( + # Model: + opt_cone: int, + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + dof_bodyid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + actuator_trnid: wp.array(dtype=wp.vec2i), + actuator_trntype_body_adr: wp.array(dtype=int), + # Data in: + ncon_in: wp.array(dtype=int), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + contact_dist_in: wp.array(dtype=float), + contact_pos_in: wp.array(dtype=wp.vec3), + contact_frame_in: wp.array(dtype=wp.mat33), + contact_includemargin_in: wp.array(dtype=float), + contact_dim_in: wp.array(dtype=int), + contact_geom_in: wp.array(dtype=wp.vec2i), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_J_in: wp.array3d(dtype=float), + # Data out: + actuator_moment_out: wp.array3d(dtype=float), + actuator_trntype_body_ncon_out: wp.array2d(dtype=int), +): + trnbodyid, conid, dofid = wp.tid() + actid = actuator_trntype_body_adr[trnbodyid] + bodyid = actuator_trnid[actid][0] + + if conid >= ncon_in[0]: + return + + worldid = contact_worldid_in[conid] + + # get geom ids + geom = contact_geom_in[conid] + g1 = geom[0] + g2 = geom[1] + + # contact involving flex, continue + if g1 < 0 or g2 < 0: + return + + # get body ids + b1 = geom_bodyid[g1] + b2 = geom_bodyid[g2] + + # irrelevant contact, continue + if b1 != bodyid and b2 != bodyid: + return + + contact_exclude = int(contact_dist_in[conid] >= contact_includemargin_in[conid]) + + if dofid == 0: + wp.atomic_add(actuator_trntype_body_ncon_out[worldid], trnbodyid, 1) + + # mark contact normals in efc_force + if contact_exclude == 0: + contact_dim = contact_dim_in[conid] + contact_efc_address = contact_efc_address_in[conid] + + if contact_dim == 1 or opt_cone == int(ConeType.ELLIPTIC.value): + efc_force = 1.0 + efcid0 = contact_efc_address[0] + wp.atomic_add(actuator_moment_out[worldid, actid], dofid, efc_J_in[worldid, efcid0, dofid] * efc_force) + + else: + npyramid = contact_dim - 1 # number of frictional directions + efc_force = 0.5 / float(npyramid) + + for j in range(2 * npyramid): + efcid = contact_efc_address[j] + wp.atomic_add(actuator_moment_out[worldid, actid], dofid, efc_J_in[worldid, efcid, dofid] * efc_force) + + # excluded contact in gap: get Jacobian, accumulate + elif contact_exclude == 1: + contact_pos = contact_pos_in[conid] + contact_frame = contact_frame_in[conid] + normal = wp.vec3(contact_frame[0, 0], contact_frame[0, 1], contact_frame[0, 2]) + + # get Jacobian difference + jacp1, _ = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + contact_pos, + b1, + dofid, + worldid, + ) + jacp2, _ = support.jac( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + contact_pos, + b2, + dofid, + worldid, + ) + jacdif = jacp2 - jacp1 + + # project Jacobian along the normal of the contact frame + wp.atomic_add(actuator_moment_out[worldid, actid], dofid, wp.dot(normal, jacdif)) + + +@wp.kernel +def _transmission_body_moment_scale( + # Model: + actuator_trntype_body_adr: wp.array(dtype=int), + # Data in: + actuator_trntype_body_ncon_in: wp.array2d(dtype=int), + # Data out: + actuator_moment_out: wp.array3d(dtype=float), +): + worldid, trnbodyid, dofid = wp.tid() + + ncon = actuator_trntype_body_ncon_in[worldid, trnbodyid] + + if ncon > 0: + actid = actuator_trntype_body_adr[trnbodyid] + actuator_moment_out[worldid, actid, dofid] /= -float(ncon) + + +@event_scope +def transmission(m: Model, d: Data): + """ + Computes actuator/transmission lengths and moments. + + Updates the actuator length and moments for all actuators in the model, including joint + and tendon transmissions. + """ + wp.launch( + _transmission, + dim=[d.nworld, m.nu], + inputs=[ + m.nv, + m.body_parentid, + m.body_rootid, + m.body_weldid, + m.body_dofnum, + m.body_dofadr, + m.jnt_type, + m.jnt_qposadr, + m.jnt_dofadr, + m.dof_bodyid, + m.dof_parentid, + m.site_bodyid, + m.site_quat, + m.actuator_trntype, + m.actuator_trnid, + m.actuator_gear, + m.actuator_cranklength, + m.tendon_adr, + m.tendon_num, + m.wrap_objid, + m.wrap_type, + d.qpos, + d.xquat, + d.site_xpos, + d.site_xmat, + d.subtree_com, + d.cdof, + d.ten_length, + d.ten_J, + ], + outputs=[d.actuator_length, d.actuator_moment], + ) + + if m.actuator_trntype_body_adr.size > 0: + # reset number of active contacts + d.actuator_trntype_body_ncon.zero_() + + # compute moments + wp.launch( + _transmission_body_moment, + dim=( + m.actuator_trntype_body_adr.size, + d.nconmax, + m.nv, + ), + inputs=[ + m.opt.cone, + m.body_parentid, + m.body_rootid, + m.dof_bodyid, + m.geom_bodyid, + m.actuator_trnid, + m.actuator_trntype_body_adr, + d.ncon, + d.subtree_com, + d.cdof, + d.contact.dist, + d.contact.pos, + d.contact.frame, + d.contact.includemargin, + d.contact.dim, + d.contact.geom, + d.contact.efc_address, + d.contact.worldid, + d.efc.J, + ], + outputs=[ + d.actuator_moment, + d.actuator_trntype_body_ncon, + ], + ) + + # scale moments + wp.launch( + _transmission_body_moment_scale, + dim=(d.nworld, m.actuator_trntype_body_adr.size, m.nv), + inputs=[ + m.actuator_trntype_body_adr, + d.actuator_trntype_body_ncon, + ], + outputs=[d.actuator_moment], + ) + + +@wp.kernel +def _solve_LD_sparse_x_acc_up( + # In: + L: wp.array3d(dtype=float), + qLD_updates_: wp.array(dtype=wp.vec3i), + # Out: + x: wp.array2d(dtype=float), +): + worldid, nodeid = wp.tid() + update = qLD_updates_[nodeid] + i, k, Madr_ki = update[0], update[1], update[2] + wp.atomic_sub(x[worldid], i, L[worldid, 0, Madr_ki] * x[worldid, k]) + + +@wp.kernel +def _solve_LD_sparse_qLDiag_mul( + # In: + D: wp.array2d(dtype=float), + # Out: + out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + out[worldid, dofid] *= D[worldid, dofid] + + +@wp.kernel +def _solve_LD_sparse_x_acc_down( + # In: + L: wp.array3d(dtype=float), + qLD_updates_: wp.array(dtype=wp.vec3i), + # Out: + x: wp.array2d(dtype=float), +): + worldid, nodeid = wp.tid() + update = qLD_updates_[nodeid] + i, k, Madr_ki = update[0], update[1], update[2] + wp.atomic_sub(x[worldid], k, L[worldid, 0, Madr_ki] * x[worldid, i]) + + +def _solve_LD_sparse( + m: Model, + d: Data, + L: wp.array3d(dtype=float), + D: wp.array2d(dtype=float), + x: wp.array2d(dtype=float), + y: wp.array2d(dtype=float), +): + """Computes sparse backsubstitution: x = inv(L'*D*L)*y""" + + wp.copy(x, y) + for qLD_updates in reversed(m.qLD_updates): + wp.launch(_solve_LD_sparse_x_acc_up, dim=(d.nworld, qLD_updates.size), inputs=[L, qLD_updates], outputs=[x]) + + wp.launch(_solve_LD_sparse_qLDiag_mul, dim=(d.nworld, m.nv), inputs=[D], outputs=[x]) + + for qLD_updates in m.qLD_updates: + wp.launch(_solve_LD_sparse_x_acc_down, dim=(d.nworld, qLD_updates.size), inputs=[L, qLD_updates], outputs=[x]) + + +@cache_kernel +def _tile_cholesky_solve(tile: TileSet): + """Returns a kernel for dense Cholesky backsubstitution of a tile.""" + + @nested_kernel + def cholesky_solve( + # In: + L: wp.array3d(dtype=float), + y: wp.array2d(dtype=float), + adr: wp.array(dtype=int), + # Out: + x: wp.array2d(dtype=float), + ): + worldid, nodeid = wp.tid() + TILE_SIZE = wp.static(tile.size) + + dofid = adr[nodeid] + y_slice = wp.tile_load(y[worldid], shape=(TILE_SIZE,), offset=(dofid,)) + L_tile = wp.tile_load(L[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid)) + x_slice = wp.tile_cholesky_solve(L_tile, y_slice) + wp.tile_store(x[worldid], x_slice, offset=(dofid,)) + + return cholesky_solve + + +def _solve_LD_dense(m: Model, d: Data, L: wp.array3d(dtype=float), x: wp.array2d(dtype=float), y: wp.array2d(dtype=float)): + """Computes dense backsubstitution: x = inv(L'*L)*y""" + for tile in m.qM_tiles: + wp.launch_tiled( + _tile_cholesky_solve(tile), + dim=(d.nworld, tile.adr.size), + inputs=[L, y, tile.adr], + outputs=[x], + block_dim=m.block_dim.cholesky_solve, + ) + + +def solve_LD( + m: Model, + d: Data, + L: wp.array3d(dtype=float), + D: wp.array2d(dtype=float), + x: wp.array2d(dtype=float), + y: wp.array2d(dtype=float), +): + """ + Computes backsubstitution to solve a linear system of the form x = inv(L'*D*L) * y, + where L and D are the factors from the Cholesky factorization of the inertia matrix. + + This function dispatches to either a sparse or dense solver depending on Model options. + + Args: + m (Model): The model containing factorization and sparsity information. + d (Data): The data object containing workspace and factorization results. + L (array3d): Lower-triangular factor from the factorization (sparse or dense). + D (array2d): Diagonal factor from the factorization (only used for sparse). + x (array2d): Output array for the solution. + y (array2d): Input right-hand side array. + """ + if m.opt.is_sparse: + _solve_LD_sparse(m, d, L, D, x, y) + else: + _solve_LD_dense(m, d, L, x, y) + + +@event_scope +def solve_m(m: Model, d: Data, x: wp.array2d(dtype=float), y: wp.array2d(dtype=float)): + """ + Computes backsubstitution: x = qLD * y. + + Args: + m (Model): The model containing inertia and factorization information. + d (Data): The data object containing factorization results. + x (array2d): Output array for the solution. + y (array2d): Input right-hand side array. + """ + solve_LD(m, d, d.qLD, d.qLDiagInv, x, y) + + +@cache_kernel +def _tile_cholesky_factorize_solve(tile: TileSet): + """Returns a kernel for dense Cholesky factorization and backsubstitution of a tile.""" + + @nested_kernel + def cholesky_factorize_solve( + # In: + M: wp.array3d(dtype=float), + y: wp.array2d(dtype=float), + adr: wp.array(dtype=int), + # Out: + x: wp.array2d(dtype=float), + ): + worldid, nodeid = wp.tid() + TILE_SIZE = wp.static(tile.size) + + dofid = adr[nodeid] + M_tile = wp.tile_load(M[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid)) + y_slice = wp.tile_load(y[worldid], shape=(TILE_SIZE,), offset=(dofid,)) + + L_tile = wp.tile_cholesky(M_tile) + x_slice = wp.tile_cholesky_solve(L_tile, y_slice) + wp.tile_store(x[worldid], x_slice, offset=(dofid,)) + + return cholesky_factorize_solve + + +def _factor_solve_i_dense( + m: Model, d: Data, M: wp.array3d(dtype=float), x: wp.array2d(dtype=float), y: wp.array2d(dtype=float) +): + for tile in m.qM_tiles: + wp.launch_tiled( + _tile_cholesky_factorize_solve(tile), + dim=(d.nworld, tile.adr.size), + inputs=[M, y, tile.adr], + outputs=[x], + block_dim=m.block_dim.cholesky_factorize_solve, + ) + + +def factor_solve_i(m, d, M, L, D, x, y): + """ + Factorizes and solves the linear system: x = inv(L'*D*L) * y or x = inv(L'*L) * y, + where M is an inertia-like matrix and L, D are its Cholesky-like factors. + + This function first factorizes the matrix M (sparse or dense depending on model options), + then solves the system for x given right-hand side y. + + Args: + m (Model): The model containing factorization and sparsity information. + d (Data): The data object containing workspace and factorization results. + M (array3d): The inertia-like matrix to factorize. + L (array3d): Output lower-triangular factor from the factorization (sparse or dense). + D (array2d): Output diagonal factor from the factorization (only used for sparse). + x (array2d): Output array for the solution. + y (array2d): Input right-hand side array. + """ + if m.opt.is_sparse: + _factor_i_sparse(m, d, M, L, D) + _solve_LD_sparse(m, d, L, D, x, y) + else: + _factor_solve_i_dense(m, d, M, x, y) + + +@wp.kernel +def _subtree_vel_forward( + # Model: + body_rootid: wp.array(dtype=int), + body_mass: wp.array2d(dtype=float), + body_inertia: wp.array2d(dtype=wp.vec3), + # Data in: + xipos_in: wp.array2d(dtype=wp.vec3), + ximat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + # Data out: + subtree_linvel_out: wp.array2d(dtype=wp.vec3), + subtree_angmom_out: wp.array2d(dtype=wp.vec3), + subtree_bodyvel_out: wp.array2d(dtype=wp.spatial_vector), +): + worldid, bodyid = wp.tid() + + cvel = cvel_in[worldid, bodyid] + ang = wp.spatial_top(cvel) + lin = wp.spatial_bottom(cvel) + xipos = xipos_in[worldid, bodyid] + ximat = ximat_in[worldid, bodyid] + subtree_com_root = subtree_com_in[worldid, body_rootid[bodyid]] + + # update linear velocity + lin -= wp.cross(xipos - subtree_com_root, ang) + + subtree_linvel_out[worldid, bodyid] = body_mass[worldid, bodyid] * lin + dv = wp.transpose(ximat) @ ang + dv[0] *= body_inertia[worldid, bodyid][0] + dv[1] *= body_inertia[worldid, bodyid][1] + dv[2] *= body_inertia[worldid, bodyid][2] + subtree_angmom_out[worldid, bodyid] = ximat @ dv + subtree_bodyvel_out[worldid, bodyid] = wp.spatial_vector(ang, lin) + + +@wp.kernel +def _linear_momentum( + # Model: + body_parentid: wp.array(dtype=int), + body_subtreemass: wp.array2d(dtype=float), + # Data in: + subtree_linvel_in: wp.array2d(dtype=wp.vec3), + # In: + body_tree_: wp.array(dtype=int), + # Data out: + subtree_linvel_out: wp.array2d(dtype=wp.vec3), +): + worldid, nodeid = wp.tid() + bodyid = body_tree_[nodeid] + if bodyid: + pid = body_parentid[bodyid] + wp.atomic_add(subtree_linvel_out[worldid], pid, subtree_linvel_in[worldid, bodyid]) + subtree_linvel_out[worldid, bodyid] /= wp.max(MJ_MINVAL, body_subtreemass[worldid, bodyid]) + + +@wp.kernel +def _angular_momentum( + # Model: + body_parentid: wp.array(dtype=int), + body_mass: wp.array2d(dtype=float), + body_subtreemass: wp.array2d(dtype=float), + # Data in: + xipos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + subtree_linvel_in: wp.array2d(dtype=wp.vec3), + subtree_bodyvel_in: wp.array2d(dtype=wp.spatial_vector), + # In: + body_tree_: wp.array(dtype=int), + # Data out: + subtree_angmom_out: wp.array2d(dtype=wp.vec3), +): + worldid, nodeid = wp.tid() + bodyid = body_tree_[nodeid] + + if bodyid == 0: + return + + pid = body_parentid[bodyid] + + xipos = xipos_in[worldid, bodyid] + com = subtree_com_in[worldid, bodyid] + com_parent = subtree_com_in[worldid, pid] + vel = subtree_bodyvel_in[worldid, bodyid] + linvel = subtree_linvel_in[worldid, bodyid] + linvel_parent = subtree_linvel_in[worldid, pid] + mass = body_mass[worldid, bodyid] + subtreemass = body_subtreemass[worldid, bodyid] + + # momentum wrt body i + dx = xipos - com + dv = wp.spatial_bottom(vel) - linvel + dp = dv * mass + dL = wp.cross(dx, dp) + + # add to subtree i + subtree_angmom_out[worldid, bodyid] += dL + + # add to parent + wp.atomic_add(subtree_angmom_out[worldid], pid, subtree_angmom_out[worldid, bodyid]) + + # momentum wrt parent + dx = com - com_parent + dv = linvel - linvel_parent + dv *= subtreemass + dL = wp.cross(dx, dv) + wp.atomic_add(subtree_angmom_out[worldid], pid, dL) + + +def subtree_vel(m: Model, d: Data): + """ + Computes subtree linear velocity and angular momentum. + + Computes the linear momentum and angular momentum for each subtree, accumulating + contributions up the kinematic tree. + """ + + # bodywise quantities + wp.launch( + _subtree_vel_forward, + dim=(d.nworld, m.nbody), + inputs=[m.body_rootid, m.body_mass, m.body_inertia, d.xipos, d.ximat, d.subtree_com, d.cvel], + outputs=[d.subtree_linvel, d.subtree_angmom, d.subtree_bodyvel], + ) + + # sum body linear momentum recursively up the kinematic tree + for body_tree in reversed(m.body_tree): + wp.launch( + _linear_momentum, + dim=[d.nworld, body_tree.size], + inputs=[m.body_parentid, m.body_subtreemass, d.subtree_linvel, body_tree], + outputs=[d.subtree_linvel], + ) + + for body_tree in reversed(m.body_tree): + wp.launch( + _angular_momentum, + dim=[d.nworld, body_tree.size], + inputs=[ + m.body_parentid, + m.body_mass, + m.body_subtreemass, + d.xipos, + d.subtree_com, + d.subtree_linvel, + d.subtree_bodyvel, + body_tree, + ], + outputs=[d.subtree_angmom], + ) + + +@wp.kernel +def _joint_tendon( + # Model: + jnt_qposadr: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + wrap_objid: wp.array(dtype=int), + wrap_prm: wp.array(dtype=float), + tendon_jnt_adr: wp.array(dtype=int), + wrap_jnt_adr: wp.array(dtype=int), + # Data in: + qpos_in: wp.array2d(dtype=float), + # Data out: + ten_length_out: wp.array2d(dtype=float), + ten_J_out: wp.array3d(dtype=float), +): + worldid, wrapid = wp.tid() + + tendon_jnt_adr_ = tendon_jnt_adr[wrapid] + wrap_jnt_adr_ = wrap_jnt_adr[wrapid] + + wrap_objid_ = wrap_objid[wrap_jnt_adr_] + prm = wrap_prm[wrap_jnt_adr_] + + # add to length + L = prm * qpos_in[worldid, jnt_qposadr[wrap_objid_]] + # TODO(team): compare atomic_add and for loop + wp.atomic_add(ten_length_out[worldid], tendon_jnt_adr_, L) + + # add to moment + ten_J_out[worldid, tendon_jnt_adr_, jnt_dofadr[wrap_objid_]] = prm + + +@wp.kernel +def _spatial_site_tendon( + # Model: + nv: int, + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + dof_bodyid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + wrap_objid: wp.array(dtype=int), + tendon_site_pair_adr: wp.array(dtype=int), + wrap_site_pair_adr: wp.array(dtype=int), + wrap_pulley_scale: wp.array(dtype=float), + # Data in: + site_xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + # Data out: + ten_length_out: wp.array2d(dtype=float), + ten_J_out: wp.array3d(dtype=float), +): + worldid, elementid = wp.tid() + + # site pairs + site_pair_adr = wrap_site_pair_adr[elementid] + ten_adr = tendon_site_pair_adr[elementid] + + # pulley scaling + pulley_scale = wrap_pulley_scale[site_pair_adr] + + id0 = wrap_objid[site_pair_adr + 0] + id1 = wrap_objid[site_pair_adr + 1] + + pnt0 = site_xpos_in[worldid, id0] + pnt1 = site_xpos_in[worldid, id1] + dif = pnt1 - pnt0 + vec, length = math.normalize_with_norm(dif) + wp.atomic_add(ten_length_out[worldid], ten_adr, length * pulley_scale) + + if length < MJ_MINVAL: + vec = wp.vec3(1.0, 0.0, 0.0) + + body0 = site_bodyid[id0] + body1 = site_bodyid[id1] + if body0 != body1: + # TODO(team): parallelize + for i in range(nv): + jacp1, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pnt0, body0, i, worldid) + jacp2, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pnt1, body1, i, worldid) + + J = wp.dot(jacp2 - jacp1, vec) + if J: + wp.atomic_add(ten_J_out[worldid, ten_adr], i, J * pulley_scale) + + +@wp.kernel +def _spatial_geom_tendon( + # Model: + nv: int, + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + dof_bodyid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + site_bodyid: wp.array(dtype=int), + wrap_objid: wp.array(dtype=int), + wrap_prm: wp.array(dtype=float), + wrap_type: wp.array(dtype=int), + tendon_geom_adr: wp.array(dtype=int), + wrap_geom_adr: wp.array(dtype=int), + wrap_pulley_scale: wp.array(dtype=float), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + site_xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + # Data out: + ten_length_out: wp.array2d(dtype=float), + ten_J_out: wp.array3d(dtype=float), + wrap_geom_xpos_out: wp.array2d(dtype=wp.spatial_vector), +): + worldid, elementid = wp.tid() + wrap_adr = wrap_geom_adr[elementid] + ten_adr = tendon_geom_adr[elementid] + + # pulley scaling + pulley_scale = wrap_pulley_scale[wrap_adr] + + # site-geom-site + wrap_objid_site0 = wrap_objid[wrap_adr - 1] + wrap_objid_geom = wrap_objid[wrap_adr + 0] + wrap_objid_site1 = wrap_objid[wrap_adr + 1] + + # get site positions before and after geom + site_pnt0 = site_xpos_in[worldid, wrap_objid_site0] + site_pnt1 = site_xpos_in[worldid, wrap_objid_site1] + + # get geom information + geom_xpos = geom_xpos_in[worldid, wrap_objid_geom] + geom_xmat = geom_xmat_in[worldid, wrap_objid_geom] + geomsize = geom_size[worldid, wrap_objid_geom][0] + geom_type = wrap_type[wrap_adr] + + # get body ids for site-geom-site instances + bodyid_site0 = site_bodyid[wrap_objid_site0] + bodyid_geom = geom_bodyid[wrap_objid_geom] + bodyid_site1 = site_bodyid[wrap_objid_site1] + + # find wrap object sidesite (if it exists) + sideid = int(wp.round(wrap_prm[wrap_adr])) + if sideid >= 0: + side = site_xpos_in[worldid, sideid] + else: + side = wp.vec3(wp.inf) + + # compute geom wrap length and connect points (if wrap occurs) + length_geomgeom, geom_pnt0, geom_pnt1 = util_misc.wrap(site_pnt0, site_pnt1, geom_xpos, geom_xmat, geomsize, geom_type, side) + + # store geom points + wrap_geom_xpos_out[worldid, elementid] = wp.spatial_vector(geom_pnt0, geom_pnt1) + + if length_geomgeom >= 0.0: + dif_sitegeom = geom_pnt0 - site_pnt0 + dif_geomsite = site_pnt1 - geom_pnt1 + vec_sitegeom, length_sitegeom = math.normalize_with_norm(dif_sitegeom) + vec_geomsite, length_geomsite = math.normalize_with_norm(dif_geomsite) + + # length + length_sitegeomsite = length_sitegeom + length_geomgeom + length_geomsite + + if length_sitegeomsite: + wp.atomic_add(ten_length_out[worldid], ten_adr, length_sitegeomsite * pulley_scale) + + # moment + if length_sitegeom < MJ_MINVAL: + vec_sitegeom = wp.vec3(1.0, 0.0, 0.0) + + if length_geomsite < MJ_MINVAL: + vec_geomsite = wp.vec3(1.0, 0.0, 0.0) + + dif_body_sitegeom = bodyid_site0 != bodyid_geom + dif_body_geomsite = bodyid_geom != bodyid_site1 + + # TODO(team): parallelize + for i in range(nv): + J = float(0.0) + # site-geom + if dif_body_sitegeom: + jacp_site0, _ = support.jac( + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt0, bodyid_site0, i, worldid + ) + + jacp_geom0, _ = support.jac( + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, geom_pnt0, bodyid_geom, i, worldid + ) + + J += wp.dot(jacp_geom0 - jacp_site0, vec_sitegeom) + + # geom-site + if dif_body_geomsite: + jacp_geom1, _ = support.jac( + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, geom_pnt1, bodyid_geom, i, worldid + ) + + jacp_site1, _ = support.jac( + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt1, bodyid_site1, i, worldid + ) + + J += wp.dot(jacp_site1 - jacp_geom1, vec_geomsite) + + if J: + wp.atomic_add(ten_J_out[worldid, ten_adr], i, J * pulley_scale) + else: + dif_sitesite = site_pnt1 - site_pnt0 + vec_sitesite, length_sitesite = math.normalize_with_norm(dif_sitesite) + + # length + if length_sitesite: + wp.atomic_add(ten_length_out[worldid], ten_adr, length_sitesite * pulley_scale) + + # moment + if length_sitesite < MJ_MINVAL: + vec_sitesite = wp.vec3(1.0, 0.0, 0.0) + + if bodyid_site0 != bodyid_site1: + # TODO(team): parallelize + for i in range(nv): + jacp1, _ = support.jac( + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt0, bodyid_site0, i, worldid + ) + jacp2, _ = support.jac( + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt1, bodyid_site1, i, worldid + ) + + J = wp.dot(jacp2 - jacp1, vec_sitesite) + + if J: + wp.atomic_add(ten_J_out[worldid, ten_adr], i, J * pulley_scale) + + +@wp.kernel +def _spatial_tendon_wrap( + # Model: + ntendon: int, + tendon_adr: wp.array(dtype=int), + tendon_num: wp.array(dtype=int), + wrap_objid: wp.array(dtype=int), + wrap_type: wp.array(dtype=int), + # Data in: + site_xpos_in: wp.array2d(dtype=wp.vec3), + wrap_geom_xpos_in: wp.array2d(dtype=wp.spatial_vector), + # Data out: + ten_wrapadr_out: wp.array2d(dtype=int), + ten_wrapnum_out: wp.array2d(dtype=int), + wrap_obj_out: wp.array2d(dtype=wp.vec2i), + wrap_xpos_out: wp.array2d(dtype=wp.spatial_vector), +): + worldid = wp.tid() + + wrapcount = int(0) + wrapgeomid = int(0) + for i in range(ntendon): + adr = tendon_adr[i] + ten_wrapadr_out[worldid, i] = wrapcount + wrapnum = int(0) + tendonnum = tendon_num[i] + + # process fixed tendon + if wrap_type[adr] == int(WrapType.JOINT.value): + continue + + # process spatial tendon + j = int(0) + while j < tendonnum - 1: + # get 1st and 2nd object + type0 = wrap_type[adr + j + 0] + type1 = wrap_type[adr + j + 1] + id0 = wrap_objid[adr + j + 0] + id1 = wrap_objid[adr + j + 1] + + # pulley + pulley0 = type0 == int(WrapType.PULLEY.value) + if pulley0 or type1 == int(WrapType.PULLEY.value): + if pulley0: + row = wrapcount // 2 + col = wrapcount % 2 + wrap_xpos_out[worldid, row][3 * col + 0] = 0.0 + wrap_xpos_out[worldid, row][3 * col + 1] = 0.0 + wrap_xpos_out[worldid, row][3 * col + 2] = 0.0 + + wrap_obj_out[worldid, row][col] = -2 + + wrapnum += 1 + wrapcount += 1 + + # move to next + j += 1 + continue + + # init sequence; assume it starts with site + wpnt_site0 = site_xpos_in[worldid, id0] + + # second object is geom: process site-geom-site + if type1 == int(WrapType.SPHERE.value) or type1 == int(WrapType.CYLINDER.value): + wrap_geom_xpos = wrap_geom_xpos_in[worldid, wrapgeomid] + wpnt_geom0 = wp.spatial_top(wrap_geom_xpos) + wrapgeomid += 1 + + wrapid = id1 + id1 = wrap_objid[adr + j + 2] + if wp.norm_l2(wpnt_geom0) < wp.inf: + wpnt_geom1 = wp.spatial_bottom(wrap_geom_xpos) + wpnt_site1 = site_xpos_in[worldid, id1] + + # assign to wrap + row0 = (wrapcount + 0) // 2 + col0 = (wrapcount + 0) % 2 + row1 = (wrapcount + 1) // 2 + col1 = (wrapcount + 1) % 2 + row2 = (wrapcount + 2) // 2 + col2 = (wrapcount + 2) % 2 + row3 = (wrapcount + 3) // 2 + col3 = (wrapcount + 3) % 2 + + wrap_xpos_out[worldid, row0][3 * col0 + 0] = wpnt_site0[0] + wrap_xpos_out[worldid, row0][3 * col0 + 1] = wpnt_site0[1] + wrap_xpos_out[worldid, row0][3 * col0 + 2] = wpnt_site0[2] + + wrap_xpos_out[worldid, row1][3 * col1 + 0] = wpnt_geom0[0] + wrap_xpos_out[worldid, row1][3 * col1 + 1] = wpnt_geom0[1] + wrap_xpos_out[worldid, row1][3 * col1 + 2] = wpnt_geom0[2] + + wrap_xpos_out[worldid, row2][3 * col2 + 0] = wpnt_geom1[0] + wrap_xpos_out[worldid, row2][3 * col2 + 1] = wpnt_geom1[1] + wrap_xpos_out[worldid, row2][3 * col2 + 2] = wpnt_geom1[2] + + wrap_xpos_out[worldid, row3][3 * col3 + 0] = wpnt_site1[0] + wrap_xpos_out[worldid, row3][3 * col3 + 1] = wpnt_site1[1] + wrap_xpos_out[worldid, row3][3 * col3 + 2] = wpnt_site1[2] + + wrap_obj_out[worldid, row0][col0] = -1 + wrap_obj_out[worldid, row1][col1] = wrapid + wrap_obj_out[worldid, row2][col2] = wrapid + + wrapnum += 3 + wrapcount += 3 + j += 2 + + else: + row0 = (wrapcount + 0) // 2 + col0 = (wrapcount + 0) % 2 + + wrap_xpos_out[worldid, row0][3 * col0 + 0] = wpnt_site0[0] + wrap_xpos_out[worldid, row0][3 * col0 + 1] = wpnt_site0[1] + wrap_xpos_out[worldid, row0][3 * col0 + 2] = wpnt_site0[2] + + wrap_obj_out[worldid, row0][col0] = -1 + + wrapnum += 1 + wrapcount += 1 + j += 2 + + else: + row0 = (wrapcount + 0) // 2 + col0 = (wrapcount + 0) % 2 + + wrap_xpos_out[worldid, row0][3 * col0 + 0] = wpnt_site0[0] + wrap_xpos_out[worldid, row0][3 * col0 + 1] = wpnt_site0[1] + wrap_xpos_out[worldid, row0][3 * col0 + 2] = wpnt_site0[2] + + wrap_obj_out[worldid, row0][col0] = -1 + + wrapnum += 1 + wrapcount += 1 + j += 1 + + # assign last site before pulley or tendon end + if adr + j + 1 < wrap_type.shape[0]: + last_before_pulley = wrap_type[adr + j + 1] == int(WrapType.PULLEY.value) + else: + last_before_pulley = False + + if j == tendonnum - 1 or last_before_pulley: + row0 = (wrapcount + 0) // 2 + col0 = (wrapcount + 0) % 2 + + wpnt_site1 = site_xpos_in[worldid, id1] + wrap_xpos_out[worldid, row0][3 * col0 + 0] = wpnt_site1[0] + wrap_xpos_out[worldid, row0][3 * col0 + 1] = wpnt_site1[1] + wrap_xpos_out[worldid, row0][3 * col0 + 2] = wpnt_site1[2] + + wrap_obj_out[worldid, row0][col0] = -1 + wrapnum += 1 + wrapcount += 1 + + ten_wrapnum_out[worldid, i] = wrapnum + + +def tendon(m: Model, d: Data): + """ + Computes tendon lengths and moments. + + Updates the tendon length and moment arrays for all tendons in the model, including joint, + site, and geom tendons. + """ + if not m.ntendon: + return + + d.ten_length.zero_() + d.ten_J.zero_() + + # process joint tendons + wp.launch( + _joint_tendon, + dim=(d.nworld, m.wrap_jnt_adr.size), + inputs=[m.jnt_qposadr, m.jnt_dofadr, m.wrap_objid, m.wrap_prm, m.tendon_jnt_adr, m.wrap_jnt_adr, d.qpos], + outputs=[d.ten_length, d.ten_J], + ) + + spatial_site = m.wrap_site_pair_adr.size > 0 + spatial_geom = m.wrap_geom_adr.size > 0 + + if spatial_site or spatial_geom: + d.wrap_xpos.zero_() + d.wrap_obj.zero_() + + # process spatial site tendons + wp.launch( + _spatial_site_tendon, + dim=(d.nworld, m.wrap_site_pair_adr.size), + inputs=[ + m.nv, + m.body_parentid, + m.body_rootid, + m.dof_bodyid, + m.site_bodyid, + m.wrap_objid, + m.tendon_site_pair_adr, + m.wrap_site_pair_adr, + m.wrap_pulley_scale, + d.site_xpos, + d.subtree_com, + d.cdof, + ], + outputs=[d.ten_length, d.ten_J], + ) + + # process spatial geom tendons + wp.launch( + _spatial_geom_tendon, + dim=(d.nworld, m.wrap_geom_adr.size), + inputs=[ + m.nv, + m.body_parentid, + m.body_rootid, + m.dof_bodyid, + m.geom_bodyid, + m.geom_size, + m.site_bodyid, + m.wrap_objid, + m.wrap_prm, + m.wrap_type, + m.tendon_geom_adr, + m.wrap_geom_adr, + m.wrap_pulley_scale, + d.geom_xpos, + d.geom_xmat, + d.site_xpos, + d.subtree_com, + d.cdof, + ], + outputs=[d.ten_length, d.ten_J, d.wrap_geom_xpos], + ) + + if spatial_site or spatial_geom: + wp.launch( + _spatial_tendon_wrap, + dim=(d.nworld,), + inputs=[m.ntendon, m.tendon_adr, m.tendon_num, m.wrap_objid, m.wrap_type, d.site_xpos, d.wrap_geom_xpos], + outputs=[d.ten_wrapadr, d.ten_wrapnum, d.wrap_obj, d.wrap_xpos], + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth_test.py new file mode 100644 index 00000000..9ceed777 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth_test.py @@ -0,0 +1,398 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Tests for smooth dynamics functions.""" + +import mujoco +import numpy as np +import warp as wp +from absl.testing import absltest +from absl.testing import parameterized + +import mujoco_warp as mjwarp + +from mujoco.mjx.third_party.mujoco_warp._src import test_util +from mujoco.mjx.third_party.mujoco_warp._src import types + +# tolerance for difference between MuJoCo and MJWarp smooth calculations - mostly +# due to float precision +_TOLERANCE = 5e-5 + + +def _assert_eq(a, b, name): + tol = _TOLERANCE * 10 # avoid test noise + err_msg = f"mismatch: {name}" + np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol) + + +class SmoothTest(parameterized.TestCase): + def test_kinematics(self): + """Tests kinematics.""" + _, mjd, m, d = test_util.fixture("pendula.xml") + + for arr in ( + d.xanchor, + d.xaxis, + d.xpos, + d.xquat, + d.xmat, + d.xipos, + d.ximat, + d.geom_xpos, + d.geom_xmat, + d.site_xpos, + d.site_xmat, + ): + arr.zero_() + + mjwarp.kinematics(m, d) + + _assert_eq(d.xanchor.numpy()[0], mjd.xanchor, "xanchor") + _assert_eq(d.xaxis.numpy()[0], mjd.xaxis, "xaxis") + _assert_eq(d.xpos.numpy()[0], mjd.xpos, "xpos") + _assert_eq(d.xquat.numpy()[0], mjd.xquat, "xquat") + _assert_eq(d.xmat.numpy()[0], mjd.xmat.reshape((-1, 3, 3)), "xmat") + _assert_eq(d.xipos.numpy()[0], mjd.xipos, "xipos") + _assert_eq(d.ximat.numpy()[0], mjd.ximat.reshape((-1, 3, 3)), "ximat") + _assert_eq(d.geom_xpos.numpy()[0], mjd.geom_xpos, "geom_xpos") + _assert_eq(d.geom_xmat.numpy()[0], mjd.geom_xmat.reshape((-1, 3, 3)), "geom_xmat") + _assert_eq(d.site_xpos.numpy()[0], mjd.site_xpos, "site_xpos") + _assert_eq(d.site_xmat.numpy()[0], mjd.site_xmat.reshape((-1, 3, 3)), "site_xmat") + + def test_com_pos(self): + """Tests com_pos.""" + _, mjd, m, d = test_util.fixture("pendula.xml") + + for arr in (d.subtree_com, d.cinert, d.cdof): + arr.zero_() + + mjwarp.com_pos(m, d) + _assert_eq(d.subtree_com.numpy()[0], mjd.subtree_com, "subtree_com") + _assert_eq(d.cinert.numpy()[0], mjd.cinert, "cinert") + _assert_eq(d.cdof.numpy()[0], mjd.cdof, "cdof") + + def test_camlight(self): + """Tests camlight.""" + _, mjd, m, d = test_util.fixture("pendula.xml") + + d.cam_xpos.zero_() + d.cam_xmat.zero_() + d.light_xpos.zero_() + d.light_xdir.zero_() + + mjwarp.camlight(m, d) + _assert_eq(d.cam_xpos.numpy()[0], mjd.cam_xpos, "cam_xpos") + _assert_eq(d.cam_xmat.numpy()[0], mjd.cam_xmat.reshape((-1, 3, 3)), "cam_xmat") + _assert_eq(d.light_xpos.numpy()[0], mjd.light_xpos, "light_xpos") + _assert_eq(d.light_xdir.numpy()[0], mjd.light_xdir, "light_xdir") + + @parameterized.parameters(True, False) + def test_crb(self, sparse: bool): + """Tests crb.""" + mjm, mjd, m, d = test_util.fixture("pendula.xml", sparse=sparse) + + d.crb.zero_() + + mjwarp.crb(m, d) + _assert_eq(d.crb.numpy()[0], mjd.crb, "crb") + + if sparse: + _assert_eq(d.qM.numpy()[0, 0], mjd.qM, "qM") + else: + qM = np.zeros((mjm.nv, mjm.nv)) + mujoco.mj_fullM(mjm, qM, mjd.qM) + _assert_eq(d.qM.numpy()[0], qM, "qM") + + @parameterized.parameters(True, False) + def test_factor_m(self, sparse: bool): + """Tests factor_m.""" + _, mjd, m, d = test_util.fixture("pendula.xml", sparse=sparse) + + qLD = d.qLD.numpy()[0].copy() + for arr in (d.qLD, d.qLDiagInv): + arr.zero_() + + mjwarp.factor_m(m, d) + + if sparse: + _assert_eq(d.qLD.numpy()[0, 0], mjd.qLD, "qLD (sparse)") + _assert_eq(d.qLDiagInv.numpy()[0], mjd.qLDiagInv, "qLDiagInv") + else: + _assert_eq(d.qLD.numpy()[0], qLD, "qLD (dense)") + + @parameterized.parameters(True, False) + def test_solve_m(self, sparse: bool): + """Tests solve_m.""" + mjm, mjd, m, d = test_util.fixture("pendula.xml", sparse=sparse) + + qfrc_smooth = np.tile(mjd.qfrc_smooth, (1, 1)) + qacc_smooth = np.zeros( + shape=( + 1, + mjm.nv, + ), + dtype=float, + ) + mujoco.mj_solveM(mjm, mjd, qacc_smooth, qfrc_smooth) + + d.qacc_smooth.zero_() + + mjwarp.solve_m(m, d, d.qacc_smooth, d.qfrc_smooth) + _assert_eq(d.qacc_smooth.numpy()[0], qacc_smooth[0], "qacc_smooth") + + @parameterized.parameters(True, False) + def test_rne(self, gravity): + """Tests rne.""" + _, mjd, m, d = test_util.fixture("pendula.xml", gravity=gravity) + + d.qfrc_bias.zero_() + + mjwarp.rne(m, d) + _assert_eq(d.qfrc_bias.numpy()[0], mjd.qfrc_bias, "qfrc_bias") + + @parameterized.parameters(True, False) + def test_rne_postconstraint(self, gravity): + """Tests rne_postconstraint.""" + mjm, mjd, m, d = test_util.fixture("pendula.xml", gravity=gravity) + + mjd.xfrc_applied = np.random.uniform(low=-0.01, high=0.01, size=mjd.xfrc_applied.shape) + d.xfrc_applied = wp.array(np.expand_dims(mjd.xfrc_applied, axis=0), dtype=wp.spatial_vector) + + mujoco.mj_rnePostConstraint(mjm, mjd) + + for arr in (d.cacc, d.cfrc_int, d.cfrc_ext): + arr.zero_() + + mjwarp.rne_postconstraint(m, d) + + _assert_eq(d.cacc.numpy()[0], mjd.cacc, "cacc") + _assert_eq(d.cfrc_int.numpy()[0], mjd.cfrc_int, "cfrc_int") + _assert_eq(d.cfrc_ext.numpy()[0], mjd.cfrc_ext, "cfrc_ext") + + _EQUALITY = """ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + """ + mjm, mjd, m, d = test_util.fixture(xml=_EQUALITY, kick=True, keyframe=0) + + mujoco.mj_rnePostConstraint(mjm, mjd) + + d.cfrc_ext.zero_() + mjwarp.rne_postconstraint(m, d) + + _assert_eq(d.cfrc_ext.numpy()[0], mjd.cfrc_ext, "cfrc_ext (equality)") + + mjm, mjd, m, d = test_util.fixture("constraints.xml", keyframe=1, equality=False) + + mujoco.mj_rnePostConstraint(mjm, mjd) + + d.cfrc_ext.zero_() + + # clear equality constraint counts + d.ne_connect.zero_() + d.ne_weld.zero_() + d.ne_jnt.zero_() + + mjwarp.rne_postconstraint(m, d) + + _assert_eq(d.cfrc_ext.numpy()[0], mjd.cfrc_ext, "cfrc_ext (contact)") + + def test_com_vel(self): + """Tests com_vel.""" + _, mjd, m, d = test_util.fixture("pendula.xml") + + for arr in (d.cvel, d.cdof_dot): + arr.zero_() + + mjwarp.com_vel(m, d) + _assert_eq(d.cvel.numpy()[0], mjd.cvel, "cvel") + _assert_eq(d.cdof_dot.numpy()[0], mjd.cdof_dot, "cdof_dot") + + @parameterized.parameters("pendula.xml", "actuation/site.xml", "actuation/slidercrank.xml") + def test_transmission(self, xml): + """Tests transmission.""" + mjm, mjd, m, d = test_util.fixture(xml) + + for arr in (d.actuator_length, d.actuator_moment): + arr.zero_() + + actuator_moment = np.zeros((mjm.nu, mjm.nv)) + mujoco.mju_sparse2dense( + actuator_moment, + mjd.actuator_moment, + mjd.moment_rownnz, + mjd.moment_rowadr, + mjd.moment_colind, + ) + + mjwarp._src.smooth.transmission(m, d) + _assert_eq(d.actuator_length.numpy()[0], mjd.actuator_length, "actuator_length") + _assert_eq(d.actuator_moment.numpy()[0], actuator_moment, "actuator_moment") + + @parameterized.product(keyframe=list(range(4)), cone=list(types.ConeType)) + def test_actuator_adhesion(self, keyframe, cone): + """Tests adhesion actuator.""" + mjm, mjd, m, d = test_util.fixture("actuation/adhesion.xml", keyframe=keyframe, cone=cone) + + d.actuator_length.zero_() + d.actuator_moment.zero_() + mjwarp._src.collision_driver.collision(m, d) # compute contact.includemargin + mjwarp._src.constraint.make_constraint(m, d) # compute contact.efc_address + mjwarp._src.smooth.transmission(m, d) + + actuator_moment = np.zeros((mjm.nu, mjm.nv)) + mujoco.mju_sparse2dense(actuator_moment, mjd.actuator_moment, mjd.moment_rownnz, mjd.moment_rowadr, mjd.moment_colind) + + _assert_eq(d.actuator_length.numpy()[0], mjd.actuator_length, "actuator_length") + _assert_eq(d.actuator_moment.numpy()[0], actuator_moment, "acutator_moment") + + def test_subtree_vel(self): + """Tests subtree_vel.""" + mjm, mjd, m, d = test_util.fixture("pendula.xml") + + for arr in (d.subtree_linvel, d.subtree_angmom): + arr.zero_() + + mujoco.mj_subtreeVel(mjm, mjd) + mjwarp.subtree_vel(m, d) + + _assert_eq(d.subtree_linvel.numpy()[0], mjd.subtree_linvel, "subtree_linvel") + _assert_eq(d.subtree_angmom.numpy()[0], mjd.subtree_angmom, "subtree_angmom") + + @parameterized.parameters( + "tendon/fixed.xml", + "tendon/site.xml", + "tendon/pulley_site.xml", + "tendon/fixed_site.xml", + "tendon/pulley_fixed_site.xml", + "tendon/site_fixed.xml", + "tendon/pulley_site_fixed.xml", + "tendon/wrap.xml", + "tendon/pulley_wrap.xml", + ) + def test_tendon(self, xml): + """Tests tendon.""" + mjm, mjd, m, d = test_util.fixture(xml, keyframe=0) + + for arr in (d.ten_length, d.ten_J, d.actuator_length, d.actuator_moment): + arr.zero_() + + mjwarp.tendon(m, d) + mjwarp.transmission(m, d) + + _assert_eq(d.ten_length.numpy()[0], mjd.ten_length, "ten_length") + _assert_eq(d.ten_J.numpy()[0], mjd.ten_J.reshape((mjm.ntendon, mjm.nv)), "ten_J") + _assert_eq(d.wrap_xpos.numpy()[0], mjd.wrap_xpos, "wrap_xpos") + _assert_eq(d.wrap_obj.numpy()[0], mjd.wrap_obj, "wrap_obj") + _assert_eq(d.ten_wrapnum.numpy()[0], mjd.ten_wrapnum, "ten_wrapnum") + _assert_eq(d.ten_wrapadr.numpy()[0], mjd.ten_wrapadr, "ten_wrapadr") + _assert_eq(d.actuator_length.numpy()[0], mjd.actuator_length, "actuator_length") + actuator_moment = np.zeros((mjm.nu, mjm.nv)) + mujoco.mju_sparse2dense( + actuator_moment, + mjd.actuator_moment, + mjd.moment_rownnz, + mjd.moment_rowadr, + mjd.moment_colind, + ) + _assert_eq(d.actuator_moment.numpy()[0], actuator_moment, "actuator_moment") + + @parameterized.parameters(True, False) + def test_factor_solve_i(self, sparse): + mjm, mjd, m, d = test_util.fixture( + xml=""" + + + + + + + + + """, + sparse=sparse, + ) + + qM = np.zeros((mjm.nv, mjm.nv)) + mujoco.mj_fullM(mjm, qM, mjd.qM) + + d.qLD.zero_() + if sparse: + d.qLDiagInv.zero_() + + res = wp.zeros((1, mjm.nv), dtype=float) + vec = wp.ones((1, mjm.nv), dtype=float) + + mjwarp._src.smooth.factor_solve_i(m, d, d.qM, d.qLD, d.qLDiagInv, res, vec) + + _assert_eq(res.numpy()[0], np.linalg.solve(qM, vec.numpy()[0]), "qM \\ 1") + + def test_tendon_armature(self): + mjm, mjd, m, d = test_util.fixture("tendon/armature.xml", keyframe=0) + + # qM + d.qM.zero_() + + mjwarp._src.smooth.crb(m, d) + mjwarp._src.smooth.tendon_armature(m, d) + + qM = np.zeros((mjm.nv, mjm.nv)) + mujoco.mj_fullM(mjm, qM, mjd.qM) + _assert_eq(d.qM.numpy()[0], qM, "qM") + + # qfrc_bias + d.qfrc_bias.zero_() + + mjwarp._src.smooth.rne(m, d) + mjwarp._src.smooth.tendon_bias(m, d, d.qfrc_bias) + _assert_eq(d.qfrc_bias.numpy()[0], mjd.qfrc_bias, "qfrc_bias") + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py new file mode 100644 index 00000000..31de24df --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py @@ -0,0 +1,2578 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +from math import ceil + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src import math +from mujoco.mjx.third_party.mujoco_warp._src import smooth +from mujoco.mjx.third_party.mujoco_warp._src import support +from mujoco.mjx.third_party.mujoco_warp._src import types +from mujoco.mjx.third_party.mujoco_warp._src.block_cholesky import create_blocked_cholesky_func +from mujoco.mjx.third_party.mujoco_warp._src.block_cholesky import create_blocked_cholesky_solve_func +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel + + +@wp.func +def _rescale(nv: int, stat_meaninertia: float, value: float) -> float: + return value / (stat_meaninertia * float(wp.max(1, nv))) + + +@wp.func +def _in_bracket(x: wp.vec3, y: wp.vec3) -> bool: + return (x[1] < y[1] and y[1] < 0.0) or (x[1] > y[1] and y[1] > 0.0) + + +@wp.func +def _eval_pt(quad: wp.vec3, alpha: float) -> wp.vec3: + return wp.vec3( + alpha * alpha * quad[2] + alpha * quad[1] + quad[0], + 2.0 * alpha * quad[2] + quad[1], + 2.0 * quad[2], + ) + + +@wp.func +def _eval_pt_elliptic( + # In: + impratio: float, + friction: types.vec5, + u0: float, + uu: float, + uv: float, + vv: float, + jv: float, + D: float, + quad: wp.vec3, + alpha: float, +) -> wp.vec3: + mu = friction[0] / wp.sqrt(impratio) + v0 = jv * mu + n = u0 + alpha * v0 + tsqr = uu + alpha * (2.0 * uv + alpha * vv) + t = wp.sqrt(tsqr) # tangential force + + bottom_zone = ((tsqr <= 0.0) and (n < 0)) or ((tsqr > 0.0) and ((mu * n + t) <= 0.0)) + middle_zone = (tsqr > 0) and (n < (mu * t)) and ((mu * n + t) > 0.0) + + # elliptic bottom zone: quadratic cose + if bottom_zone: + pt = _eval_pt(quad, alpha) + else: + pt = wp.vec3(0.0) + + # elliptic middle zone + if t == 0.0: + t += types.MJ_MINVAL + + if tsqr == 0.0: + tsqr += types.MJ_MINVAL + + n1 = v0 + t1 = (uv + alpha * vv) / t + t2 = vv / t - (uv + alpha * vv) * t1 / tsqr + + if middle_zone: + mu2 = mu * mu + dm = D / wp.max(mu2 * (1.0 + mu2), types.MJ_MINVAL) + nmt = n - mu * t + n1mut1 = n1 - mu * t1 + + pt += wp.vec3( + 0.5 * dm * nmt * nmt, + dm * nmt * n1mut1, + dm * (n1mut1 * n1mut1 - nmt * mu * t2), + ) + + return pt + + +@wp.kernel +def linesearch_iterative_init_gtol_p0_gauss( + # Model: + nv: int, + opt_tolerance: wp.array(dtype=float), + opt_ls_tolerance: wp.array(dtype=float), + stat_meaninertia: float, + # Data in: + efc_search_dot_in: wp.array(dtype=float), + efc_quad_gauss_in: wp.array(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_gtol_out: wp.array(dtype=float), + efc_p0_out: wp.array(dtype=wp.vec3), +): + worldid = wp.tid() + + if efc_done_in[worldid]: + return + + tolerance = opt_tolerance[worldid] + ls_tolerance = opt_ls_tolerance[worldid] + snorm = wp.math.sqrt(efc_search_dot_in[worldid]) + scale = stat_meaninertia * wp.float(wp.max(1, nv)) + efc_gtol_out[worldid] = tolerance * ls_tolerance * snorm * scale + + quad = efc_quad_gauss_in[worldid] + efc_p0_out[worldid] = wp.vec3(quad[0], quad[1], 2.0 * quad[2]) + + +@wp.kernel +def linesearch_iterative_init_p0_elliptic0( + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nl_in: wp.array(dtype=int), + nefc_in: wp.array(dtype=int), + efc_Jaref_in: wp.array2d(dtype=float), + efc_quad_in: wp.array2d(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + efc_condim_in: wp.array2d(dtype=int), + # Data out: + efc_p0_out: wp.array(dtype=wp.vec3), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + if efc_done_in[worldid]: + return + + active = efc_Jaref_in[worldid, efcid] < 0.0 + + nef = ne_in[worldid] + nf_in[worldid] + nefl = nef + nl_in[worldid] + if efcid < nef: + active = True + elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1: + active = False + + if active: + quad = efc_quad_in[worldid, efcid] + wp.atomic_add(efc_p0_out, worldid, wp.vec3(quad[0], quad[1], 2.0 * quad[2])) + + +@wp.kernel +def linesearch_iterative_init_p0_elliptic1( + # Model: + opt_impratio: wp.array(dtype=float), + # Data in: + ncon_in: wp.array(dtype=int), + contact_friction_in: wp.array(dtype=types.vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_D_in: wp.array2d(dtype=float), + efc_jv_in: wp.array2d(dtype=float), + efc_quad_in: wp.array2d(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + efc_u_in: wp.array(dtype=types.vec6), + efc_uu_in: wp.array(dtype=float), + efc_uv_in: wp.array(dtype=float), + efc_vv_in: wp.array(dtype=float), + # Data out: + efc_p0_out: wp.array(dtype=wp.vec3), +): + conid = wp.tid() + + if conid >= ncon_in[0]: + return + + worldid = contact_worldid_in[conid] + if efc_done_in[worldid]: + return + + if contact_dim_in[conid] < 2: + return + + efcid = contact_efc_address_in[conid, 0] + + pt = _eval_pt_elliptic( + opt_impratio[worldid], + contact_friction_in[conid], + efc_u_in[conid][0], + efc_uu_in[conid], + efc_uv_in[conid], + efc_vv_in[conid], + efc_jv_in[worldid, efcid], + efc_D_in[worldid, efcid], + efc_quad_in[worldid, efcid], + 0.0, + ) + + wp.atomic_add(efc_p0_out, worldid, pt) + + +@wp.kernel +def linesearch_iterative_init_p0_pyramidal( + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nefc_in: wp.array(dtype=int), + efc_Jaref_in: wp.array2d(dtype=float), + efc_quad_in: wp.array2d(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_p0_out: wp.array(dtype=wp.vec3), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + if efc_done_in[worldid]: + return + + if efc_Jaref_in[worldid, efcid] >= 0.0 and efcid >= ne_in[worldid] + nf_in[worldid]: + return + + quad = efc_quad_in[worldid, efcid] + + wp.atomic_add(efc_p0_out, worldid, wp.vec3(quad[0], quad[1], 2.0 * quad[2])) + + +@wp.kernel +def linesearch_iterative_init_lo_gauss( + # Data in: + efc_quad_gauss_in: wp.array(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + efc_p0_in: wp.array(dtype=wp.vec3), + # Data out: + efc_lo_out: wp.array(dtype=wp.vec3), + efc_lo_alpha_out: wp.array(dtype=float), +): + worldid = wp.tid() + + if efc_done_in[worldid]: + return + + p0 = efc_p0_in[worldid] + alpha = -math.safe_div(p0[1], p0[2]) + efc_lo_out[worldid] = _eval_pt(efc_quad_gauss_in[worldid], alpha) + efc_lo_alpha_out[worldid] = alpha + + +@wp.kernel +def linesearch_iterative_init_lo_elliptic0( + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nl_in: wp.array(dtype=int), + nefc_in: wp.array(dtype=int), + efc_Jaref_in: wp.array2d(dtype=float), + efc_jv_in: wp.array2d(dtype=float), + efc_quad_in: wp.array2d(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + efc_lo_alpha_in: wp.array(dtype=float), + efc_condim_in: wp.array2d(dtype=int), + # Data out: + efc_lo_out: wp.array(dtype=wp.vec3), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + if efc_done_in[worldid]: + return + + alpha = efc_lo_alpha_in[worldid] + + active = efc_Jaref_in[worldid, efcid] + alpha * efc_jv_in[worldid, efcid] < 0.0 + + nef = ne_in[worldid] + nf_in[worldid] + nefl = nef + nl_in[worldid] + if efcid < nef: + active = True + elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1: + active = False + + if active: + wp.atomic_add(efc_lo_out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha)) + + +@wp.kernel +def linesearch_iterative_init_lo_elliptic1( + # Model: + opt_impratio: wp.array(dtype=float), + # Data in: + ncon_in: wp.array(dtype=int), + contact_friction_in: wp.array(dtype=types.vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_D_in: wp.array2d(dtype=float), + efc_jv_in: wp.array2d(dtype=float), + efc_quad_in: wp.array2d(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + efc_lo_alpha_in: wp.array(dtype=float), + efc_u_in: wp.array(dtype=types.vec6), + efc_uu_in: wp.array(dtype=float), + efc_uv_in: wp.array(dtype=float), + efc_vv_in: wp.array(dtype=float), + # Data out: + efc_lo_out: wp.array(dtype=wp.vec3), +): + conid = wp.tid() + + if conid >= ncon_in[0]: + return + + worldid = contact_worldid_in[conid] + if efc_done_in[worldid]: + return + + if contact_dim_in[conid] < 2: + return + + efcid = contact_efc_address_in[conid, 0] + alpha = efc_lo_alpha_in[worldid] + pt = _eval_pt_elliptic( + opt_impratio[worldid], + contact_friction_in[conid], + efc_u_in[conid][0], + efc_uu_in[conid], + efc_uv_in[conid], + efc_vv_in[conid], + efc_jv_in[worldid, efcid], + efc_D_in[worldid, efcid], + efc_quad_in[worldid, efcid], + alpha, + ) + wp.atomic_add(efc_lo_out, worldid, pt) + + +@wp.kernel +def linesearch_iterative_init_lo_pyramidal( + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nefc_in: wp.array(dtype=int), + efc_Jaref_in: wp.array2d(dtype=float), + efc_jv_in: wp.array2d(dtype=float), + efc_quad_in: wp.array2d(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + efc_lo_alpha_in: wp.array(dtype=float), + # Data out: + efc_lo_out: wp.array(dtype=wp.vec3), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + if efc_done_in[worldid]: + return + + alpha = efc_lo_alpha_in[worldid] + + if efc_Jaref_in[worldid, efcid] + alpha * efc_jv_in[worldid, efcid] < 0.0 or (efcid < ne_in[worldid] + nf_in[worldid]): + wp.atomic_add(efc_lo_out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha)) + + +@wp.kernel +def linesearch_iterative_init_bounds( + # Data in: + efc_done_in: wp.array(dtype=bool), + efc_p0_in: wp.array(dtype=wp.vec3), + efc_lo_in: wp.array(dtype=wp.vec3), + efc_lo_alpha_in: wp.array(dtype=float), + # Data out: + efc_lo_out: wp.array(dtype=wp.vec3), + efc_lo_alpha_out: wp.array(dtype=float), + efc_hi_out: wp.array(dtype=wp.vec3), + efc_hi_alpha_out: wp.array(dtype=float), +): + worldid = wp.tid() + + if efc_done_in[worldid]: + return + + p0 = efc_p0_in[worldid] + lo = efc_lo_in[worldid] + lo_alpha = efc_lo_alpha_in[worldid] + lo_less = lo[1] < p0[1] + + efc_lo_out[worldid] = wp.where(lo_less, lo, p0) + efc_lo_alpha_out[worldid] = wp.where(lo_less, lo_alpha, 0.0) + efc_hi_out[worldid] = wp.where(lo_less, p0, lo) + efc_hi_alpha_out[worldid] = wp.where(lo_less, 0.0, lo_alpha) + + +@wp.kernel +def linesearch_iterative_next_alpha_gauss( + # Data in: + efc_quad_gauss_in: wp.array(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + efc_ls_done_in: wp.array(dtype=bool), + efc_lo_in: wp.array(dtype=wp.vec3), + efc_lo_alpha_in: wp.array(dtype=float), + efc_hi_in: wp.array(dtype=wp.vec3), + efc_hi_alpha_in: wp.array(dtype=float), + # Data out: + efc_lo_next_out: wp.array(dtype=wp.vec3), + efc_lo_next_alpha_out: wp.array(dtype=float), + efc_hi_next_out: wp.array(dtype=wp.vec3), + efc_hi_next_alpha_out: wp.array(dtype=float), + efc_mid_out: wp.array(dtype=wp.vec3), + efc_mid_alpha_out: wp.array(dtype=float), +): + worldid = wp.tid() + + if efc_ls_done_in[worldid]: + return + + if efc_done_in[worldid]: + return + + quad = efc_quad_gauss_in[worldid] + + lo = efc_lo_in[worldid] + lo_alpha = efc_lo_alpha_in[worldid] + lo_next_alpha = lo_alpha - math.safe_div(lo[1], lo[2]) + efc_lo_next_out[worldid] = _eval_pt(quad, lo_next_alpha) + efc_lo_next_alpha_out[worldid] = lo_next_alpha + + hi = efc_hi_in[worldid] + hi_alpha = efc_hi_alpha_in[worldid] + hi_next_alpha = hi_alpha - math.safe_div(hi[1], hi[2]) + efc_hi_next_out[worldid] = _eval_pt(quad, hi_next_alpha) + efc_hi_next_alpha_out[worldid] = hi_next_alpha + + mid_alpha = 0.5 * (lo_alpha + hi_alpha) + efc_mid_out[worldid] = _eval_pt(quad, mid_alpha) + efc_mid_alpha_out[worldid] = mid_alpha + + +@wp.kernel +def linesearch_iterative_next_quad_elliptic0( + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nl_in: wp.array(dtype=int), + nefc_in: wp.array(dtype=int), + efc_Jaref_in: wp.array2d(dtype=float), + efc_jv_in: wp.array2d(dtype=float), + efc_quad_in: wp.array2d(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + efc_ls_done_in: wp.array(dtype=bool), + efc_lo_next_alpha_in: wp.array(dtype=float), + efc_hi_next_alpha_in: wp.array(dtype=float), + efc_mid_alpha_in: wp.array(dtype=float), + efc_condim_in: wp.array2d(dtype=int), + # Data out: + efc_lo_next_out: wp.array(dtype=wp.vec3), + efc_hi_next_out: wp.array(dtype=wp.vec3), + efc_mid_out: wp.array(dtype=wp.vec3), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + if efc_done_in[worldid]: + return + + if efc_ls_done_in[worldid]: + return + + nef = ne_in[worldid] + nf_in[worldid] + nefl = nef + nl_in[worldid] + + quad = efc_quad_in[worldid, efcid] + jaref = efc_Jaref_in[worldid, efcid] + jv = efc_jv_in[worldid, efcid] + + alpha = efc_lo_next_alpha_in[worldid] + + active = jaref + alpha * jv < 0.0 + if efcid < nef: + active = True + elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1: + active = False + + if active: + wp.atomic_add(efc_lo_next_out, worldid, _eval_pt(quad, alpha)) + + alpha = efc_hi_next_alpha_in[worldid] + + active = jaref + alpha * jv < 0.0 + if efcid < nef: + active = True + elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1: + active = False + + if active: + wp.atomic_add(efc_hi_next_out, worldid, _eval_pt(quad, alpha)) + + alpha = efc_mid_alpha_in[worldid] + + active = jaref + alpha * jv < 0.0 + if efcid < nef: + active = True + elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1: + active = False + + if active: + wp.atomic_add(efc_mid_out, worldid, _eval_pt(quad, alpha)) + + +@wp.kernel +def linesearch_iterative_next_quad_elliptic1( + # Model: + opt_impratio: wp.array(dtype=float), + # Data in: + ncon_in: wp.array(dtype=int), + contact_friction_in: wp.array(dtype=types.vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_D_in: wp.array2d(dtype=float), + efc_jv_in: wp.array2d(dtype=float), + efc_quad_in: wp.array2d(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + efc_lo_next_alpha_in: wp.array(dtype=float), + efc_hi_next_alpha_in: wp.array(dtype=float), + efc_mid_alpha_in: wp.array(dtype=float), + efc_u_in: wp.array(dtype=types.vec6), + efc_uu_in: wp.array(dtype=float), + efc_uv_in: wp.array(dtype=float), + efc_vv_in: wp.array(dtype=float), + # Data out: + efc_lo_next_out: wp.array(dtype=wp.vec3), + efc_hi_next_out: wp.array(dtype=wp.vec3), + efc_mid_out: wp.array(dtype=wp.vec3), +): + conid = wp.tid() + + if conid >= ncon_in[0]: + return + + worldid = contact_worldid_in[conid] + + if efc_done_in[worldid]: + return + + if contact_dim_in[conid] < 2: + return + + efcid = contact_efc_address_in[conid, 0] + impratio = opt_impratio[worldid] + friction = contact_friction_in[conid] + u = efc_u_in[conid][0] + uu = efc_uu_in[conid] + uv = efc_uv_in[conid] + vv = efc_vv_in[conid] + jv = efc_jv_in[worldid, efcid] + d = efc_D_in[worldid, efcid] + quad = efc_quad_in[worldid, efcid] + + alpha = efc_lo_next_alpha_in[worldid] + pt = _eval_pt_elliptic(impratio, friction, u, uu, uv, vv, jv, d, quad, alpha) + wp.atomic_add(efc_lo_next_out, worldid, pt) + + alpha = efc_hi_next_alpha_in[worldid] + pt = _eval_pt_elliptic(impratio, friction, u, uu, uv, vv, jv, d, quad, alpha) + wp.atomic_add(efc_hi_next_out, worldid, pt) + + alpha = efc_mid_alpha_in[worldid] + pt = _eval_pt_elliptic(impratio, friction, u, uu, uv, vv, jv, d, quad, alpha) + wp.atomic_add(efc_mid_out, worldid, pt) + + +@wp.kernel +def linesearch_iterative_next_quad_pyramidal( + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nefc_in: wp.array(dtype=int), + efc_Jaref_in: wp.array2d(dtype=float), + efc_jv_in: wp.array2d(dtype=float), + efc_quad_in: wp.array2d(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + efc_ls_done_in: wp.array(dtype=bool), + efc_lo_next_alpha_in: wp.array(dtype=float), + efc_hi_next_alpha_in: wp.array(dtype=float), + efc_mid_alpha_in: wp.array(dtype=float), + # Data out: + efc_lo_next_out: wp.array(dtype=wp.vec3), + efc_hi_next_out: wp.array(dtype=wp.vec3), + efc_mid_out: wp.array(dtype=wp.vec3), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + if efc_done_in[worldid]: + return + + if efc_ls_done_in[worldid]: + return + + nef_active = efcid < ne_in[worldid] + nf_in[worldid] + + quad = efc_quad_in[worldid, efcid] + jaref = efc_Jaref_in[worldid, efcid] + jv = efc_jv_in[worldid, efcid] + + alpha = efc_lo_next_alpha_in[worldid] + if jaref + alpha * jv < 0.0 or nef_active: + wp.atomic_add(efc_lo_next_out, worldid, _eval_pt(quad, alpha)) + + alpha = efc_hi_next_alpha_in[worldid] + if jaref + alpha * jv < 0.0 or nef_active: + wp.atomic_add(efc_hi_next_out, worldid, _eval_pt(quad, alpha)) + + alpha = efc_mid_alpha_in[worldid] + if jaref + alpha * jv < 0.0 or nef_active: + wp.atomic_add(efc_mid_out, worldid, _eval_pt(quad, alpha)) + + +@wp.kernel +def linesearch_iterative_swap( + # Data in: + efc_gtol_in: wp.array(dtype=float), + efc_done_in: wp.array(dtype=bool), + efc_ls_done_in: wp.array(dtype=bool), + efc_p0_in: wp.array(dtype=wp.vec3), + efc_lo_in: wp.array(dtype=wp.vec3), + efc_lo_alpha_in: wp.array(dtype=float), + efc_hi_in: wp.array(dtype=wp.vec3), + efc_hi_alpha_in: wp.array(dtype=float), + efc_lo_next_in: wp.array(dtype=wp.vec3), + efc_lo_next_alpha_in: wp.array(dtype=float), + efc_hi_next_in: wp.array(dtype=wp.vec3), + efc_hi_next_alpha_in: wp.array(dtype=float), + efc_mid_in: wp.array(dtype=wp.vec3), + efc_mid_alpha_in: wp.array(dtype=float), + # Data out: + efc_alpha_out: wp.array(dtype=float), + efc_ls_done_out: wp.array(dtype=bool), + efc_lo_out: wp.array(dtype=wp.vec3), + efc_lo_alpha_out: wp.array(dtype=float), + efc_hi_out: wp.array(dtype=wp.vec3), + efc_hi_alpha_out: wp.array(dtype=float), +): + worldid = wp.tid() + + if efc_done_in[worldid]: + return + + if efc_ls_done_in[worldid]: + return + + lo = efc_lo_in[worldid] + lo_alpha = efc_lo_alpha_in[worldid] + hi = efc_hi_in[worldid] + hi_alpha = efc_hi_alpha_in[worldid] + lo_next = efc_lo_next_in[worldid] + lo_next_alpha = efc_lo_next_alpha_in[worldid] + hi_next = efc_hi_next_in[worldid] + hi_next_alpha = efc_hi_next_alpha_in[worldid] + mid = efc_mid_in[worldid] + mid_alpha = efc_mid_alpha_in[worldid] + + # swap lo: + swap_lo_lo_next = _in_bracket(lo, lo_next) + lo = wp.where(swap_lo_lo_next, lo_next, lo) + lo_alpha = wp.where(swap_lo_lo_next, lo_next_alpha, lo_alpha) + swap_lo_mid = _in_bracket(lo, mid) + lo = wp.where(swap_lo_mid, mid, lo) + lo_alpha = wp.where(swap_lo_mid, mid_alpha, lo_alpha) + swap_lo_hi_next = _in_bracket(lo, hi_next) + lo = wp.where(swap_lo_hi_next, hi_next, lo) + lo_alpha = wp.where(swap_lo_hi_next, hi_next_alpha, lo_alpha) + efc_lo_out[worldid] = lo + efc_lo_alpha_out[worldid] = lo_alpha + swap_lo = swap_lo_lo_next or swap_lo_mid or swap_lo_hi_next + + # swap hi: + swap_hi_hi_next = _in_bracket(hi, hi_next) + hi = wp.where(swap_hi_hi_next, hi_next, hi) + hi_alpha = wp.where(swap_hi_hi_next, hi_next_alpha, hi_alpha) + swap_hi_mid = _in_bracket(hi, mid) + hi = wp.where(swap_hi_mid, mid, hi) + hi_alpha = wp.where(swap_hi_mid, mid_alpha, hi_alpha) + swap_hi_lo_next = _in_bracket(hi, lo_next) + hi = wp.where(swap_hi_lo_next, lo_next, hi) + hi_alpha = wp.where(swap_hi_lo_next, lo_next_alpha, hi_alpha) + efc_hi_out[worldid] = hi + efc_hi_alpha_out[worldid] = hi_alpha + swap_hi = swap_hi_hi_next or swap_hi_mid or swap_hi_lo_next + + # if we did not adjust the interval, we are done + # also done if either low or hi slope is nearly flat + gtol = efc_gtol_in[worldid] + efc_ls_done_out[worldid] = (not swap_lo and not swap_hi) or (lo[1] < 0 and lo[1] > -gtol) or (hi[1] > 0 and hi[1] < gtol) + + # update alpha if we have an improvement + p0 = efc_p0_in[worldid] + alpha = 0.0 + improved = lo[0] < p0[0] or hi[0] < p0[0] + lo_better = lo[0] < hi[0] + alpha = wp.where(improved and lo_better, lo_alpha, alpha) + alpha = wp.where(improved and not lo_better, hi_alpha, alpha) + efc_alpha_out[worldid] = alpha + + +def _linesearch_iterative(m: types.Model, d: types.Data): + """Iterative linesearch.""" + d.efc.ls_done.zero_() + + wp.launch( + linesearch_iterative_init_gtol_p0_gauss, + dim=(d.nworld,), + inputs=[ + m.nv, m.opt.tolerance, m.opt.ls_tolerance, m.stat.meaninertia, d.efc.search_dot, + d.efc.quad_gauss, d.efc.done + ], + outputs=[d.efc.gtol, d.efc.p0]) # fmt: skip + + if m.opt.cone == types.ConeType.ELLIPTIC: + wp.launch( + linesearch_iterative_init_p0_elliptic0, + dim=(d.nworld, d.njmax,), + inputs=[ + d.ne, d.nf, d.nl, d.nefc, d.efc.Jaref, d.efc.quad, d.efc.done, + d.efc.condim + ], + outputs=[d.efc.p0]) # fmt: skip + wp.launch( + linesearch_iterative_init_p0_elliptic1, + dim=(d.nconmax), + inputs=[ + m.opt.impratio, d.ncon, d.contact.friction, d.contact.dim, + d.contact.efc_address, d.contact.worldid, d.efc.D, d.efc.jv, d.efc.quad, + d.efc.done, d.efc.u, d.efc.uu, d.efc.uv, d.efc.vv + ], + outputs=[d.efc.p0]) # fmt: skip + else: + wp.launch( + linesearch_iterative_init_p0_pyramidal, + dim=(d.nworld, d.njmax,), + inputs=[ + d.ne, d.nf, d.nefc, d.efc.Jaref, d.efc.quad, d.efc.done + ], outputs=[d.efc.p0]) # fmt: skip + + wp.launch( + linesearch_iterative_init_lo_gauss, + dim=(d.nworld,), + inputs=[ + d.efc.quad_gauss, d.efc.done, d.efc.p0 + ], + outputs=[d.efc.lo, d.efc.lo_alpha]) # fmt: skip + + if m.opt.cone == types.ConeType.ELLIPTIC: + wp.launch( + linesearch_iterative_init_lo_elliptic0, + dim=(d.nworld, d.njmax,), + inputs=[ + d.ne, d.nf, d.nl, d.nefc, d.efc.Jaref, d.efc.jv, d.efc.quad, + d.efc.done, d.efc.lo_alpha, d.efc.condim + ], + outputs=[d.efc.lo]) # fmt: skip + wp.launch( + linesearch_iterative_init_lo_elliptic1, + dim=(d.nconmax), + inputs=[ + m.opt.impratio, d.ncon, d.contact.friction, d.contact.dim, + d.contact.efc_address, d.contact.worldid, d.efc.D, d.efc.jv, d.efc.quad, + d.efc.done, d.efc.lo_alpha, d.efc.u, d.efc.uu, d.efc.uv, d.efc.vv + ], + outputs=[d.efc.lo]) # fmt: skip + else: + wp.launch( + linesearch_iterative_init_lo_pyramidal, + dim=(d.nworld, d.njmax,), + inputs=[ + d.ne, d.nf, d.nefc, d.efc.Jaref, d.efc.jv, + d.efc.quad, d.efc.done, d.efc.lo_alpha + ], + outputs=[d.efc.lo]) # fmt: skip + + # set the lo/hi interval bounds + + wp.launch( + linesearch_iterative_init_bounds, + dim=(d.nworld,), + inputs=[d.efc.done, d.efc.p0, d.efc.lo, d.efc.lo_alpha], + outputs=[d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha]) # fmt: skip + + for _ in range(m.opt.ls_iterations): + # NOTE: we always launch ls_iterations kernels, but the kernels may early exit if done + # is true. this preserves cudagraph requirements (no dynamic kernel launching) at the + # expense of extra launches + wp.launch( + linesearch_iterative_next_alpha_gauss, + dim=(d.nworld,), + inputs=[ + d.efc.quad_gauss, d.efc.done, d.efc.ls_done, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha + ], + outputs=[ + d.efc.lo_next, d.efc.lo_next_alpha, d.efc.hi_next, d.efc.hi_next_alpha, d.efc.mid, d.efc.mid_alpha + ]) # fmt: skip + + if m.opt.cone == types.ConeType.ELLIPTIC: + wp.launch( + linesearch_iterative_next_quad_elliptic0, + dim=(d.nworld, d.njmax,), + inputs=[ + d.ne, d.nf, d.nl, d.nefc, d.efc.Jaref, d.efc.jv, d.efc.quad, d.efc.done, d.efc.ls_done, + d.efc.lo_next_alpha, d.efc.hi_next_alpha, d.efc.mid_alpha, d.efc.condim + ], + outputs=[d.efc.lo_next, d.efc.hi_next, d.efc.mid]) # fmt: skip + wp.launch( + linesearch_iterative_next_quad_elliptic1, + dim=(d.nconmax), + inputs=[ + m.opt.impratio, d.ncon, d.contact.friction, d.contact.dim, d.contact.efc_address, d.contact.worldid, d.efc.D, + d.efc.jv, d.efc.quad, d.efc.done, d.efc.lo_next_alpha, d.efc.hi_next_alpha, d.efc.mid_alpha, d.efc.u, d.efc.uu, + d.efc.uv, d.efc.vv + ], + outputs=[d.efc.lo_next, d.efc.hi_next, d.efc.mid]) # fmt: skip + else: + wp.launch( + linesearch_iterative_next_quad_pyramidal, + dim=(d.nworld, d.njmax,), + inputs=[ + d.ne, d.nf, d.nefc, d.efc.Jaref, d.efc.jv, d.efc.quad, d.efc.done, d.efc.ls_done, + d.efc.lo_next_alpha, d.efc.hi_next_alpha, d.efc.mid_alpha + ], + outputs=[d.efc.lo_next, d.efc.hi_next, d.efc.mid]) # fmt: skip + + wp.launch( + linesearch_iterative_swap, + dim=(d.nworld,), + inputs=[ + d.efc.gtol, d.efc.done, d.efc.ls_done, d.efc.p0, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha, + d.efc.lo_next, d.efc.lo_next_alpha, d.efc.hi_next, d.efc.hi_next_alpha, d.efc.mid, d.efc.mid_alpha + ], + outputs=[ + d.efc.alpha, d.efc.ls_done, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha + ]) # fmt: skip + + +@wp.kernel +def linesearch_parallel_fused( + # Model: + nlsp: int, + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nefc_in: wp.array(dtype=int), + efc_Jaref_in: wp.array2d(dtype=float), + efc_jv_in: wp.array2d(dtype=float), + efc_quad_in: wp.array2d(dtype=wp.vec3), + efc_quad_gauss_in: wp.array(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_cost_candidate_out: wp.array2d(dtype=float), +): + worldid, alphaid = wp.tid() + + if efc_done_in[worldid]: + return + + efc_quad_total_candidate = efc_quad_gauss_in[worldid] + + alpha = float(alphaid) / float(nlsp - 1) + ne = ne_in[worldid] + nf = nf_in[worldid] + for efcid in range(nefc_in[worldid]): + Jaref = efc_Jaref_in[worldid, efcid] + jv = efc_jv_in[worldid, efcid] + quad = efc_quad_in[worldid, efcid] + + if (Jaref + alpha * jv) < 0.0 or (efcid < ne + nf): + efc_quad_total_candidate += quad + + alpha_sq = alpha * alpha + quad_total0 = efc_quad_total_candidate[0] + quad_total1 = efc_quad_total_candidate[1] + quad_total2 = efc_quad_total_candidate[2] + + efc_cost_candidate_out[worldid, alphaid] = alpha_sq * quad_total2 + alpha * quad_total1 + quad_total0 + + +@wp.kernel +def linesearch_parallel_best_alpha( + # Model: + nlsp: int, + # Data in: + efc_done_in: wp.array(dtype=bool), + efc_cost_candidate_in: wp.array2d(dtype=float), + # Data out: + efc_alpha_out: wp.array(dtype=float), +): + worldid = wp.tid() + + if efc_done_in[worldid]: + return + + # TODO(team): investigate alternatives to wp.argmin + # TODO(thowell): how did this use to work? + bestid = int(0) + best_cost = float(wp.inf) + for i in range(nlsp): + cost = efc_cost_candidate_in[worldid, i] + if cost < best_cost: + best_cost = cost + bestid = i + + efc_alpha_out[worldid] = float(bestid) / float(nlsp - 1) + + +def _linesearch_parallel(m: types.Model, d: types.Data): + wp.launch( + linesearch_parallel_fused, + dim=(d.nworld, m.nlsp), + inputs=[ + m.nlsp, + d.ne, + d.nf, + d.nefc, + d.efc.Jaref, + d.efc.jv, + d.efc.quad, + d.efc.quad_gauss, + d.efc.done, + ], + outputs=[d.efc.cost_candidate], + ) + + wp.launch( + linesearch_parallel_best_alpha, + dim=(d.nworld), + inputs=[m.nlsp, d.efc.done, d.efc.cost_candidate], + outputs=[d.efc.alpha], + ) + + +@wp.kernel +def linesearch_zero_jv( + # Data in: + nefc_in: wp.array(dtype=int), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_jv_out: wp.array2d(dtype=float), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + if efc_done_in[worldid]: + return + + efc_jv_out[worldid, efcid] = 0.0 + + +@cache_kernel +def linesearch_jv_fused(nv: int, dofs_per_thread: int): + @nested_kernel + def kernel( + # Data in: + nefc_in: wp.array(dtype=int), + efc_J_in: wp.array3d(dtype=float), + efc_search_in: wp.array2d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_jv_out: wp.array2d(dtype=float), + ): + worldid, efcid, dofstart = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + if efc_done_in[worldid]: + return + + jv_out = float(0.0) + + if wp.static(dofs_per_thread >= nv): + for i in range(wp.static(min(dofs_per_thread, nv))): + jv_out += efc_J_in[worldid, efcid, i] * efc_search_in[worldid, i] + efc_jv_out[worldid, efcid] = jv_out + + else: + for i in range(wp.static(dofs_per_thread)): + ii = dofstart * wp.static(dofs_per_thread) + i + if ii < nv: + jv_out += efc_J_in[worldid, efcid, ii] * efc_search_in[worldid, ii] + wp.atomic_add(efc_jv_out, worldid, efcid, jv_out) + + return kernel + + +@wp.kernel +def linesearch_init_quad_gauss( + # Model: + nv: int, + # Data in: + qfrc_smooth_in: wp.array2d(dtype=float), + efc_Ma_in: wp.array2d(dtype=float), + efc_search_in: wp.array2d(dtype=float), + efc_gauss_in: wp.array(dtype=float), + efc_mv_in: wp.array2d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_quad_gauss_out: wp.array(dtype=wp.vec3), +): + worldid = wp.tid() + if efc_done_in[worldid]: + return + + quad_gauss_0 = efc_gauss_in[worldid] + quad_gauss_1 = float(0.0) + quad_gauss_2 = float(0.0) + for i in range(nv): + search = efc_search_in[worldid, i] + quad_gauss_1 += search * (efc_Ma_in[worldid, i] - qfrc_smooth_in[worldid, i]) + quad_gauss_2 += 0.5 * search * efc_mv_in[worldid, i] + + efc_quad_gauss_out[worldid] = wp.vec3(quad_gauss_0, quad_gauss_1, quad_gauss_2) + + +@wp.kernel +def linesearch_init_quad( + # Data in: + nefc_in: wp.array(dtype=int), + efc_D_in: wp.array2d(dtype=float), + efc_frictionloss_in: wp.array2d(dtype=float), + efc_Jaref_in: wp.array2d(dtype=float), + efc_jv_in: wp.array2d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # In: + disable_floss: bool, + # Data out: + efc_quad_out: wp.array2d(dtype=wp.vec3), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + if efc_done_in[worldid]: + return + + Jaref = efc_Jaref_in[worldid, efcid] + jv = efc_jv_in[worldid, efcid] + efc_D = efc_D_in[worldid, efcid] + floss = efc_frictionloss_in[worldid, efcid] + + if floss > 0.0 and not disable_floss: + rf = math.safe_div(floss, efc_D) + if Jaref <= -rf: + efc_quad_out[worldid, efcid] = wp.vec3(floss * (-0.5 * rf - Jaref), -floss * jv, 0.0) + return + elif Jaref >= rf: + efc_quad_out[worldid, efcid] = wp.vec3(floss * (-0.5 * rf + Jaref), floss * jv, 0.0) + return + + efc_quad_out[worldid, efcid] = wp.vec3(0.5 * Jaref * Jaref * efc_D, jv * Jaref * efc_D, 0.5 * jv * jv * efc_D) + + +@wp.kernel +def linesearch_quad_elliptic( + # Data in: + ncon_in: wp.array(dtype=int), + contact_friction_in: wp.array(dtype=types.vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_jv_in: wp.array2d(dtype=float), + efc_quad_in: wp.array2d(dtype=wp.vec3), + efc_done_in: wp.array(dtype=bool), + efc_u_in: wp.array(dtype=types.vec6), + # Data out: + efc_quad_out: wp.array2d(dtype=wp.vec3), + efc_uv_out: wp.array(dtype=float), + efc_vv_out: wp.array(dtype=float), +): + conid, dimid = wp.tid() + dimid += 1 + + if conid >= ncon_in[0]: + return + + worldid = contact_worldid_in[conid] + if efc_done_in[worldid]: + return + + condim = contact_dim_in[conid] + + if condim == 1 or (dimid >= condim): + return + + efcid0 = contact_efc_address_in[conid, 0] + efcid = contact_efc_address_in[conid, dimid] + + # complete vector quadratic (for bottom zone) + wp.atomic_add(efc_quad_out, worldid, efcid0, efc_quad_in[worldid, efcid]) + + # rescale to make primal cone circular + u = efc_u_in[conid][dimid] + v = efc_jv_in[worldid, efcid] * contact_friction_in[conid][dimid - 1] + wp.atomic_add(efc_uv_out, conid, u * v) + wp.atomic_add(efc_vv_out, conid, v * v) + + +@wp.kernel +def linesearch_qacc_ma( + # Data in: + efc_search_in: wp.array2d(dtype=float), + efc_mv_in: wp.array2d(dtype=float), + efc_alpha_in: wp.array(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + qacc_out: wp.array2d(dtype=float), + efc_Ma_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + + if efc_done_in[worldid]: + return + + alpha = efc_alpha_in[worldid] + qacc_out[worldid, dofid] += alpha * efc_search_in[worldid, dofid] + efc_Ma_out[worldid, dofid] += alpha * efc_mv_in[worldid, dofid] + + +@wp.kernel +def linesearch_jaref( + # Data in: + nefc_in: wp.array(dtype=int), + efc_jv_in: wp.array2d(dtype=float), + efc_alpha_in: wp.array(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_Jaref_out: wp.array2d(dtype=float), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + if efc_done_in[worldid]: + return + + efc_Jaref_out[worldid, efcid] += efc_alpha_in[worldid] * efc_jv_in[worldid, efcid] + + +@event_scope +def _linesearch(m: types.Model, d: types.Data): + # mv = qM @ search + support.mul_m(m, d, d.efc.mv, d.efc.search, d.efc.done) + + # jv = efc_J @ search + # TODO(team): is there a better way of doing batched matmuls with dynamic array sizes? + + # if we are only using 1 thread, it makes sense to do more dofs as we can also skip the + # init kernel. For more than 1 thread, dofs_per_thread is lower for better load balancing. + + if m.nv > 50: + dofs_per_thread = 20 + else: + dofs_per_thread = 50 + + threads_per_efc = ceil(m.nv / dofs_per_thread) + # we need to clear the jv array if we're doing atomic adds. + if threads_per_efc > 1: + wp.launch( + linesearch_zero_jv, + dim=(d.nworld, d.njmax), + inputs=[d.nefc, d.efc.done], + outputs=[d.efc.jv], + ) + + wp.launch( + linesearch_jv_fused(m.nv, dofs_per_thread), + dim=(d.nworld, d.njmax, threads_per_efc), + inputs=[d.nefc, d.efc.J, d.efc.search, d.efc.done], + outputs=[d.efc.jv], + ) + + # prepare quadratics + # quad_gauss = [gauss, search.T @ Ma - search.T @ qfrc_smooth, 0.5 * search.T @ mv] + wp.launch( + linesearch_init_quad_gauss, + dim=(d.nworld), + inputs=[m.nv, d.qfrc_smooth, d.efc.Ma, d.efc.search, d.efc.gauss, d.efc.mv, d.efc.done], + outputs=[d.efc.quad_gauss], + ) + + # quad = [0.5 * Jaref * Jaref * efc_D, jv * Jaref * efc_D, 0.5 * jv * jv * efc_D] + + disable_floss = m.opt.disableflags & types.DisableBit.FRICTIONLOSS + wp.launch( + linesearch_init_quad, + dim=(d.nworld, d.njmax), + inputs=[ + d.nefc, + d.efc.D, + d.efc.frictionloss, + d.efc.Jaref, + d.efc.jv, + d.efc.done, + disable_floss, + ], + outputs=[d.efc.quad], + ) + + if m.opt.cone == types.ConeType.ELLIPTIC: + d.efc.uv.zero_() + d.efc.vv.zero_() + wp.launch( + linesearch_quad_elliptic, + dim=(d.nconmax, m.condim_max - 1), + inputs=[ + d.ncon, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.contact.worldid, + d.efc.jv, + d.efc.quad, + d.efc.done, + d.efc.u, + ], + outputs=[d.efc.quad, d.efc.uv, d.efc.vv], + ) + + if m.opt.ls_parallel: + _linesearch_parallel(m, d) + else: + _linesearch_iterative(m, d) + + wp.launch( + linesearch_qacc_ma, + dim=(d.nworld, m.nv), + inputs=[d.efc.search, d.efc.mv, d.efc.alpha, d.efc.done], + outputs=[d.qacc, d.efc.Ma], + ) + + wp.launch( + linesearch_jaref, + dim=(d.nworld, d.njmax), + inputs=[d.nefc, d.efc.jv, d.efc.alpha, d.efc.done], + outputs=[d.efc.Jaref], + ) + + +@wp.kernel +def solve_init_efc( + # Data out: + solver_niter_out: wp.array(dtype=int), + efc_search_dot_out: wp.array(dtype=float), + efc_cost_out: wp.array(dtype=float), + efc_done_out: wp.array(dtype=bool), +): + worldid = wp.tid() + efc_cost_out[worldid] = wp.inf + solver_niter_out[worldid] = 0 + efc_done_out[worldid] = False + efc_search_dot_out[worldid] = 0.0 + + +@wp.kernel +def solve_init_jaref( + # Model: + nv: int, + # Data in: + nefc_in: wp.array(dtype=int), + qacc_in: wp.array2d(dtype=float), + efc_J_in: wp.array3d(dtype=float), + efc_aref_in: wp.array2d(dtype=float), + # Data out: + efc_Jaref_out: wp.array2d(dtype=float), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + jaref = float(0.0) + for i in range(nv): + jaref += efc_J_in[worldid, efcid, i] * qacc_in[worldid, i] + + efc_Jaref_out[worldid, efcid] = jaref - efc_aref_in[worldid, efcid] + + +@wp.kernel +def solve_init_search( + # Data in: + efc_Mgrad_in: wp.array2d(dtype=float), + # Data out: + efc_search_out: wp.array2d(dtype=float), + efc_search_dot_out: wp.array(dtype=float), +): + worldid, dofid = wp.tid() + search = -1.0 * efc_Mgrad_in[worldid, dofid] + efc_search_out[worldid, dofid] = search + wp.atomic_add(efc_search_dot_out, worldid, search * search) + + +@wp.kernel +def update_constraint_init_cost( + # Data in: + efc_cost_in: wp.array(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_gauss_out: wp.array(dtype=float), + efc_cost_out: wp.array(dtype=float), + efc_prev_cost_out: wp.array(dtype=float), +): + worldid = wp.tid() + + if efc_done_in[worldid]: + return + + efc_gauss_out[worldid] = 0.0 + efc_prev_cost_out[worldid] = efc_cost_in[worldid] + efc_cost_out[worldid] = 0.0 + + +@wp.kernel +def update_constraint_efc_pyramidal( + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nefc_in: wp.array(dtype=int), + efc_D_in: wp.array2d(dtype=float), + efc_frictionloss_in: wp.array2d(dtype=float), + efc_Jaref_in: wp.array2d(dtype=float), + # In: + disable_floss: int, + # Data out: + efc_force_out: wp.array2d(dtype=float), + efc_cost_out: wp.array(dtype=float), + efc_active_out: wp.array2d(dtype=bool), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + efc_D = efc_D_in[worldid, efcid] + Jaref = efc_Jaref_in[worldid, efcid] + + cost = 0.5 * efc_D * Jaref * Jaref + efc_force = -efc_D * Jaref + + ne = ne_in[worldid] + nf = nf_in[worldid] + + if efcid < ne: + # equality + pass + elif efcid < ne + nf and not disable_floss: + # friction + f = efc_frictionloss_in[worldid, efcid] + if f > 0.0: + rf = math.safe_div(f, efc_D) + if Jaref <= -rf: + efc_force_out[worldid, efcid] = f + efc_active_out[worldid, efcid] = False + wp.atomic_add(efc_cost_out, worldid, -0.5 * rf - Jaref) + return + elif Jaref >= rf: + efc_force_out[worldid, efcid] = -f + efc_active_out[worldid, efcid] = False + wp.atomic_add(efc_cost_out, worldid, -0.5 * rf + Jaref) + return + else: + # limit, contact + if Jaref >= 0.0: + efc_force_out[worldid, efcid] = 0.0 + efc_active_out[worldid, efcid] = False + return + + efc_force_out[worldid, efcid] = efc_force + efc_active_out[worldid, efcid] = True + wp.atomic_add(efc_cost_out, worldid, cost) + + +@wp.kernel +def update_constraint_u_elliptic( + # Model: + opt_impratio: wp.array(dtype=float), + # Data in: + ncon_in: wp.array(dtype=int), + contact_friction_in: wp.array(dtype=types.vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_Jaref_in: wp.array2d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_u_out: wp.array(dtype=types.vec6), + efc_uu_out: wp.array(dtype=float), + efc_condim_out: wp.array2d(dtype=int), +): + conid, dimid = wp.tid() + + if conid >= ncon_in[0]: + return + + worldid = contact_worldid_in[conid] + if efc_done_in[worldid]: + return + + efcid = contact_efc_address_in[conid, dimid] + + condim = contact_dim_in[conid] + efc_condim_out[worldid, efcid] = condim + + if condim == 1: + return + + if dimid < condim: + if dimid == 0: + fri = contact_friction_in[conid][0] / wp.sqrt(opt_impratio[worldid]) + else: + fri = contact_friction_in[conid][dimid - 1] + u = efc_Jaref_in[worldid, efcid] * fri + efc_u_out[conid][dimid] = u + if dimid > 0: + wp.atomic_add(efc_uu_out, conid, u * u) + + +@wp.kernel +def update_constraint_active_elliptic_bottom_zone( + # Model: + opt_impratio: wp.array(dtype=float), + # Data in: + ncon_in: wp.array(dtype=int), + contact_friction_in: wp.array(dtype=types.vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_done_in: wp.array(dtype=bool), + efc_u_in: wp.array(dtype=types.vec6), + efc_uu_in: wp.array(dtype=float), + # Data out: + efc_active_out: wp.array2d(dtype=bool), +): + conid, dimid = wp.tid() + + if conid >= ncon_in[0]: + return + + worldid = contact_worldid_in[conid] + if efc_done_in[worldid]: + return + + condim = contact_dim_in[conid] + if condim == 1: + return + + mu = contact_friction_in[conid][0] / wp.sqrt(opt_impratio[worldid]) + n = efc_u_in[conid][0] + tt = efc_uu_in[conid] + if tt <= 0.0: + t = 0.0 + else: + t = wp.sqrt(tt) + + # bottom zone: quadratic + bottom_zone = ((t <= 0.0) and (n < 0.0)) or ((t > 0.0) and ((mu * n + t) <= 0.0)) + + # update active + efcid = contact_efc_address_in[conid, dimid] + efc_active_out[worldid, efcid] = bottom_zone + + +@wp.kernel +def update_constraint_efc_elliptic0( + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nl_in: wp.array(dtype=int), + nefc_in: wp.array(dtype=int), + efc_D_in: wp.array2d(dtype=float), + efc_frictionloss_in: wp.array2d(dtype=float), + efc_Jaref_in: wp.array2d(dtype=float), + efc_active_in: wp.array2d(dtype=bool), + efc_done_in: wp.array(dtype=bool), + # In: + disable_floss: int, + # Data out: + efc_force_out: wp.array2d(dtype=float), + efc_cost_out: wp.array(dtype=float), + efc_active_out: wp.array2d(dtype=bool), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + if efc_done_in[worldid]: + return + + efc_D = efc_D_in[worldid, efcid] + Jaref = efc_Jaref_in[worldid, efcid] + + ne = ne_in[worldid] + nf = nf_in[worldid] + nl = nl_in[worldid] + + if efcid < ne: + # equality + efc_active_out[worldid, efcid] = True + elif efcid < ne + nf and not disable_floss: + # friction + f = efc_frictionloss_in[worldid, efcid] + if f > 0.0: + rf = math.safe_div(f, efc_D) + if Jaref <= -rf: + efc_force_out[worldid, efcid] = f + efc_active_out[worldid, efcid] = False + wp.atomic_add(efc_cost_out, worldid, -0.5 * rf - Jaref) + return + elif Jaref >= rf: + efc_force_out[worldid, efcid] = -f + efc_active_out[worldid, efcid] = False + wp.atomic_add(efc_cost_out, worldid, -0.5 * rf + Jaref) + return + elif efcid < ne + nf + nl: + # limits + if Jaref < 0.0: + efc_active_out[worldid, efcid] = True + else: + efc_force_out[worldid, efcid] = 0.0 + efc_active_out[worldid, efcid] = False + return + else: + # contact + if not efc_active_in[worldid, efcid]: # calculated by solve_active_elliptic_bottom_zone + efc_force_out[worldid, efcid] = 0.0 + return + + efc_force_out[worldid, efcid] = -efc_D * Jaref + wp.atomic_add(efc_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref) + + +@wp.kernel +def update_constraint_efc_elliptic1( + # Model: + opt_impratio: wp.array(dtype=float), + # Data in: + ncon_in: wp.array(dtype=int), + contact_friction_in: wp.array(dtype=types.vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_D_in: wp.array2d(dtype=float), + efc_done_in: wp.array(dtype=bool), + efc_u_in: wp.array(dtype=types.vec6), + efc_uu_in: wp.array(dtype=float), + # Data out: + efc_force_out: wp.array2d(dtype=float), + efc_cost_out: wp.array(dtype=float), +): + conid, dimid = wp.tid() + + if conid >= ncon_in[0]: + return + + worldid = contact_worldid_in[conid] + if efc_done_in[worldid]: + return + + condim = contact_dim_in[conid] + + if condim == 1 or dimid >= condim: + return + + friction = contact_friction_in[conid] + efcid = contact_efc_address_in[conid, dimid] + + mu = friction[0] / wp.sqrt(opt_impratio[worldid]) + n = efc_u_in[conid][0] + tt = efc_uu_in[conid] + if tt <= 0.0: + t = 0.0 + else: + t = wp.sqrt(tt) + + # middle zone: cone + middle_zone = (t > 0.0) and (n < (mu * t)) and ((mu * n + t) > 0.0) + + # tangent and friction for middle zone: + if middle_zone: + efcid0 = contact_efc_address_in[conid, 0] + mu2 = mu * mu + dm = efc_D_in[worldid, efcid0] / wp.max(mu2 * float(1.0 + mu2), types.MJ_MINVAL) + + nmt = n - mu * t + + force = -dm * nmt * mu + if dimid > 0: + force_fri = -force / t + force_fri *= efc_u_in[conid][dimid] * friction[dimid - 1] + efc_force_out[worldid, efcid] += force_fri + else: + efc_force_out[worldid, efcid] += force + worldid = contact_worldid_in[conid] + wp.atomic_add(efc_cost_out, worldid, 0.5 * dm * nmt * nmt) + + +@wp.kernel +def update_constraint_zero_qfrc_constraint( + # Data in: + efc_done_in: wp.array(dtype=bool), + # Data out: + qfrc_constraint_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + + if efc_done_in[worldid]: + return + + qfrc_constraint_out[worldid, dofid] = 0.0 + + +@wp.kernel +def update_constraint_init_qfrc_constraint( + # Model: + nv: int, + # Data in: + nefc_in: wp.array(dtype=int), + efc_J_in: wp.array3d(dtype=float), + efc_force_in: wp.array2d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + qfrc_constraint_out: wp.array2d(dtype=float), +): + worldid, efcid = wp.tid() + + if efcid >= nefc_in[worldid]: + return + + if efc_done_in[worldid]: + return + + force = efc_force_in[worldid, efcid] + for i in range(nv): + wp.atomic_add( + qfrc_constraint_out[worldid], + i, + efc_J_in[worldid, efcid, i] * force, + ) + + +@wp.kernel +def update_constraint_gauss_cost( + # Data in: + qacc_in: wp.array2d(dtype=float), + qfrc_smooth_in: wp.array2d(dtype=float), + qacc_smooth_in: wp.array2d(dtype=float), + efc_Ma_in: wp.array2d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_gauss_out: wp.array(dtype=float), + efc_cost_out: wp.array(dtype=float), +): + worldid, dofid = wp.tid() + + if efc_done_in[worldid]: + return + + gauss_cost = ( + 0.5 + * (efc_Ma_in[worldid, dofid] - qfrc_smooth_in[worldid, dofid]) + * (qacc_in[worldid, dofid] - qacc_smooth_in[worldid, dofid]) + ) + wp.atomic_add(efc_gauss_out, worldid, gauss_cost) + wp.atomic_add(efc_cost_out, worldid, gauss_cost) + + +def _update_constraint(m: types.Model, d: types.Data): + """Update constraint arrays after each solve iteration.""" + + disable_floss = m.opt.disableflags & types.DisableBit.FRICTIONLOSS + + wp.launch( + update_constraint_init_cost, + dim=(d.nworld), + inputs=[d.efc.cost, d.efc.done], + outputs=[d.efc.gauss, d.efc.cost, d.efc.prev_cost], + ) + + if m.opt.cone == types.ConeType.PYRAMIDAL: + wp.launch( + update_constraint_efc_pyramidal, + dim=(d.nworld, d.njmax), + inputs=[ + d.ne, + d.nf, + d.nefc, + d.efc.D, + d.efc.frictionloss, + d.efc.Jaref, + disable_floss, + ], + outputs=[d.efc.force, d.efc.cost, d.efc.active], + ) + elif m.opt.cone == types.ConeType.ELLIPTIC: + d.efc.uu.zero_() + d.efc.active.zero_() + d.efc.condim.fill_(-1) + wp.launch( + update_constraint_u_elliptic, + dim=(d.nconmax, m.condim_max), + inputs=[ + m.opt.impratio, + d.ncon, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.contact.worldid, + d.efc.Jaref, + d.efc.done, + ], + outputs=[d.efc.u, d.efc.uu, d.efc.condim], + ) + wp.launch( + update_constraint_active_elliptic_bottom_zone, + dim=(d.nconmax, m.condim_max), + inputs=[ + m.opt.impratio, + d.ncon, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.contact.worldid, + d.efc.done, + d.efc.u, + d.efc.uu, + ], + outputs=[d.efc.active], + ) + wp.launch( + update_constraint_efc_elliptic0, + dim=(d.nworld, d.njmax), + inputs=[ + d.ne, + d.nf, + d.nl, + d.nefc, + d.efc.D, + d.efc.frictionloss, + d.efc.Jaref, + d.efc.active, + d.efc.done, + disable_floss, + ], + outputs=[d.efc.force, d.efc.cost, d.efc.active], + ) + wp.launch( + update_constraint_efc_elliptic1, + dim=(d.nconmax, m.condim_max), + inputs=[ + m.opt.impratio, + d.ncon, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.contact.worldid, + d.efc.D, + d.efc.done, + d.efc.u, + d.efc.uu, + ], + outputs=[d.efc.force, d.efc.cost], + ) + else: + raise ValueError(f"Unknown cone type: {m.opt.cone}") + + # qfrc_constraint = efc_J.T @ efc_force + wp.launch( + update_constraint_zero_qfrc_constraint, + dim=(d.nworld, m.nv), + inputs=[d.efc.done], + outputs=[d.qfrc_constraint], + ) + + wp.launch( + update_constraint_init_qfrc_constraint, + dim=(d.nworld, d.njmax), + inputs=[m.nv, d.nefc, d.efc.J, d.efc.force, d.efc.done], + outputs=[d.qfrc_constraint], + ) + + # gauss = 0.5 * (Ma - qfrc_smooth).T @ (qacc - qacc_smooth) + wp.launch( + update_constraint_gauss_cost, + dim=(d.nworld, m.nv), + inputs=[d.qacc, d.qfrc_smooth, d.qacc_smooth, d.efc.Ma, d.efc.done], + outputs=[d.efc.gauss, d.efc.cost], + ) + + +@wp.kernel +def update_gradient_zero_grad_dot( + # Data in: + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_grad_dot_out: wp.array(dtype=float), +): + worldid = wp.tid() + + if efc_done_in[worldid]: + return + + efc_grad_dot_out[worldid] = 0.0 + + +@wp.kernel +def update_gradient_grad( + # Data in: + qfrc_smooth_in: wp.array2d(dtype=float), + qfrc_constraint_in: wp.array2d(dtype=float), + efc_Ma_in: wp.array2d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_grad_out: wp.array2d(dtype=float), + efc_grad_dot_out: wp.array(dtype=float), +): + worldid, dofid = wp.tid() + + if efc_done_in[worldid]: + return + + grad = efc_Ma_in[worldid, dofid] - qfrc_smooth_in[worldid, dofid] - qfrc_constraint_in[worldid, dofid] + efc_grad_out[worldid, dofid] = grad + wp.atomic_add(efc_grad_dot_out, worldid, grad * grad) + + +@wp.kernel +def update_gradient_zero_h_lower( + # Model: + dof_tri_row: wp.array(dtype=int), + dof_tri_col: wp.array(dtype=int), + # Data in: + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_h_out: wp.array3d(dtype=float), +): + worldid, elementid = wp.tid() + + if efc_done_in[worldid]: + return + + rowid = dof_tri_row[elementid] + colid = dof_tri_col[elementid] + efc_h_out[worldid, rowid, colid] = 0.0 + + +@wp.kernel +def update_gradient_set_h_qM_lower_sparse( + # Model: + qM_fullm_i: wp.array(dtype=int), + qM_fullm_j: wp.array(dtype=int), + # Data in: + qM_in: wp.array3d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_h_out: wp.array3d(dtype=float), +): + worldid, elementid = wp.tid() + + if efc_done_in[worldid]: + return + + i = qM_fullm_i[elementid] + j = qM_fullm_j[elementid] + efc_h_out[worldid, i, j] = qM_in[worldid, 0, elementid] + + +@wp.kernel +def update_gradient_copy_lower_triangle( + # Model: + dof_tri_row: wp.array(dtype=int), + dof_tri_col: wp.array(dtype=int), + # Data in: + qM_in: wp.array3d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_h_out: wp.array3d(dtype=float), +): + worldid, elementid = wp.tid() + + if efc_done_in[worldid]: + return + + rowid = dof_tri_row[elementid] + colid = dof_tri_col[elementid] + efc_h_out[worldid, rowid, colid] = qM_in[worldid, rowid, colid] + + +@wp.kernel +def update_gradient_JTDAJ( + # Model: + dof_tri_row: wp.array(dtype=int), + dof_tri_col: wp.array(dtype=int), + # Data in: + nefc_in: wp.array(dtype=int), + efc_J_in: wp.array3d(dtype=float), + efc_D_in: wp.array2d(dtype=float), + efc_active_in: wp.array2d(dtype=bool), + efc_done_in: wp.array(dtype=bool), + # In: + # Data out: + efc_h_out: wp.array3d(dtype=float), +): + worldid, elementid = wp.tid() + + if efc_done_in[worldid]: + return + + nefc = nefc_in[worldid] + + dofi = dof_tri_row[elementid] + dofj = dof_tri_col[elementid] + + for efcid in range(nefc): + efc_D = efc_D_in[worldid, efcid] + active = efc_active_in[worldid, efcid] + + if efc_D == 0.0 or not active: + continue + + # TODO(team): sparse efc_J + value = efc_J_in[worldid, efcid, dofi] * efc_J_in[worldid, efcid, dofj] * efc_D + if value != 0.0: + wp.atomic_add(efc_h_out[worldid, dofi], dofj, value) + + +@wp.kernel +def update_gradient_JTCJ( + # Model: + opt_impratio: wp.array(dtype=float), + dof_tri_row: wp.array(dtype=int), + dof_tri_col: wp.array(dtype=int), + # Data in: + nconmax_in: int, + ncon_in: wp.array(dtype=int), + contact_friction_in: wp.array(dtype=types.vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_J_in: wp.array3d(dtype=float), + efc_D_in: wp.array2d(dtype=float), + efc_done_in: wp.array(dtype=bool), + efc_u_in: wp.array(dtype=types.vec6), + efc_uu_in: wp.array(dtype=float), + # In: + nblocks_perblock: int, + dim_block: int, + # Data out: + efc_h_out: wp.array3d(dtype=float), +): + conid_start, elementid = wp.tid() + + dof1id = dof_tri_row[elementid] + dof2id = dof_tri_col[elementid] + + for i in range(nblocks_perblock): + conid = conid_start + i * dim_block + + if conid >= min(ncon_in[0], nconmax_in): + return + + worldid = contact_worldid_in[conid] + if efc_done_in[worldid]: + continue + + condim = contact_dim_in[conid] + + if condim == 1: + continue + + fri = contact_friction_in[conid] + mu = fri[0] / wp.sqrt(opt_impratio[worldid]) + n = efc_u_in[conid][0] + tt = efc_uu_in[conid] + if tt <= 0.0: + t = 0.0 + else: + t = wp.sqrt(tt) + + middle_zone = (t > 0) and (n < (mu * t)) and ((mu * n + t) > 0.0) + + if not middle_zone: + continue + + t = wp.max(t, types.MJ_MINVAL) + ttt = wp.max(t * t * t, types.MJ_MINVAL) + + mu2 = mu * mu + efc0 = contact_efc_address_in[conid, 0] + dm = efc_D_in[worldid, efc0] / wp.max(mu2 * (1.0 + mu2), types.MJ_MINVAL) + + if dm == 0.0: + continue + + u = efc_u_in[conid] + + efc_h = float(0.0) + + for dim1id in range(condim): + if dim1id == 0: + efcid1 = efc0 + else: + efcid1 = contact_efc_address_in[conid, dim1id] + + efc_J11 = efc_J_in[worldid, efcid1, dof1id] + efc_J12 = efc_J_in[worldid, efcid1, dof2id] + + ui = u[dim1id] + + for dim2id in range(0, dim1id + 1): + if dim2id == 0: + efcid2 = efc0 + else: + efcid2 = contact_efc_address_in[conid, dim2id] + + efc_J21 = efc_J_in[worldid, efcid2, dof1id] + efc_J22 = efc_J_in[worldid, efcid2, dof2id] + + uj = u[dim2id] + + # set first row/column: (1, -mu/t * u) + if dim1id == 0 and dim2id == 0: + hcone = 1.0 + elif dim1id == 0: + hcone = -mu / t * uj + elif dim2id == 0: + hcone = -mu / t * ui + else: + hcone = mu * n / ttt * ui * uj + + # add to diagonal: mu^2 - mu * n / t + if dim1id == dim2id: + hcone += mu2 - mu * n / t + + # pre and post multiply by diag(mu, friction) scale by dm + if dim1id == 0: + fri1 = mu + else: + fri1 = fri[dim1id - 1] + + if dim2id == 0: + fri2 = mu + else: + fri2 = fri[dim2id - 1] + + hcone *= dm * fri1 * fri2 + + if hcone != 0.0: + efc_h += hcone * efc_J11 * efc_J22 + + if dim1id != dim2id: + efc_h += hcone * efc_J12 * efc_J21 + + worldid = contact_worldid_in[conid] + efc_h_out[worldid, dof1id, dof2id] += efc_h + + +@cache_kernel +def update_gradient_cholesky(tile_size: int): + @nested_kernel + def kernel( + # Data in: + efc_grad_in: wp.array2d(dtype=float), + efc_h_in: wp.array3d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_Mgrad_out: wp.array2d(dtype=float), + ): + worldid = wp.tid() + TILE_SIZE = wp.static(tile_size) + + if efc_done_in[worldid]: + return + + mat_tile = wp.tile_load(efc_h_in[worldid], shape=(TILE_SIZE, TILE_SIZE)) + fact_tile = wp.tile_cholesky(mat_tile) + input_tile = wp.tile_load(efc_grad_in[worldid], shape=TILE_SIZE) + output_tile = wp.tile_cholesky_solve(fact_tile, input_tile) + wp.tile_store(efc_Mgrad_out[worldid], output_tile) + + return kernel + + +@cache_kernel +def update_gradient_cholesky_blocked(tile_size: int): + @nested_kernel + def kernel( + # Data in: + efc_grad_in: wp.array3d(dtype=float), + efc_h_in: wp.array3d(dtype=float), + efc_done_in: wp.array(dtype=bool), + matrix_size: int, + cholesky_L_tmp: wp.array3d(dtype=float), + cholesky_y_tmp: wp.array3d(dtype=float), + # Data out: + efc_Mgrad_out: wp.array3d(dtype=float), + ): + worldid, tid_block = wp.tid() + TILE_SIZE = wp.static(tile_size) + + if efc_done_in[worldid]: + return + + wp.static(create_blocked_cholesky_func(TILE_SIZE))(tid_block, efc_h_in[worldid], matrix_size, cholesky_L_tmp[worldid]) + wp.static(create_blocked_cholesky_solve_func(TILE_SIZE))( + tid_block, cholesky_L_tmp[worldid], efc_grad_in[worldid], cholesky_y_tmp[worldid], matrix_size, efc_Mgrad_out[worldid] + ) + + return kernel + + +def _update_gradient(m: types.Model, d: types.Data): + # grad = Ma - qfrc_smooth - qfrc_constraint + wp.launch(update_gradient_zero_grad_dot, dim=(d.nworld), inputs=[d.efc.done], outputs=[d.efc.grad_dot]) + + wp.launch( + update_gradient_grad, + dim=(d.nworld, m.nv), + inputs=[d.qfrc_smooth, d.qfrc_constraint, d.efc.Ma, d.efc.done], + outputs=[d.efc.grad, d.efc.grad_dot], + ) + + if m.opt.solver == types.SolverType.CG: + smooth.solve_m(m, d, d.efc.Mgrad, d.efc.grad) + elif m.opt.solver == types.SolverType.NEWTON: + # h = qM + (efc_J.T * efc_D * active) @ efc_J + if m.opt.is_sparse: + wp.launch( + update_gradient_zero_h_lower, + dim=(d.nworld, m.dof_tri_row.size), + inputs=[m.dof_tri_row, m.dof_tri_col, d.efc.done], + outputs=[d.efc.h], + ) + wp.launch( + update_gradient_set_h_qM_lower_sparse, + dim=(d.nworld, m.qM_fullm_i.size), + inputs=[m.qM_fullm_i, m.qM_fullm_j, d.qM, d.efc.done], + outputs=[d.efc.h], + ) + else: + wp.launch( + update_gradient_copy_lower_triangle, + dim=(d.nworld, m.dof_tri_row.size), + inputs=[m.dof_tri_row, m.dof_tri_col, d.qM, d.efc.done], + outputs=[d.efc.h], + ) + + wp.launch( + update_gradient_JTDAJ, + dim=(d.nworld, m.dof_tri_row.size), + inputs=[ + m.dof_tri_row, + m.dof_tri_col, + d.nefc, + d.efc.J, + d.efc.D, + d.efc.active, + d.efc.done, + ], + outputs=[d.efc.h], + ) + + if m.opt.cone == types.ConeType.ELLIPTIC: + # Optimization: launching update_gradient_JTCJ with limited number of blocks on a GPU. + # Profiling suggests that only a fraction of blocks out of the original + # d.njmax blocks do the actual work. It aims to minimize #CTAs with no + # effective work. It launches with #blocks that's proportional to the number + # of SMs on the GPU. We can now query the SM count: + # https://github.com/NVIDIA/warp/commit/f3814e7e5459e5fd13032cf0fddb3daddd510f30 + + # make dim_block and nblocks_perblock static for update_gradient_JTCJ to allow + # loop unrolling + if wp.get_device().is_cuda: + sm_count = wp.get_device().sm_count + + # Here we assume one block has 256 threads. We use a factor of 6, which + # can be changed in the future to fine-tune the perf. The optimal factor will + # depend on the kernel's occupancy, which determines how many blocks can + # simultaneously run on the SM. TODO: This factor can be tuned further. + dim_block = ceil((sm_count * 6 * 256) / m.dof_tri_row.size) + else: + # fall back for CPU + dim_block = d.nconmax + + nblocks_perblock = int((d.nconmax + dim_block - 1) / dim_block) + + wp.launch( + update_gradient_JTCJ, + dim=(dim_block, m.dof_tri_row.size), + inputs=[ + m.opt.impratio, + m.dof_tri_row, + m.dof_tri_col, + d.nconmax, + d.ncon, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.contact.worldid, + d.efc.J, + d.efc.D, + d.efc.done, + d.efc.u, + d.efc.uu, + nblocks_perblock, + dim_block, + ], + outputs=[d.efc.h], + ) + + # TODO(team): Define good threshold for blocked vs non-blocked cholesky + if m.nv < 32: + wp.launch_tiled( + update_gradient_cholesky(m.nv), + dim=(d.nworld,), + inputs=[d.efc.grad, d.efc.h, d.efc.done], + outputs=[d.efc.Mgrad], + block_dim=m.block_dim.update_gradient_cholesky, + ) + else: + wp.launch_tiled( + update_gradient_cholesky_blocked(32), + dim=(d.nworld,), + inputs=[ + d.efc.grad.reshape(shape=(d.nworld, m.nv, 1)), + d.efc.h, + d.efc.done, + m.nv, + d.efc.cholesky_L_tmp, + d.efc.cholesky_y_tmp.reshape(shape=(d.nworld, m.nv, 1)), + ], + outputs=[d.efc.Mgrad.reshape(shape=(d.nworld, m.nv, 1))], + block_dim=m.block_dim.update_gradient_cholesky, + ) + else: + raise ValueError(f"Unknown solver type: {m.opt.solver}") + + +@wp.kernel +def solve_prev_grad_Mgrad( + # Data in: + efc_grad_in: wp.array2d(dtype=float), + efc_Mgrad_in: wp.array2d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_prev_grad_out: wp.array2d(dtype=float), + efc_prev_Mgrad_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + + if efc_done_in[worldid]: + return + + efc_prev_grad_out[worldid, dofid] = efc_grad_in[worldid, dofid] + efc_prev_Mgrad_out[worldid, dofid] = efc_Mgrad_in[worldid, dofid] + + +@wp.kernel +def solve_zero_beta_num_den( + # Data in: + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_beta_num_out: wp.array(dtype=float), + efc_beta_den_out: wp.array(dtype=float), +): + worldid = wp.tid() + + if efc_done_in[worldid]: + return + + efc_beta_num_out[worldid] = 0.0 + efc_beta_den_out[worldid] = 0.0 + + +@wp.kernel +def solve_beta_num_den( + # Data in: + efc_grad_in: wp.array2d(dtype=float), + efc_Mgrad_in: wp.array2d(dtype=float), + efc_prev_grad_in: wp.array2d(dtype=float), + efc_prev_Mgrad_in: wp.array2d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_beta_num_out: wp.array(dtype=float), + efc_beta_den_out: wp.array(dtype=float), +): + worldid, dofid = wp.tid() + + if efc_done_in[worldid]: + return + + prev_Mgrad = efc_prev_Mgrad_in[worldid][dofid] + wp.atomic_add( + efc_beta_num_out, + worldid, + efc_grad_in[worldid, dofid] * (efc_Mgrad_in[worldid, dofid] - prev_Mgrad), + ) + wp.atomic_add(efc_beta_den_out, worldid, efc_prev_grad_in[worldid, dofid] * prev_Mgrad) + + +@wp.kernel +def solve_beta( + # Data in: + efc_beta_num_in: wp.array(dtype=float), + efc_beta_den_in: wp.array(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_beta_out: wp.array(dtype=float), +): + worldid = wp.tid() + + if efc_done_in[worldid]: + return + + efc_beta_out[worldid] = wp.max(0.0, efc_beta_num_in[worldid] / wp.max(types.MJ_MINVAL, efc_beta_den_in[worldid])) + + +@wp.kernel +def solve_zero_search_dot( + # Data in: + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_search_dot_out: wp.array(dtype=float), +): + worldid = wp.tid() + + if efc_done_in[worldid]: + return + + efc_search_dot_out[worldid] = 0.0 + + +@wp.kernel +def solve_search_update( + # Model: + opt_solver: int, + # Data in: + efc_Mgrad_in: wp.array2d(dtype=float), + efc_search_in: wp.array2d(dtype=float), + efc_beta_in: wp.array(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_search_out: wp.array2d(dtype=float), + efc_search_dot_out: wp.array(dtype=float), +): + worldid, dofid = wp.tid() + + if efc_done_in[worldid]: + return + + search = -1.0 * efc_Mgrad_in[worldid, dofid] + + if opt_solver == wp.static(types.SolverType.CG.value): + search += efc_beta_in[worldid] * efc_search_in[worldid, dofid] + + efc_search_out[worldid, dofid] = search + wp.atomic_add(efc_search_dot_out, worldid, search * search) + + +@wp.kernel +def solve_done( + # Model: + nv: int, + opt_tolerance: wp.array(dtype=float), + opt_iterations: int, + stat_meaninertia: float, + # Data in: + efc_grad_dot_in: wp.array(dtype=float), + efc_cost_in: wp.array(dtype=float), + efc_prev_cost_in: wp.array(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + solver_niter_out: wp.array(dtype=int), + nsolving_out: wp.array(dtype=int), + efc_done_out: wp.array(dtype=bool), +): + worldid = wp.tid() + + if efc_done_in[worldid]: + return + + solver_niter_out[worldid] += 1 + tolerance = opt_tolerance[worldid] + + improvement = _rescale(nv, stat_meaninertia, efc_prev_cost_in[worldid] - efc_cost_in[worldid]) + gradient = _rescale(nv, stat_meaninertia, wp.math.sqrt(efc_grad_dot_in[worldid])) + done = (improvement < tolerance) or (gradient < tolerance) + if done or solver_niter_out[worldid] == opt_iterations: + # if the solver has converged or the maximum number of iterations has been reached then + # mark this world as done and remove it from the number of unconverged worlds + efc_done_out[worldid] = True + wp.atomic_add(nsolving_out, 0, -1) + + +@event_scope +def _solver_iteration( + m: types.Model, + d: types.Data, +): + _linesearch(m, d) + + if m.opt.solver == types.SolverType.CG: + wp.launch( + solve_prev_grad_Mgrad, + dim=(d.nworld, m.nv), + inputs=[d.efc.grad, d.efc.Mgrad, d.efc.done], + outputs=[d.efc.prev_grad, d.efc.prev_Mgrad], + ) + + _update_constraint(m, d) + _update_gradient(m, d) + + # polak-ribiere + if m.opt.solver == types.SolverType.CG: + wp.launch( + solve_zero_beta_num_den, + dim=(d.nworld), + inputs=[d.efc.done], + outputs=[d.efc.beta_num, d.efc.beta_den], + ) + + wp.launch( + solve_beta_num_den, + dim=(d.nworld, m.nv), + inputs=[d.efc.grad, d.efc.Mgrad, d.efc.prev_grad, d.efc.prev_Mgrad, d.efc.done], + outputs=[d.efc.beta_num, d.efc.beta_den], + ) + + wp.launch( + solve_beta, + dim=(d.nworld,), + inputs=[d.efc.beta_num, d.efc.beta_den, d.efc.done], + outputs=[d.efc.beta], + ) + + wp.launch(solve_zero_search_dot, dim=(d.nworld), inputs=[d.efc.done], outputs=[d.efc.search_dot]) + + wp.launch( + solve_search_update, + dim=(d.nworld, m.nv), + inputs=[m.opt.solver, d.efc.Mgrad, d.efc.search, d.efc.beta, d.efc.done], + outputs=[d.efc.search, d.efc.search_dot], + ) + + wp.launch( + solve_done, + dim=(d.nworld,), + inputs=[ + m.nv, + m.opt.tolerance, + m.opt.iterations, + m.stat.meaninertia, + d.efc.grad_dot, + d.efc.cost, + d.efc.prev_cost, + d.efc.done, + ], + outputs=[d.solver_niter, d.nsolving, d.efc.done], + ) + + +def create_context(m: types.Model, d: types.Data, grad: bool = True): + # initialize some efc arrays + wp.launch( + solve_init_efc, + dim=(d.nworld), + outputs=[d.solver_niter, d.efc.search_dot, d.efc.cost, d.efc.done], + ) + + # jaref = d.efc_J @ d.qacc - d.efc_aref + wp.launch( + solve_init_jaref, + dim=(d.nworld, d.njmax), + inputs=[m.nv, d.nefc, d.qacc, d.efc.J, d.efc.aref], + outputs=[d.efc.Jaref], + ) + + # Ma = qM @ qacc + support.mul_m(m, d, d.efc.Ma, d.qacc, d.efc.done) + + _update_constraint(m, d) + + if grad: + _update_gradient(m, d) + + +def _copy_acc(m: types.Model, d: types.Data): + wp.copy(d.qacc, d.qacc_smooth) + wp.copy(d.qacc_warmstart, d.qacc_smooth) + d.solver_niter.fill_(0) + + +@event_scope +def solve(m: types.Model, d: types.Data): + if d.njmax == 0: + _copy_acc(m, d) + else: + _solve(m, d) + + +def _solve(m: types.Model, d: types.Data): + """Finds forces that satisfy constraints.""" + # warmstart + wp.copy(d.qacc, d.qacc_warmstart) + + # create context + create_context(m, d, grad=True) + + # search = -Mgrad + wp.launch( + solve_init_search, + dim=(d.nworld, m.nv), + inputs=[d.efc.Mgrad], + outputs=[d.efc.search, d.efc.search_dot], + ) + + if m.opt.iterations != 0 and m.opt.graph_conditional: + # Note: the iteration kernel (indicated by while_body) is repeatedly launched + # as long as condition_iteration is not zero. + # condition_iteration is a warp array of size 1 and type int, it counts the number + # of worlds that are not converged, it becomes 0 when all worlds are converged. + # When the number of iterations reaches m.opt.iterations, solver_niter + # becomes zero and all worlds are marked as converged to avoid an infinite loop. + # note: we only launch the iteration kernel if everything is not done + d.nsolving.fill_(d.nworld) + wp.capture_while( + d.nsolving, + while_body=_solver_iteration, + m=m, + d=d, + ) + else: + # This branch is mostly for when JAX is used as it is currently not compatible + # with CUDA graph conditional. + # It should be removed when JAX becomes compatible. + for i in range(m.opt.iterations): + _solver_iteration(m, d) + + wp.copy(d.qacc_warmstart, d.qacc) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py new file mode 100644 index 00000000..11d043f6 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py @@ -0,0 +1,344 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Tests for solver functions.""" + +import mujoco +import numpy as np +import warp as wp +from absl.testing import absltest +from absl.testing import parameterized + +import mujoco_warp as mjwarp + +from mujoco.mjx.third_party.mujoco_warp._src import solver +from mujoco.mjx.third_party.mujoco_warp._src import test_util +from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType +from mujoco.mjx.third_party.mujoco_warp._src.types import SolverType + +# tolerance for difference between MuJoCo and MJWarp solver calculations - mostly +# due to float precision +_TOLERANCE = 5e-3 + + +def _assert_eq(a, b, name): + tol = _TOLERANCE * 10 # avoid test noise + err_msg = f"mismatch: {name}" + np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol) + + +class SolverTest(parameterized.TestCase): + @parameterized.product(cone=tuple(ConeType), solver_=tuple(SolverType)) + def test_cost(self, cone, solver_): + """Tests cost function is correct.""" + for keyframe in range(3): + mjm, mjd, m, d = test_util.fixture( + "constraints.xml", + keyframe=keyframe, + cone=cone, + solver=solver_, + iterations=0, + ) + + def cost(qacc): + jaref = np.zeros(mjd.nefc, dtype=float) + cost = np.zeros(1) + mujoco.mj_mulJacVec(mjm, mjd, jaref, qacc) + mujoco.mj_constraintUpdate(mjm, mjd, jaref - mjd.efc_aref, cost, 0) + return cost + + mj_cost = cost(mjd.qacc) + + # solve with 0 iterations just initializes constraints and costs and then exits + mjwarp.solve(m, d) + + mjwarp_cost = d.efc.cost.numpy()[0] - d.efc.gauss.numpy()[0] + + _assert_eq(mjwarp_cost, mj_cost, name="cost") + + @parameterized.parameters( + (ConeType.PYRAMIDAL, SolverType.CG, 5, 5, False, False), + (ConeType.ELLIPTIC, SolverType.CG, 5, 5, False, False), + (ConeType.PYRAMIDAL, SolverType.NEWTON, 2, 4, False, False), + (ConeType.ELLIPTIC, SolverType.NEWTON, 2, 5, False, False), + (ConeType.PYRAMIDAL, SolverType.NEWTON, 2, 4, True, True), + (ConeType.ELLIPTIC, SolverType.NEWTON, 3, 16, True, True), + ) + def test_solve(self, cone, solver_, iterations, ls_iterations, sparse, ls_parallel): + """Tests solve.""" + for keyframe in range(3): + mjm, mjd, m, d = test_util.fixture( + "constraints.xml", + keyframe=keyframe, + sparse=sparse, + cone=cone, + solver=solver_, + iterations=iterations, + ls_iterations=ls_iterations, + ls_parallel=ls_parallel, + ) + + qacc_warmstart = mjd.qacc_warmstart.copy() + mujoco.mj_forward(mjm, mjd) + mjd.qacc_warmstart = qacc_warmstart + + d.qacc.zero_() + d.qfrc_constraint.zero_() + d.efc.force.zero_() + + if solver_ == mujoco.mjtSolver.mjSOL_CG: + mjwarp.factor_m(m, d) + mjwarp.solve(m, d) + + def cost(qacc): + jaref = np.zeros(mjd.nefc, dtype=float) + cost = np.zeros(1) + mujoco.mj_mulJacVec(mjm, mjd, jaref, qacc) + mujoco.mj_constraintUpdate(mjm, mjd, jaref - mjd.efc_aref, cost, 0) + return cost + + mj_cost = cost(mjd.qacc) + mjwarp_cost = cost(d.qacc.numpy()[0]) + self.assertLessEqual(mjwarp_cost, mj_cost * 1.025) + + if m.opt.solver == mujoco.mjtSolver.mjSOL_NEWTON: + _assert_eq(d.qacc.numpy()[0], mjd.qacc, "qacc") + _assert_eq(d.qfrc_constraint.numpy()[0], mjd.qfrc_constraint, "qfrc_constraint") + _assert_eq(d.efc.force.numpy()[0, : mjd.nefc], mjd.efc_force, "efc_force") + + @parameterized.parameters( + (ConeType.PYRAMIDAL, SolverType.CG, 25, 5), + (ConeType.PYRAMIDAL, SolverType.NEWTON, 2, 4), + ) + def test_solve_batch(self, cone, solver_, iterations, ls_iterations): + """Tests solve (batch).""" + + mjm0, mjd0, _, _ = test_util.fixture( + "humanoid/humanoid.xml", + keyframe=0, + sparse=False, + cone=cone, + solver=solver_, + iterations=iterations, + ls_iterations=ls_iterations, + ) + qacc_warmstart0 = mjd0.qacc_warmstart.copy() + mujoco.mj_forward(mjm0, mjd0) + mjd0.qacc_warmstart = qacc_warmstart0 + + mjm1, mjd1, _, _ = test_util.fixture( + "humanoid/humanoid.xml", + keyframe=2, + sparse=False, + cone=cone, + solver=solver_, + iterations=iterations, + ls_iterations=ls_iterations, + ) + qacc_warmstart1 = mjd1.qacc_warmstart.copy() + mujoco.mj_forward(mjm1, mjd1) + mjd1.qacc_warmstart = qacc_warmstart1 + + mjm2, mjd2, _, _ = test_util.fixture( + "humanoid/humanoid.xml", + keyframe=1, + sparse=False, + cone=cone, + solver=solver_, + iterations=iterations, + ls_iterations=ls_iterations, + ) + qacc_warmstart2 = mjd2.qacc_warmstart.copy() + mujoco.mj_forward(mjm2, mjd2) + mjd2.qacc_warmstart = qacc_warmstart2 + + nefc_active = mjd0.nefc + mjd1.nefc + mjd2.nefc + ne_active = mjd0.ne + mjd1.ne + mjd2.ne + + mjm, mjd, m, _ = test_util.fixture( + "humanoid/humanoid.xml", + sparse=False, + cone=cone, + solver=solver_, + iterations=iterations, + ls_iterations=ls_iterations, + ) + d = mjwarp.put_data(mjm, mjd, nworld=3, njmax=2 * nefc_active) + + d.nefc = wp.array([nefc_active, nefc_active, nefc_active], dtype=wp.int32, ndim=1) + d.ne = wp.array([ne_active, ne_active, ne_active], dtype=wp.int32, ndim=1) + + qacc_warmstart = np.vstack( + [ + np.expand_dims(qacc_warmstart0, axis=0), + np.expand_dims(qacc_warmstart1, axis=0), + np.expand_dims(qacc_warmstart2, axis=0), + ] + ) + + qM0 = np.zeros((mjm0.nv, mjm0.nv)) + mujoco.mj_fullM(mjm0, qM0, mjd0.qM) + qM1 = np.zeros((mjm1.nv, mjm1.nv)) + mujoco.mj_fullM(mjm1, qM1, mjd1.qM) + qM2 = np.zeros((mjm2.nv, mjm2.nv)) + mujoco.mj_fullM(mjm2, qM2, mjd2.qM) + + qM = np.vstack( + [ + np.expand_dims(qM0, axis=0), + np.expand_dims(qM1, axis=0), + np.expand_dims(qM2, axis=0), + ] + ) + qacc_smooth = np.vstack( + [ + np.expand_dims(mjd0.qacc_smooth, axis=0), + np.expand_dims(mjd1.qacc_smooth, axis=0), + np.expand_dims(mjd2.qacc_smooth, axis=0), + ] + ) + qfrc_smooth = np.vstack( + [ + np.expand_dims(mjd0.qfrc_smooth, axis=0), + np.expand_dims(mjd1.qfrc_smooth, axis=0), + np.expand_dims(mjd2.qfrc_smooth, axis=0), + ] + ) + + # Reshape the Jacobians + efc_J0 = mjd0.efc_J.reshape((mjd0.nefc, mjm0.nv)) + efc_J1 = mjd1.efc_J.reshape((mjd1.nefc, mjm1.nv)) + efc_J2 = mjd2.efc_J.reshape((mjd2.nefc, mjm2.nv)) + + efc_J_fill = np.zeros((3, d.njmax, m.nv)) + efc_J_fill[0, : mjd0.nefc, :] = efc_J0 + efc_J_fill[1, : mjd1.nefc, :] = efc_J1 + efc_J_fill[2, : mjd2.nefc, :] = efc_J2 + + # Similarly for D and aref values + efc_D0 = mjd0.efc_D[: mjd0.nefc] + efc_D1 = mjd1.efc_D[: mjd1.nefc] + efc_D2 = mjd2.efc_D[: mjd2.nefc] + + efc_D_fill = np.zeros((3, d.njmax)) + efc_D_fill[0, : mjd0.nefc] = efc_D0 + efc_D_fill[1, : mjd1.nefc] = efc_D1 + efc_D_fill[2, : mjd2.nefc] = efc_D2 + + efc_aref0 = mjd0.efc_aref[: mjd0.nefc] + efc_aref1 = mjd1.efc_aref[: mjd1.nefc] + efc_aref2 = mjd2.efc_aref[: mjd2.nefc] + + efc_aref_fill = np.zeros((3, d.njmax)) + efc_aref_fill[0, : mjd0.nefc] = efc_aref0 + efc_aref_fill[1, : mjd1.nefc] = efc_aref1 + efc_aref_fill[2, : mjd2.nefc] = efc_aref2 + + d.qacc_warmstart = wp.from_numpy(qacc_warmstart, dtype=wp.float32) + d.qM = wp.from_numpy(qM, dtype=wp.float32) + d.qacc_smooth = wp.from_numpy(qacc_smooth, dtype=wp.float32) + d.qfrc_smooth = wp.from_numpy(qfrc_smooth, dtype=wp.float32) + d.efc.J = wp.from_numpy(efc_J_fill, dtype=wp.float32) + d.efc.D = wp.from_numpy(efc_D_fill, dtype=wp.float32) + d.efc.aref = wp.from_numpy(efc_aref_fill, dtype=wp.float32) + + if solver_ == SolverType.CG: + m0 = mjwarp.put_model(mjm0) + d0 = mjwarp.put_data(mjm0, mjd0) + mjwarp.factor_m(m0, d0) + qLD0 = d0.qLD.numpy() + + m1 = mjwarp.put_model(mjm1) + d1 = mjwarp.put_data(mjm1, mjd1) + mjwarp.factor_m(m1, d1) + qLD1 = d1.qLD.numpy() + + m2 = mjwarp.put_model(mjm2) + d2 = mjwarp.put_data(mjm2, mjd2) + mjwarp.factor_m(m2, d2) + qLD2 = d2.qLD.numpy() + + qLD = np.vstack([qLD0, qLD1, qLD2]) + d.qLD = wp.from_numpy(qLD, dtype=wp.float32) + + d.qacc.zero_() + d.qfrc_constraint.zero_() + d.efc.force.zero_() + solver.solve(m, d) + + def cost(m, d, qacc): + jaref = np.zeros(d.nefc, dtype=float) + cost = np.zeros(1) + mujoco.mj_mulJacVec(m, d, jaref, qacc) + mujoco.mj_constraintUpdate(m, d, jaref - d.efc_aref, cost, 0) + return cost + + mj_cost0 = cost(mjm0, mjd0, mjd0.qacc) + mjwarp_cost0 = cost(mjm0, mjd0, d.qacc.numpy()[0]) + self.assertLessEqual(mjwarp_cost0, mj_cost0 * 1.025) + + mj_cost1 = cost(mjm1, mjd1, mjd1.qacc) + mjwarp_cost1 = cost(mjm1, mjd1, d.qacc.numpy()[1]) + self.assertLessEqual(mjwarp_cost1, mj_cost1 * 1.025) + + mj_cost2 = cost(mjm2, mjd2, mjd2.qacc) + mjwarp_cost2 = cost(mjm2, mjd2, d.qacc.numpy()[2]) + self.assertLessEqual(mjwarp_cost2, mj_cost2 * 1.025) + + if m.opt.solver == SolverType.NEWTON: + _assert_eq(d.qacc.numpy()[0], mjd0.qacc, "qacc0") + _assert_eq(d.qacc.numpy()[1], mjd1.qacc, "qacc1") + _assert_eq(d.qacc.numpy()[2], mjd2.qacc, "qacc2") + + _assert_eq(d.qfrc_constraint.numpy()[0], mjd0.qfrc_constraint, "qfrc_constraint0") + _assert_eq(d.qfrc_constraint.numpy()[1], mjd1.qfrc_constraint, "qfrc_constraint1") + _assert_eq(d.qfrc_constraint.numpy()[2], mjd2.qfrc_constraint, "qfrc_constraint2") + + # Get world 0 forces - equality constraints at start, inequality constraints later + nieq0 = mjd0.nefc - mjd0.ne + nieq1 = mjd1.nefc - mjd1.ne + nieq2 = mjd2.nefc - mjd2.ne + world0_eq_forces = d.efc.force.numpy()[0, : mjd0.ne] + world0_ineq_forces = d.efc.force.numpy()[0, ne_active : ne_active + nieq0] + world0_forces = np.concatenate([world0_eq_forces, world0_ineq_forces]) + _assert_eq(world0_forces, mjd0.efc_force, "efc_force0") + + # Get world 1 forces + world1_eq_forces = d.efc.force.numpy()[1, : mjd1.ne] + world1_ineq_forces = d.efc.force.numpy()[1, ne_active : ne_active + nieq1] + world1_forces = np.concatenate([world1_eq_forces, world1_ineq_forces]) + _assert_eq(world1_forces, mjd1.efc_force, "efc_force1") + + # Get world 2 forces + world2_eq_forces = d.efc.force.numpy()[2, : mjd2.ne] + world2_ineq_forces = d.efc.force.numpy()[2, ne_active : ne_active + nieq2] + world2_forces = np.concatenate([world2_eq_forces, world2_ineq_forces]) + _assert_eq(world2_forces, mjd2.efc_force, "efc_force2") + + def test_frictionloss(self): + """Tests solver with frictionloss.""" + for keyframe in range(3): + _, mjd, m, d = test_util.fixture("constraints.xml", keyframe=keyframe) + mjwarp.solve(m, d) + + _assert_eq(d.nf.numpy()[0], mjd.nf, "nf") + _assert_eq(d.qacc.numpy()[0], mjd.qacc, "qacc") + _assert_eq(d.qfrc_constraint.numpy()[0], mjd.qfrc_constraint, "qfrc_constraint") + _assert_eq(d.efc.force.numpy()[0, : mjd.nefc], mjd.efc_force, "efc_force") + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py new file mode 100644 index 00000000..b59ffadd --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py @@ -0,0 +1,554 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +from typing import Tuple + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src.math import motion_cross +from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import JointType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import TileSet +from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel + +wp.set_module_options({"enable_backward": False}) + + +@wp.kernel +def mul_m_sparse_diag( + # Model: + dof_Madr: wp.array(dtype=int), + # Data in: + qM_in: wp.array3d(dtype=float), + # In: + vec: wp.array2d(dtype=float), + skip: wp.array(dtype=bool), + # Out: + res: wp.array2d(dtype=float), +): + """Diagonal update for sparse matmul.""" + worldid, dofid = wp.tid() + + if skip[worldid]: + return + + res[worldid, dofid] = qM_in[worldid, 0, dof_Madr[dofid]] * vec[worldid, dofid] + + +@wp.kernel +def mul_m_sparse_ij( + # Model: + qM_mulm_i: wp.array(dtype=int), + qM_mulm_j: wp.array(dtype=int), + qM_madr_ij: wp.array(dtype=int), + # Data in: + qM_in: wp.array3d(dtype=float), + # In: + vec: wp.array2d(dtype=float), + skip: wp.array(dtype=bool), + # Out: + res: wp.array2d(dtype=float), +): + """Off-diagonal update for sparse matmul.""" + worldid, elementid = wp.tid() + + if skip[worldid]: + return + + i = qM_mulm_i[elementid] + j = qM_mulm_j[elementid] + madr_ij = qM_madr_ij[elementid] + + qM_ij = qM_in[worldid, 0, madr_ij] + + wp.atomic_add(res[worldid], i, qM_ij * vec[worldid, j]) + wp.atomic_add(res[worldid], j, qM_ij * vec[worldid, i]) + + +@cache_kernel +def mul_m_dense(tile: TileSet): + """Returns a matmul kernel for some tile size""" + + @nested_kernel + def kernel( + # Data In: + qM_in: wp.array3d(dtype=float), + # In: + adr: wp.array(dtype=int), + vec: wp.array3d(dtype=float), + skip: wp.array(dtype=bool), + # Out: + res: wp.array3d(dtype=float), + ): + worldid, nodeid = wp.tid() + TILE_SIZE = wp.static(tile.size) + + if skip[worldid]: + return + + dofid = adr[nodeid] + qM_tile = wp.tile_load(qM_in[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid)) + vec_tile = wp.tile_load(vec[worldid], shape=(TILE_SIZE, 1), offset=(dofid, 0)) + res_tile = wp.tile_matmul(qM_tile, vec_tile) + wp.tile_store(res[worldid], res_tile, offset=(dofid, 0)) + + return kernel + + +@event_scope +def mul_m( + m: Model, + d: Data, + res: wp.array2d(dtype=float), + vec: wp.array2d(dtype=float), + skip: wp.array(dtype=bool), + M: wp.array3d(dtype=float) = None, +): + """Multiply vectors by inertia matrix. + + Args: + m (Model): The model containing kinematic and dynamic information (device). + d (Data): The data object containing the current state and output arrays (device). + res (wp.array2d(dtype=float)): Result: qM @ vec. + vec (wp.array2d(dtype=float)): Input vector to multiply by qM. + skip (wp.array(dtype=flooat)): Skip output. + M (wp.array3d(dtype=float), optional): Input matrix: M @ vec. + """ + + if M is None: + M = d.qM + + if m.opt.is_sparse: + wp.launch( + mul_m_sparse_diag, + dim=(d.nworld, m.nv), + inputs=[m.dof_Madr, M, vec, skip], + outputs=[res], + ) + + wp.launch( + mul_m_sparse_ij, + dim=(d.nworld, m.qM_madr_ij.size), + inputs=[m.qM_mulm_i, m.qM_mulm_j, m.qM_madr_ij, M, vec, skip], + outputs=[res], + ) + + else: + for tile in m.qM_tiles: + wp.launch_tiled( + mul_m_dense(tile), + dim=(d.nworld, tile.adr.size), + inputs=[ + M, + tile.adr, + # note reshape: tile_matmul expects 2d input + vec.reshape(vec.shape + (1,)), + skip, + ], + outputs=[res.reshape(res.shape + (1,))], + block_dim=m.block_dim.mul_m_dense, + ) + + +@wp.kernel +def xfrc_accumulate_kernel( + # Model: + nbody: int, + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + dof_bodyid: wp.array(dtype=int), + # Data in: + xfrc_applied_in: wp.array2d(dtype=wp.spatial_vector), + xipos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + # Out: + out: wp.array2d(dtype=float), +): + """Accumulate applied forces on the subtree of a dof.""" + worldid, dofid = wp.tid() + cdof = cdof_in[worldid, dofid] + rotational_cdof = wp.spatial_top(cdof) + jac = wp.spatial_vector(cdof[3], cdof[4], cdof[5], cdof[0], cdof[1], cdof[2]) + + bodyid = dof_bodyid[dofid] + accumul = float(0.0) + + for child in range(bodyid, nbody): + # any body that is in the subtree of dof_bodyid is part of the jacobian + parentid = child + while parentid != 0 and parentid != bodyid: + parentid = body_parentid[parentid] + if parentid == 0: + continue # body is not part of the subtree + offset = xipos_in[worldid, child] - subtree_com_in[worldid, body_rootid[child]] + cross_term = wp.cross(rotational_cdof, offset) + xfrc_applied = xfrc_applied_in[worldid, child] + accumul += wp.dot(jac, xfrc_applied) + wp.dot(cross_term, wp.spatial_top(xfrc_applied)) + + out[worldid, dofid] += accumul + + +@wp.kernel +def _apply_ft( + # Model: + nbody: int, + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + dof_bodyid: wp.array(dtype=int), + # Data in: + xipos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + # In: + ft_in: wp.array2d(dtype=wp.spatial_vector), + flg_add: bool, + # Out: + qfrc_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + cdof = cdof_in[worldid, dofid] + rotational_cdof = wp.vec3(cdof[0], cdof[1], cdof[2]) + jac = wp.spatial_vector(cdof[3], cdof[4], cdof[5], cdof[0], cdof[1], cdof[2]) + + dofbodyid = dof_bodyid[dofid] + accumul = float(0.0) + + for bodyid in range(dofbodyid, nbody): + # any body that is in the subtree of dofbodyid is part of the jacobian + parentid = bodyid + while parentid != 0 and parentid != dofbodyid: + parentid = body_parentid[parentid] + if parentid == 0: + continue # body is not part of the subtree + offset = xipos_in[worldid, bodyid] - subtree_com_in[worldid, body_rootid[bodyid]] + cross_term = wp.cross(rotational_cdof, offset) + ft_body = ft_in[worldid, bodyid] + accumul += wp.dot(jac, ft_body) + wp.dot(cross_term, wp.spatial_top(ft_body)) + + if flg_add: + qfrc_out[worldid, dofid] += accumul + else: + qfrc_out[worldid, dofid] = accumul + + +def apply_ft(m: Model, d: Data, ft: wp.array2d(dtype=wp.spatial_vector), qfrc: wp.array2d(dtype=float), flg_add: bool): + wp.launch( + kernel=_apply_ft, + dim=(d.nworld, m.nv), + inputs=[m.nbody, m.body_parentid, m.body_rootid, m.dof_bodyid, d.xipos, d.subtree_com, d.cdof, ft, flg_add], + outputs=[qfrc], + ) + + +@event_scope +def xfrc_accumulate(m: Model, d: Data, qfrc: wp.array2d(dtype=float)): + """ + Map applied forces at each body via Jacobians to dof space and accumulate. + + Args: + m (Model): The model containing kinematic and dynamic information (device). + d (Data): The data object containing the current state and output arrays (device). + qfrc (wp.array2d(dtype=float)): Total applied force mapped to dof space. + """ + apply_ft(m, d, d.xfrc_applied, qfrc, True) + + +@wp.func +def all_same(v0: wp.vec3, v1: wp.vec3) -> wp.bool: + dx = abs(v0[0] - v1[0]) + dy = abs(v0[1] - v1[1]) + dz = abs(v0[2] - v1[2]) + + return ( + (dx <= 1.0e-9 or dx <= max(abs(v0[0]), abs(v1[0])) * 1.0e-9) + and (dy <= 1.0e-9 or dy <= max(abs(v0[1]), abs(v1[1])) * 1.0e-9) + and (dz <= 1.0e-9 or dz <= max(abs(v0[2]), abs(v1[2])) * 1.0e-9) + ) + + +@wp.func +def any_different(v0: wp.vec3, v1: wp.vec3) -> wp.bool: + dx = abs(v0[0] - v1[0]) + dy = abs(v0[1] - v1[1]) + dz = abs(v0[2] - v1[2]) + + return ( + (dx > 1.0e-9 and dx > max(abs(v0[0]), abs(v1[0])) * 1.0e-9) + or (dy > 1.0e-9 and dy > max(abs(v0[1]), abs(v1[1])) * 1.0e-9) + or (dz > 1.0e-9 and dz > max(abs(v0[2]), abs(v1[2])) * 1.0e-9) + ) + + +@wp.func +def _decode_pyramid(pyramid: wp.array(dtype=float), efc_address: int, mu: vec5, condim: int) -> wp.spatial_vector: + """Converts pyramid representation to contact force.""" + force = wp.spatial_vector() + + if condim == 1: + force[0] = pyramid[efc_address] + return force + + force[0] = float(0.0) + for i in range(condim - 1): + dir1 = pyramid[2 * i + efc_address] + dir2 = pyramid[2 * i + efc_address + 1] + force[0] += dir1 + dir2 + force[i + 1] = (dir1 - dir2) * mu[i] + + return force + + +@wp.func +def contact_force_fn( + # Model: + opt_cone: int, + # Data in: + ncon_in: wp.array(dtype=int), + contact_frame_in: wp.array(dtype=wp.mat33), + contact_friction_in: wp.array(dtype=vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + efc_force_in: wp.array2d(dtype=float), + # In: + worldid: int, + contact_id: int, + to_world_frame: bool, +) -> wp.spatial_vector: + """Extract 6D force:torque for one contact, in contact frame by default.""" + force = wp.spatial_vector(0.0, 0.0, 0.0, 0.0, 0.0, 0.0) + condim = contact_dim_in[contact_id] + efc_address = contact_efc_address_in[contact_id, 0] + + if contact_id >= 0 and contact_id <= ncon_in[0] and efc_address >= 0: + if opt_cone == int(ConeType.PYRAMIDAL.value): + force = _decode_pyramid( + efc_force_in[worldid], + efc_address, + contact_friction_in[contact_id], + condim, + ) + else: + for i in range(condim): + force[i] = efc_force_in[worldid, contact_efc_address_in[contact_id, i]] + + if to_world_frame: + # Transform both top and bottom parts of spatial vector by the full contact frame matrix + t = wp.spatial_top(force) @ contact_frame_in[contact_id] + b = wp.spatial_bottom(force) @ contact_frame_in[contact_id] + force = wp.spatial_vector(t, b) + + return force + + +@wp.kernel +def contact_force_kernel( + # Model: + opt_cone: int, + # Data in: + ncon_in: wp.array(dtype=int), + contact_frame_in: wp.array(dtype=wp.mat33), + contact_friction_in: wp.array(dtype=vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_force_in: wp.array2d(dtype=float), + # In: + contact_ids: wp.array(dtype=int), + to_world_frame: bool, + # Out: + out: wp.array(dtype=wp.spatial_vector), +): + tid = wp.tid() + + contactid = contact_ids[tid] + + if contactid >= ncon_in[0]: + return + + worldid = contact_worldid_in[contactid] + + out[tid] = contact_force_fn( + opt_cone, + ncon_in, + contact_frame_in, + contact_friction_in, + contact_dim_in, + contact_efc_address_in, + efc_force_in, + worldid, + contactid, + to_world_frame, + ) + + +def contact_force( + m: Model, + d: Data, + contact_ids: wp.array(dtype=int), + to_world_frame: bool, + force: wp.array(dtype=wp.spatial_vector), +): + """ + Compute forces for contacts in Data. + + Args: + m (Model): The model containing kinematic and dynamic information (device). + d (Data): The data object containing the current state and output arrays (device). + contact_ids (wp.array(dtype=int)): IDs for each contact. + to_world_frame (bool): If True, map force from contact to world frame. + force (wp.array(dtype=wp.spatial_vector)): Contact forces. + """ + wp.launch( + contact_force_kernel, + dim=(contact_ids.size,), + inputs=[ + m.opt.cone, + d.ncon, + d.contact.frame, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.contact.worldid, + d.efc.force, + contact_ids, + to_world_frame, + ], + outputs=[force], + ) + + +@wp.func +def transform_force(force: wp.vec3, torque: wp.vec3, offset: wp.vec3) -> wp.spatial_vector: + return wp.spatial_vector(torque - wp.cross(offset, force), force) + + +@wp.func +def transform_force(frc: wp.spatial_vector, offset: wp.vec3) -> wp.spatial_vector: + force = wp.spatial_top(frc) + torque = wp.spatial_bottom(frc) + return transform_force(force, torque, offset) + + +@wp.func +def jac( + # Model: + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + dof_bodyid: wp.array(dtype=int), + # Data in: + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + # In: + point: wp.vec3, + bodyid: int, + dofid: int, + worldid: int, +) -> Tuple[wp.vec3, wp.vec3]: + dof_bodyid_ = dof_bodyid[dofid] + in_tree = int(dof_bodyid_ == 0) + parentid = bodyid + while parentid != 0: + if parentid == dof_bodyid_: + in_tree = 1 + break + parentid = body_parentid[parentid] + + if not in_tree: + return wp.vec3(0.0), wp.vec3(0.0) + + offset = point - wp.vec3(subtree_com_in[worldid, body_rootid[bodyid]]) + + cdof = cdof_in[worldid, dofid] + cdof_ang = wp.spatial_top(cdof) + cdof_lin = wp.spatial_bottom(cdof) + + jacp = cdof_lin + wp.cross(cdof_ang, offset) + jacr = cdof_ang + + return jacp, jacr + + +@wp.func +def jac_dot( + # Model: + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + jnt_type: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + dof_bodyid: wp.array(dtype=int), + dof_jntid: wp.array(dtype=int), + # Data in: + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + cdof_dot_in: wp.array2d(dtype=wp.spatial_vector), + # In: + point: wp.vec3, + bodyid: int, + dofid: int, + worldid: int, +) -> Tuple[wp.vec3, wp.vec3]: + dof_bodyid_ = dof_bodyid[dofid] + in_tree = int(dof_bodyid_ == 0) + parentid = bodyid + while parentid != 0: + if parentid == dof_bodyid_: + in_tree = 1 + break + parentid = body_parentid[parentid] + + if not in_tree: + return wp.vec3(0.0), wp.vec3(0.0) + + com = subtree_com_in[worldid, body_rootid[bodyid]] + offset = point - com + + # transform spatial + cvel = cvel_in[worldid, bodyid] + pvel_lin = wp.spatial_bottom(cvel) - wp.cross(offset, wp.spatial_top(cvel)) + + cdof = cdof_in[worldid, dofid] + cdof_dot = cdof_dot_in[worldid, dofid] + + # check for quaternion + dofjntid = dof_jntid[dofid] + jnttype = jnt_type[dofjntid] + jntdofadr = jnt_dofadr[dofjntid] + + if (jnttype == int(JointType.BALL.value)) or ((jnttype == int(JointType.FREE.value)) and dofid >= jntdofadr + 3): + # compute cdof_dot for quaternion (use current body cvel) + cvel = cvel_in[worldid, dof_bodyid[dofid]] + cdof_dot = motion_cross(cvel, cdof) + + cdof_dot_ang = wp.spatial_top(cdof_dot) + cdof_dot_lin = wp.spatial_bottom(cdof_dot) + + # construct translational Jacobian (correct for rotation) + # first correction term, account for varying cdof + correction1 = wp.cross(cdof_dot_ang, offset) + + # second correction term, account for point translational velocity + correction2 = wp.cross(wp.spatial_top(cdof), pvel_lin) + + jacp = cdof_dot_lin + correction1 + correction2 + jacr = cdof_dot_ang + + return jacp, jacr diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support_test.py new file mode 100644 index 00000000..7c639e3b --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support_test.py @@ -0,0 +1,123 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Tests for support functions.""" + +import mujoco +import numpy as np +import warp as wp +from absl.testing import absltest +from absl.testing import parameterized + +import mujoco_warp as mjwarp + +from mujoco.mjx.third_party.mujoco_warp._src import test_util +from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType + +# tolerance for difference between MuJoCo and MJWarp support calculations - mostly +# due to float precision +_TOLERANCE = 5e-5 + + +def _assert_eq(a, b, name): + tol = _TOLERANCE * 10 # avoid test noise + err_msg = f"mismatch: {name}" + np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol) + + +class SupportTest(parameterized.TestCase): + @parameterized.parameters(True, False) + def test_mul_m(self, sparse): + """Tests mul_m.""" + mjm, mjd, m, d = test_util.fixture("pendula.xml", sparse=sparse) + + mj_res = np.zeros(mjm.nv) + mj_vec = np.random.uniform(low=-1.0, high=1.0, size=mjm.nv) + mujoco.mj_mulM(mjm, mjd, mj_res, mj_vec) + + res = wp.zeros((1, mjm.nv), dtype=wp.float32) + vec = wp.from_numpy(np.expand_dims(mj_vec, axis=0), dtype=wp.float32) + skip = wp.zeros((d.nworld), dtype=bool) + mjwarp.mul_m(m, d, res, vec, skip) + + _assert_eq(res.numpy()[0], mj_res, f"mul_m ({'sparse' if sparse else 'dense'})") + + def test_xfrc_accumulated(self): + """Tests that xfrc_accumulate output matches mj_xfrcAccumulate.""" + mjm, mjd, m, d = test_util.fixture("pendula.xml") + xfrc = np.random.randn(*d.xfrc_applied.numpy().shape) + d.xfrc_applied = wp.from_numpy(xfrc, dtype=wp.spatial_vector) + qfrc = wp.zeros((1, mjm.nv), dtype=wp.float32) + mjwarp.xfrc_accumulate(m, d, qfrc) + + qfrc_expected = np.zeros(m.nv) + xfrc = xfrc[0] + for i in range(1, m.nbody): + mujoco.mj_applyFT(mjm, mjd, xfrc[i, :3], xfrc[i, 3:], mjd.xipos[i], i, qfrc_expected) + np.testing.assert_almost_equal(qfrc.numpy()[0], qfrc_expected, 6) + + @parameterized.parameters( + (ConeType.PYRAMIDAL, 1, False), + (ConeType.PYRAMIDAL, 3, False), + (ConeType.PYRAMIDAL, 4, False), + (ConeType.PYRAMIDAL, 6, False), + (ConeType.PYRAMIDAL, 1, True), + (ConeType.PYRAMIDAL, 3, True), + (ConeType.PYRAMIDAL, 4, True), + (ConeType.PYRAMIDAL, 6, True), + (ConeType.ELLIPTIC, 1, False), + (ConeType.ELLIPTIC, 3, False), + (ConeType.ELLIPTIC, 4, False), + (ConeType.ELLIPTIC, 6, False), + (ConeType.ELLIPTIC, 1, True), + (ConeType.ELLIPTIC, 3, True), + (ConeType.ELLIPTIC, 4, True), + (ConeType.ELLIPTIC, 6, True), + ) + def test_contact_force(self, cone, condim, to_world_frame): + _CONTACT = f""" + + + + + + + + + + + + + """ + mjm, mjd, m, d = test_util.fixture(xml=_CONTACT, cone=cone, keyframe=0) + + mj_force = np.zeros(6, dtype=float) + mujoco.mj_contactForce(mjm, mjd, 0, mj_force) + + contact_ids = wp.zeros(1, dtype=int) + force = wp.zeros(1, dtype=wp.spatial_vector) + + mjwarp.contact_force(m, d, contact_ids, to_world_frame, force) + + if to_world_frame: + frame = mjd.contact.frame[0].reshape((3, 3)) + mj_force = np.concatenate([frame.T @ mj_force[:3], frame.T @ mj_force[3:]]) + + _assert_eq(force.numpy()[0], mj_force, "contact force") + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/test_util.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/test_util.py new file mode 100644 index 00000000..4d9ae7b1 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/test_util.py @@ -0,0 +1,267 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Utilities for testing.""" + +import time +from typing import Callable, Optional, Tuple + +import mujoco +import numpy as np +import warp as wp +from etils import epath + +from mujoco.mjx.third_party.mujoco_warp._src import io +from mujoco.mjx.third_party.mujoco_warp._src import warp_util +from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import SolverType +from mujoco.mjx.third_party.mujoco_warp._src.util_misc import halton + + +def fixture( + fname: Optional[str] = None, + xml: Optional[str] = None, + keyframe: int = -1, + actuation: bool = True, + contact: bool = True, + constraint: bool = True, + equality: bool = True, + passive: bool = True, + gravity: bool = True, + clampctrl: bool = True, + filterparent: bool = True, + qpos0: bool = False, + kick: bool = False, + energy: bool = False, + eulerdamp: Optional[bool] = None, + cone: Optional[ConeType] = None, + integrator: Optional[IntegratorType] = None, + solver: Optional[SolverType] = None, + iterations: Optional[int] = None, + ls_iterations: Optional[int] = None, + ls_parallel: Optional[bool] = None, + sparse: Optional[bool] = None, + disableflags: Optional[int] = None, + enableflags: Optional[int] = None, + applied: bool = False, + nstep: int = 3, + seed: int = 42, + nworld: int = None, + nconmax: int = None, + njmax: int = None, +): + np.random.seed(seed) + if fname is not None: + path = epath.resource_path("mjx") / "third_party/mujoco_warp" / "test_data" / fname + mjm = mujoco.MjModel.from_xml_path(path.as_posix()) + elif xml is not None: + mjm = mujoco.MjModel.from_xml_string(xml) + else: + raise ValueError("either fname or xml must be provided") + + if not actuation: + mjm.opt.disableflags |= DisableBit.ACTUATION + if not contact: + mjm.opt.disableflags |= DisableBit.CONTACT + if not constraint: + mjm.opt.disableflags |= DisableBit.CONSTRAINT + if not equality: + mjm.opt.disableflags |= DisableBit.EQUALITY + if not passive: + mjm.opt.disableflags |= DisableBit.PASSIVE + if not gravity: + mjm.opt.disableflags |= DisableBit.GRAVITY + if not clampctrl: + mjm.opt.disableflags |= DisableBit.CLAMPCTRL + if not eulerdamp: + mjm.opt.disableflags |= DisableBit.EULERDAMP + if not filterparent: + mjm.opt.disableflags |= DisableBit.FILTERPARENT + + if energy: + mjm.opt.enableflags |= EnableBit.ENERGY + + if cone is not None: + mjm.opt.cone = cone + if integrator is not None: + mjm.opt.integrator = integrator + if disableflags is not None: + mjm.opt.disableflags |= disableflags + if enableflags is not None: + mjm.opt.enableflags |= enableflags + if solver is not None: + mjm.opt.solver = solver + if iterations is not None: + mjm.opt.iterations = iterations + if ls_iterations is not None: + mjm.opt.ls_iterations = ls_iterations + if sparse is not None: + if sparse: + mjm.opt.jacobian = mujoco.mjtJacobian.mjJAC_SPARSE + else: + mjm.opt.jacobian = mujoco.mjtJacobian.mjJAC_DENSE + + mjd = mujoco.MjData(mjm) + if keyframe > -1: + mujoco.mj_resetDataKeyframe(mjm, mjd, keyframe) + elif qpos0: + mjd.qpos[:] = mjm.qpos0 + else: + # set random qpos, underlying code should gracefully handle un-normalized quats + mjd.qpos[:] = np.random.random(mjm.nq) + + if kick: + # give the system a little kick to ensure we have non-identity rotations + mjd.qvel = np.random.uniform(-0.01, 0.01, mjm.nv) + mjd.ctrl = np.random.uniform(-0.1, 0.1, size=mjm.nu) + if applied: + mjd.qfrc_applied = np.random.uniform(-0.1, 0.1, size=mjm.nv) + mjd.xfrc_applied = np.random.uniform(-0.1, 0.1, size=mjd.xfrc_applied.shape) + if kick or applied: + mujoco.mj_step(mjm, mjd, nstep) # let dynamics get state significantly non-zero + + if mjm.nmocap: + mjd.mocap_pos = np.random.random(mjd.mocap_pos.shape) + mocap_quat = np.random.random(mjd.mocap_quat.shape) + mjd.mocap_quat = mocap_quat + + mujoco.mj_forward(mjm, mjd) + m = io.put_model(mjm) + if ls_parallel is not None: + m.opt.ls_parallel = ls_parallel + + d = io.put_data(mjm, mjd, nworld=nworld, nconmax=nconmax, njmax=njmax) + return mjm, mjd, m, d + + +def _sum(stack1, stack2): + ret = {} + for k in stack1: + times1, sub_stack1 = stack1[k] + times2, sub_stack2 = stack2[k] + times = [t1 + t2 for t1, t2 in zip(times1, times2)] + ret[k] = (times, _sum(sub_stack1, sub_stack2)) + return ret + + +@wp.kernel +def ctrl_noise( + # Model: + actuator_ctrllimited: wp.array(dtype=bool), + actuator_ctrlrange: wp.array2d(dtype=wp.vec2), + # In: + step: int, + ctrlnoise: float, + # Data out: + ctrl_out: wp.array2d(dtype=float), +): + worldid, actid = wp.tid() + + center = 0.0 + radius = 1.0 + ctrlrange = actuator_ctrlrange[0, actid] + if actuator_ctrllimited[actid]: + center = (ctrlrange[1] + ctrlrange[0]) / 2.0 + radius = (ctrlrange[1] - ctrlrange[0]) / 2.0 + radius *= ctrlnoise + noise = 2.0 * halton((step + 1) * (worldid + 1), actid + 2) - 1.0 + ctrl_out[worldid, actid] = center + radius * noise + + +def benchmark( + fn: Callable[[Model, Data], None], + m: Model, + d: Data, + nstep: int, + event_trace: bool = False, + measure_alloc: bool = False, + measure_solver_niter: bool = False, +) -> Tuple[float, float, dict, list, list, list]: + """Benchmark a function of Model and Data. + + Args: + fn (Callable[[Model, Data], None]): Function to benchmark. + m (Model): The model containing kinematic and dynamic information (device). + d (Data): The data object containing the current state and output information (device). + nstep (int): Number of timesteps. + event_trace (bool, optional): If True, time routines decorated with @event_scope. + Default is False. + measure_alloc (bool, optional): If True, record number of contacts and constraints. + Default is False. + measure_solver_niter (bool, False): If True, record the number of solver iterations. + Default is False. + Returns: + float: Time to JIT fn. + float: Total time to run the benchmark. + dict: Trace. + list: Number of contacts. + list: Number of constraints. + list: Number of solver iterations. + """ + jit_beg = time.perf_counter() + + fn(m, d) + + jit_end = time.perf_counter() + jit_duration = jit_end - jit_beg + wp.synchronize() + + trace = {} + ncon, nefc, solver_niter = [], [], [] + + with warp_util.EventTracer(enabled=event_trace) as tracer: + # capture the whole function as a CUDA graph + with wp.ScopedCapture() as capture: + fn(m, d) + graph = capture.graph + + time_vec = np.zeros(nstep) + for i in range(nstep): + with wp.ScopedStream(wp.get_stream()): + wp.launch( + ctrl_noise, + dim=(d.nworld, m.nu), + inputs=[ + m.actuator_ctrllimited, m.actuator_ctrlrange, i, 0.01 + ], + outputs=[d.ctrl]) # fmt: skip + + run_beg = time.perf_counter() + wp.capture_launch(graph) + wp.synchronize() + + run_end = time.perf_counter() + time_vec[i] = run_end - run_beg + if trace: + trace = _sum(trace, tracer.trace()) + else: + trace = tracer.trace() + if measure_alloc or measure_solver_niter: + wp.synchronize() + if measure_alloc: + ncon.append(d.ncon.numpy()[0]) + nefc.append(np.sum(d.nefc.numpy())) + if measure_solver_niter: + solver_niter.append(d.solver_niter.numpy()) + + wp.synchronize() + run_duration = np.sum(time_vec) + + return jit_duration, run_duration, trace, ncon, nefc, solver_niter diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py new file mode 100644 index 00000000..f1e08dca --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py @@ -0,0 +1,1667 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +import dataclasses +import enum + +import mujoco +import warp as wp + +MJ_MINVAL = mujoco.mjMINVAL +MJ_MAXVAL = mujoco.mjMAXVAL +MJ_MINIMP = mujoco.mjMINIMP # minimum constraint impedance +MJ_MAXIMP = mujoco.mjMAXIMP # maximum constraint impedance +MJ_MAXCONPAIR = mujoco.mjMAXCONPAIR +MJ_MINMU = mujoco.mjMINMU # minimum friction + + +# TODO(team): add check that all wp.launch_tiled 'block_dim' settings are configurable +@dataclasses.dataclass +class BlockDim: + """ + Block dimension 'block_dim' settings for wp.launch_tiled. + + TODO(team): experimental and may be removed + """ + + # collision_driver + segmented_sort: int = 128 + # derivative + qderiv_actuator_passive_actuation: int = 64 + qderiv_actuator_passive_no_actuation: int = 256 + # forward + euler_dense: int = 256 + actuator_velocity: int = 32 + tendon_velocity: int = 256 + qfrc_actuator: int = 256 + # ray + ray: int = 64 + # sensor + energy_vel_kinetic: int = 256 + # smooth + cholesky_factorize: int = 256 + cholesky_solve: int = 256 + cholesky_factorize_solve: int = 256 + # solver + update_gradient_cholesky: int = 256 + # support + mul_m_dense: int = 256 + + +class BroadphaseFilter(enum.IntFlag): + """Bitmask specifying which collision functions to run during broadphase. + + Attributes: + PLANE: collision between bounding sphere and plane. + SPHERE: collision between bounding spheres. + AABB: collision between axis-aligned bounding boxes. + OBB: collision between oriented bounding boxes. + """ + + PLANE = 1 + SPHERE = 2 + AABB = 4 + OBB = 8 + + +class CamLightType(enum.IntEnum): + """Type of camera light. + + Attributes: + FIXED: pos and rot fixed in body + TRACK: pos tracks body, rot fixed in global + TRACKCOM: pos tracks subtree com, rot fixed in body + TARGETBODY: pos fixed in body, rot tracks target body + TARGETBODYCOM: pos fixed in body, rot tracks target subtree com + """ + + FIXED = mujoco.mjtCamLight.mjCAMLIGHT_FIXED + TRACK = mujoco.mjtCamLight.mjCAMLIGHT_TRACK + TRACKCOM = mujoco.mjtCamLight.mjCAMLIGHT_TRACKCOM + TARGETBODY = mujoco.mjtCamLight.mjCAMLIGHT_TARGETBODY + TARGETBODYCOM = mujoco.mjtCamLight.mjCAMLIGHT_TARGETBODYCOM + + +class DataType(enum.IntFlag): + """Sensor data types. + + Attributes: + REAL: real values, no constraints + POSITIVE: positive values, 0 or negative: inactive + """ + + REAL = mujoco.mjtDataType.mjDATATYPE_REAL + POSITIVE = mujoco.mjtDataType.mjDATATYPE_POSITIVE + # unsupported: AXIS, QUATERNION + + +class DisableBit(enum.IntFlag): + """Disable default feature bitflags. + + Attributes: + CONSTRAINT: entire constraint solver + EQUALITY: equality constraints + FRICTIONLOSS: joint and tendon frictionloss constraints + LIMIT: joint and tendon limit constraints + CONTACT: contact constraints + PASSIVE: passive forces + GRAVITY: gravitational forces + CLAMPCTRL: clamp control to specified range + ACTUATION: apply actuation forces + REFSAFE: integrator safety: make ref[0]>=2*timestep + EULERDAMP: implicit damping for Euler integration + FILTERPARENT: disable collisions between parent and child bodies + SENSOR: sensors + """ + + CONSTRAINT = mujoco.mjtDisableBit.mjDSBL_CONSTRAINT + EQUALITY = mujoco.mjtDisableBit.mjDSBL_EQUALITY + FRICTIONLOSS = mujoco.mjtDisableBit.mjDSBL_FRICTIONLOSS + LIMIT = mujoco.mjtDisableBit.mjDSBL_LIMIT + CONTACT = mujoco.mjtDisableBit.mjDSBL_CONTACT + PASSIVE = mujoco.mjtDisableBit.mjDSBL_PASSIVE + GRAVITY = mujoco.mjtDisableBit.mjDSBL_GRAVITY + CLAMPCTRL = mujoco.mjtDisableBit.mjDSBL_CLAMPCTRL + ACTUATION = mujoco.mjtDisableBit.mjDSBL_ACTUATION + REFSAFE = mujoco.mjtDisableBit.mjDSBL_REFSAFE + EULERDAMP = mujoco.mjtDisableBit.mjDSBL_EULERDAMP + FILTERPARENT = mujoco.mjtDisableBit.mjDSBL_FILTERPARENT + SENSOR = mujoco.mjtDisableBit.mjDSBL_SENSOR + # unsupported: MIDPHASE, WARMSTART + + +class EnableBit(enum.IntFlag): + """Enable optional feature bitflags. + + Attributes: + ENERGY: energy computation + INVDISCRETE: discrete-time inverse dynamics + """ + + ENERGY = mujoco.mjtEnableBit.mjENBL_ENERGY + INVDISCRETE = mujoco.mjtEnableBit.mjENBL_INVDISCRETE + # unsupported: OVERRIDE, FWDINV, MULTICCD, ISLAND + + +class TrnType(enum.IntEnum): + """Type of actuator transmission. + + Attributes: + JOINT: force on joint + JOINTINPARENT: force on joint, expressed in parent frame + SLIDERCRANK: force via slider-crank linkage + TENDON: force on tendon + BODY: adhesion force on body's geoms + SITE: force on site + """ + + JOINT = mujoco.mjtTrn.mjTRN_JOINT + JOINTINPARENT = mujoco.mjtTrn.mjTRN_JOINTINPARENT + SLIDERCRANK = mujoco.mjtTrn.mjTRN_SLIDERCRANK + TENDON = mujoco.mjtTrn.mjTRN_TENDON + BODY = mujoco.mjtTrn.mjTRN_BODY + SITE = mujoco.mjtTrn.mjTRN_SITE + + +class DynType(enum.IntEnum): + """Type of actuator dynamics. + + Attributes: + NONE: no internal dynamics; ctrl specifies force + INTEGRATOR: integrator: da/dt = u + FILTER: linear filter: da/dt = (u-a) / tau + FILTEREXACT: linear filter: da/dt = (u-a) / tau, with exact integration + MUSCLE: piece-wise linear filter with two time constants + """ + + NONE = mujoco.mjtDyn.mjDYN_NONE + INTEGRATOR = mujoco.mjtDyn.mjDYN_INTEGRATOR + FILTER = mujoco.mjtDyn.mjDYN_FILTER + FILTEREXACT = mujoco.mjtDyn.mjDYN_FILTEREXACT + MUSCLE = mujoco.mjtDyn.mjDYN_MUSCLE + # unsupported: USER + + +class GainType(enum.IntEnum): + """Type of actuator gain. + + Attributes: + FIXED: fixed gain + AFFINE: const + kp*length + kv*velocity + MUSCLE: muscle FLV curve computed by muscle_gain + """ + + FIXED = mujoco.mjtGain.mjGAIN_FIXED + AFFINE = mujoco.mjtGain.mjGAIN_AFFINE + MUSCLE = mujoco.mjtGain.mjGAIN_MUSCLE + # unsupported: USER + + +class BiasType(enum.IntEnum): + """Type of actuator bias. + + Attributes: + NONE: no bias + AFFINE: const + kp*length + kv*velocity + MUSCLE: muscle passive force computed by muscle_bias + """ + + NONE = mujoco.mjtBias.mjBIAS_NONE + AFFINE = mujoco.mjtBias.mjBIAS_AFFINE + MUSCLE = mujoco.mjtBias.mjBIAS_MUSCLE + # unsupported: USER + + +class JointType(enum.IntEnum): + """Type of degree of freedom. + + Attributes: + FREE: global position and orientation (quat) (7,) + BALL: orientation (quat) relative to parent (4,) + SLIDE: sliding distance along body-fixed axis (1,) + HINGE: rotation angle (rad) around body-fixed axis (1,) + """ + + FREE = mujoco.mjtJoint.mjJNT_FREE + BALL = mujoco.mjtJoint.mjJNT_BALL + SLIDE = mujoco.mjtJoint.mjJNT_SLIDE + HINGE = mujoco.mjtJoint.mjJNT_HINGE + + def dof_width(self) -> int: + return {0: 6, 1: 3, 2: 1, 3: 1}[self.value] + + def qpos_width(self) -> int: + return {0: 7, 1: 4, 2: 1, 3: 1}[self.value] + + +class ConeType(enum.IntEnum): + """Type of friction cone. + + Attributes: + PYRAMIDAL: pyramidal + ELLIPTIC: elliptic + """ + + PYRAMIDAL = mujoco.mjtCone.mjCONE_PYRAMIDAL + ELLIPTIC = mujoco.mjtCone.mjCONE_ELLIPTIC + + +class IntegratorType(enum.IntEnum): + """Integrator mode. + + Attributes: + EULER: semi-implicit Euler + RK4: 4th-order Runge Kutta + IMPLICITFAST: implicit in velocity, no rne derivative + """ + + EULER = mujoco.mjtIntegrator.mjINT_EULER + RK4 = mujoco.mjtIntegrator.mjINT_RK4 + IMPLICITFAST = mujoco.mjtIntegrator.mjINT_IMPLICITFAST + # unsupported: IMPLICIT + + +class GeomType(enum.IntEnum): + """Type of geometry. + + Attributes: + PLANE: plane + HFIELD: heightfield + SPHERE: sphere + CAPSULE: capsule + ELLIPSOID: ellipsoid + CYLINDER: cylinder + BOX: box + MESH: mesh + SDF: sdf + """ + + PLANE = mujoco.mjtGeom.mjGEOM_PLANE + HFIELD = mujoco.mjtGeom.mjGEOM_HFIELD + SPHERE = mujoco.mjtGeom.mjGEOM_SPHERE + CAPSULE = mujoco.mjtGeom.mjGEOM_CAPSULE + ELLIPSOID = mujoco.mjtGeom.mjGEOM_ELLIPSOID + CYLINDER = mujoco.mjtGeom.mjGEOM_CYLINDER + BOX = mujoco.mjtGeom.mjGEOM_BOX + MESH = mujoco.mjtGeom.mjGEOM_MESH + SDF = mujoco.mjtGeom.mjGEOM_SDF + # unsupported: NGEOMTYPES, ARROW*, LINE, SKIN, LABEL, NONE + + +class SolverType(enum.IntEnum): + """Constraint solver algorithm. + + Attributes: + CG: Conjugate gradient (primal) + NEWTON: Newton (primal) + """ + + CG = mujoco.mjtSolver.mjSOL_CG + NEWTON = mujoco.mjtSolver.mjSOL_NEWTON + # unsupported: PGS + + +class ConstraintType(enum.IntEnum): + """Type of constraint. + + Attributes: + EQUALITY: equality constraint + FRICTION_DOF: dof friction + FRICTION_TENDON: tendon friction + LIMIT_JOINT: joint limit + LIMIT_TENDON: tendon limit + CONTACT_FRICTIONLESS: frictionless contact + CONTACT_PYRAMIDAL: frictional contact, pyramidal friction cone + CONTACT_ELLIPTIC: frictional contact, elliptic friction cone + """ + + EQUALITY = mujoco.mjtConstraint.mjCNSTR_EQUALITY + FRICTION_DOF = mujoco.mjtConstraint.mjCNSTR_FRICTION_DOF + FRICTION_TENDON = mujoco.mjtConstraint.mjCNSTR_FRICTION_TENDON + LIMIT_JOINT = mujoco.mjtConstraint.mjCNSTR_LIMIT_JOINT + LIMIT_TENDON = mujoco.mjtConstraint.mjCNSTR_LIMIT_TENDON + CONTACT_FRICTIONLESS = mujoco.mjtConstraint.mjCNSTR_CONTACT_FRICTIONLESS + CONTACT_PYRAMIDAL = mujoco.mjtConstraint.mjCNSTR_CONTACT_PYRAMIDAL + CONTACT_ELLIPTIC = mujoco.mjtConstraint.mjCNSTR_CONTACT_ELLIPTIC + + +class SensorType(enum.IntEnum): + """Type of sensor. + + Attributes: + MAGNETOMETER: magnetometer + CAMPROJECTION: camera projection + RANGEFINDER: scalar distance to nearest geom or site along z-axis + JOINTPOS: joint position + TENDONPOS: scalar tendon position + ACTUATORPOS: actuator position + BALLQUAT: ball joint orientation + JOINTLIMITPOS: joint limit distance-margin + TENDONLIMITPOS: tendon limit distance-margin + FRAMEPOS: frame position + FRAMEXAXIS: frame x-axis + FRAMEYAXIS: frame y-axis + FRAMEZAXIS: frame z-axis + FRAMEQUAT: frame orientation, represented as quaternion + SUBTREECOM: subtree center of mass + E_POTENTIAL: potential energy + E_KINETIC: kinetic energy + CLOCK: simulation time + VELOCIMETER: 3D linear velocity, in local frame + GYRO: 3D angular velocity, in local frame + JOINTVEL: joint velocity + TENDONVEL: scalar tendon velocity + ACTUATORVEL: actuator velocity + BALLANGVEL: ball joint angular velocity + JOINTLIMITVEL: joint limit velocity + TENDONLIMITVEL: tendon limit velocity + FRAMELINVEL: 3D linear velocity + FRAMEANGVEL: 3D angular velocity + SUBTREELINVEL: subtree linear velocity + SUBTREEANGMOM: subtree angular momentum + TOUCH: scalar contact normal forces summed over sensor zone + ACCELEROMETER: accelerometer + FORCE: force + TORQUE: torque + ACTUATORFRC: scalar actuator force, measured at the joint + TENDONACTFRC: scalar actuator force, measured at the tendon + JOINTACTFRC: scalar actuator force, measured at the joint + JOINTLIMITFRC: joint limit force + TENDONLIMITFRC: tendon limit force + FRAMELINACC: 3D linear acceleration + FRAMEANGACC: 3D angular acceleration + """ + + MAGNETOMETER = mujoco.mjtSensor.mjSENS_MAGNETOMETER + CAMPROJECTION = mujoco.mjtSensor.mjSENS_CAMPROJECTION + RANGEFINDER = mujoco.mjtSensor.mjSENS_RANGEFINDER + JOINTPOS = mujoco.mjtSensor.mjSENS_JOINTPOS + TENDONPOS = mujoco.mjtSensor.mjSENS_TENDONPOS + ACTUATORPOS = mujoco.mjtSensor.mjSENS_ACTUATORPOS + BALLQUAT = mujoco.mjtSensor.mjSENS_BALLQUAT + JOINTLIMITPOS = mujoco.mjtSensor.mjSENS_JOINTLIMITPOS + TENDONLIMITPOS = mujoco.mjtSensor.mjSENS_TENDONLIMITPOS + FRAMEPOS = mujoco.mjtSensor.mjSENS_FRAMEPOS + FRAMEXAXIS = mujoco.mjtSensor.mjSENS_FRAMEXAXIS + FRAMEYAXIS = mujoco.mjtSensor.mjSENS_FRAMEYAXIS + FRAMEZAXIS = mujoco.mjtSensor.mjSENS_FRAMEZAXIS + FRAMEQUAT = mujoco.mjtSensor.mjSENS_FRAMEQUAT + SUBTREECOM = mujoco.mjtSensor.mjSENS_SUBTREECOM + E_POTENTIAL = mujoco.mjtSensor.mjSENS_E_POTENTIAL + E_KINETIC = mujoco.mjtSensor.mjSENS_E_KINETIC + CLOCK = mujoco.mjtSensor.mjSENS_CLOCK + VELOCIMETER = mujoco.mjtSensor.mjSENS_VELOCIMETER + GYRO = mujoco.mjtSensor.mjSENS_GYRO + JOINTVEL = mujoco.mjtSensor.mjSENS_JOINTVEL + TENDONVEL = mujoco.mjtSensor.mjSENS_TENDONVEL + ACTUATORVEL = mujoco.mjtSensor.mjSENS_ACTUATORVEL + BALLANGVEL = mujoco.mjtSensor.mjSENS_BALLANGVEL + JOINTLIMITVEL = mujoco.mjtSensor.mjSENS_JOINTLIMITVEL + TENDONLIMITVEL = mujoco.mjtSensor.mjSENS_TENDONLIMITVEL + FRAMELINVEL = mujoco.mjtSensor.mjSENS_FRAMELINVEL + FRAMEANGVEL = mujoco.mjtSensor.mjSENS_FRAMEANGVEL + SUBTREELINVEL = mujoco.mjtSensor.mjSENS_SUBTREELINVEL + SUBTREEANGMOM = mujoco.mjtSensor.mjSENS_SUBTREEANGMOM + TOUCH = mujoco.mjtSensor.mjSENS_TOUCH + ACCELEROMETER = mujoco.mjtSensor.mjSENS_ACCELEROMETER + FORCE = mujoco.mjtSensor.mjSENS_FORCE + TORQUE = mujoco.mjtSensor.mjSENS_TORQUE + ACTUATORFRC = mujoco.mjtSensor.mjSENS_ACTUATORFRC + TENDONACTFRC = mujoco.mjtSensor.mjSENS_TENDONACTFRC + JOINTACTFRC = mujoco.mjtSensor.mjSENS_JOINTACTFRC + JOINTLIMITFRC = mujoco.mjtSensor.mjSENS_JOINTLIMITFRC + TENDONLIMITFRC = mujoco.mjtSensor.mjSENS_TENDONLIMITFRC + FRAMELINACC = mujoco.mjtSensor.mjSENS_FRAMELINACC + FRAMEANGACC = mujoco.mjtSensor.mjSENS_FRAMEANGACC + + +class ObjType(enum.IntEnum): + """Type of object. + + Attributes: + UNKNOWN: unknown object type + BODY: body + XBODY: body, used to access regular frame instead of i-frame + GEOM: geom + SITE: site + CAMERA: camera + """ + + UNKNOWN = mujoco.mjtObj.mjOBJ_UNKNOWN + BODY = mujoco.mjtObj.mjOBJ_BODY + XBODY = mujoco.mjtObj.mjOBJ_XBODY + GEOM = mujoco.mjtObj.mjOBJ_GEOM + SITE = mujoco.mjtObj.mjOBJ_SITE + CAMERA = mujoco.mjtObj.mjOBJ_CAMERA + + +class EqType(enum.IntEnum): + """Type of equality constraint. + + Attributes: + CONNECT: connect two bodies at a point (ball joint) + JOINT: couple the values of two scalar joints with cubic + WELD: fix relative position and orientation of two bodies + """ + + CONNECT = mujoco.mjtEq.mjEQ_CONNECT + WELD = mujoco.mjtEq.mjEQ_WELD + JOINT = mujoco.mjtEq.mjEQ_JOINT + TENDON = mujoco.mjtEq.mjEQ_TENDON + # unsupported: FLEX, DISTANCE + + +class WrapType(enum.IntEnum): + """Type of tendon wrapping object. + + Attributes: + JOINT: constant moment arm + PULLEY: pulley used to split tendon + SITE: pass through site + SPHERE: wrap around sphere + CYLINDER: wrap around (infinite) cylinder + """ + + JOINT = mujoco.mjtWrap.mjWRAP_JOINT + PULLEY = mujoco.mjtWrap.mjWRAP_PULLEY + SITE = mujoco.mjtWrap.mjWRAP_SITE + SPHERE = mujoco.mjtWrap.mjWRAP_SPHERE + CYLINDER = mujoco.mjtWrap.mjWRAP_CYLINDER + + +class vec5f(wp.types.vector(length=5, dtype=float)): + pass + + +class vec6f(wp.types.vector(length=6, dtype=float)): + pass + + +class vec10f(wp.types.vector(length=10, dtype=float)): + pass + + +class vec11f(wp.types.vector(length=11, dtype=float)): + pass + + +vec5 = vec5f +vec6 = vec6f +vec10 = vec10f +vec11 = vec11f + + +class BroadphaseType(enum.IntEnum): + """Type of broadphase algorithm. + + Attributes: + NXN: Broad phase checking all pairs + SAP_TILE: Sweep and prune broad phase using tile sort + SAP_SEGMENTED: Sweep and prune broad phase using segment sort + """ + + NXN = 0 + SAP_TILE = 1 + SAP_SEGMENTED = 2 + + +@dataclasses.dataclass +class Option: + """Physics options. + + Attributes: + timestep: simulation timestep + impratio: ratio of friction-to-normal contact impedance + tolerance: main solver tolerance + ls_tolerance: CG/Newton linesearch tolerance + gravity: gravitational acceleration + magnetic: global magnetic flux + integrator: integration mode (mjtIntegrator) + cone: type of friction cone (mjtCone) + solver: solver algorithm (mjtSolver) + iterations: number of main solver iterations + ls_iterations: maximum number of CG/Newton linesearch iterations + disableflags: bit flags for disabling standard features + enableflags: bit flags for enabling optional features + is_sparse: whether to use sparse representations + gjk_iterations: number of Gjk iterations in the convex narrowphase + epa_iterations: number of Epa iterations in the convex narrowphase + ls_parallel: evaluate engine solver step sizes in parallel + wind: wind (for lift, drag, and viscosity) + has_fluid: True if wind, density, or viscosity are non-zero at put_model time + density: density of medium + viscosity: viscosity of medium + broadphase: broadphase type, 0: nxn, 1: sap_tile, 2: sap_segmented + broadphase_filter: broadphase filter bitflag + graph_conditional: flag to use cuda graph conditional, should be False when JAX is used + sdf_initpoints: number of starting points for gradient descent + sdf_iterations: max number of iterations for gradient descent + run_collision_detection: if False, skips collision detection and allows user-populated + contacts during the physics step (as opposed to DisableBit.CONTACT which explicitly + zeros out the contacts at each step) + """ + + timestep: wp.array(dtype=float) + impratio: wp.array(dtype=float) + tolerance: wp.array(dtype=float) + ls_tolerance: wp.array(dtype=float) + gravity: wp.array(dtype=wp.vec3) + magnetic: wp.array(dtype=wp.vec3) + integrator: int + cone: int + solver: int + iterations: int + ls_iterations: int + disableflags: int + enableflags: int + is_sparse: bool + gjk_iterations: int # warp only + epa_iterations: int # warp only + ls_parallel: bool + wind: wp.array(dtype=wp.vec3) + has_fluid: bool + density: wp.array(dtype=float) + viscosity: wp.array(dtype=float) + broadphase: int # warp only + broadphase_filter: int # warp only + graph_conditional: bool # warp only + sdf_initpoints: int + sdf_iterations: int + run_collision_detection: bool # warp only + + +@dataclasses.dataclass +class Statistic: + """Model statistics (in qpos0). + + Attributes: + meaninertia: mean diagonal inertia + """ + + meaninertia: float + + +@dataclasses.dataclass +class Constraint: + """Constraint data. + + Attributes: + type: constraint type (mjtConstraint) (nworld, njmax) + id: id of object of specific type (nworld, njmax) + J: constraint Jacobian (nworld, njmax, nv) + pos: constraint position (equality, contact) (nworld, njmax) + margin: inclusion margin (contact) (nworld, njmax) + D: constraint mass (nworld, njmax) + vel: velocity in constraint space: J*qvel (nworld, njmax) + aref: reference pseudo-acceleration (nworld, njmax) + frictionloss: frictionloss (friction) (nworld, njmax) + force: constraint force in constraint space (nworld, njmax) + Jaref: Jac*qacc - aref (nworld, njmax) + Ma: M*qacc (nworld, nv) + grad: gradient of master cost (nworld, nv) + grad_dot: dot(grad, grad) (nworld,) + Mgrad: M / grad (nworld, nv) + search: linesearch vector (nworld, nv) + search_dot: dot(search, search) (nworld,) + gauss: gauss Cost (nworld,) + cost: constraint + Gauss cost (nworld,) + prev_cost: cost from previous iter (nworld,) + active: active (quadratic) constraints (nworld, njmax) + gtol: linesearch termination tolerance (nworld,) + mv: qM @ search (nworld, nv) + jv: efc_J @ search (nworld, njmax) + quad: quadratic cost coefficients (nworld, njmax, 3) + quad_gauss: quadratic cost gauss coefficients (nworld, 3) + h: cone hessian (nworld, nv, nv) + alpha: line search step size (nworld,) + prev_grad: previous grad (nworld, nv) + prev_Mgrad: previous Mgrad (nworld, nv) + beta: polak-ribiere beta (nworld,) + beta_num: numerator of beta (nworld,) + beta_den: denominator of beta (nworld,) + done: solver done (nworld,) + ls_done: linesearch done (nworld,) + p0: initial point (nworld, 3) + lo: low point bounding the line search interval (nworld, 3) + lo_alpha: alpha for low point (nworld,) + hi: high point bounding the line search interval (nworld, 3) + hi_alpha: alpha for high point (nworld,) + lo_next: next low point (nworld, 3) + lo_next_alpha: alpha for next low point (nworld,) + hi_next: next high point (nworld, 3) + hi_next_alpha: alpha for next high point (nworld,) + mid: loss at mid_alpha (nworld, 3) + mid_alpha: midpoint between lo_alpha and hi_alpha (nworld,) + cost_candidate: costs associated with step sizes (nworld, nlsp) + u: friction cone (normal and tangents) (nconmax, 6) + uu: elliptic cone variables (nconmax,) + uv: elliptic cone variables (nconmax,) + vv: elliptic cone variables (nconmax,) + condim: if contact: condim, else: -1 (nworld, njmax) + """ + + type: wp.array2d(dtype=int) + id: wp.array2d(dtype=int) + J: wp.array3d(dtype=float) + pos: wp.array2d(dtype=float) + margin: wp.array2d(dtype=float) + D: wp.array2d(dtype=float) + vel: wp.array2d(dtype=float) + aref: wp.array2d(dtype=float) + frictionloss: wp.array2d(dtype=float) + force: wp.array2d(dtype=float) + Jaref: wp.array2d(dtype=float) + Ma: wp.array2d(dtype=float) + grad: wp.array2d(dtype=float) + cholesky_L_tmp: wp.array3d(dtype=float) + cholesky_y_tmp: wp.array2d(dtype=float) + grad_dot: wp.array(dtype=float) + Mgrad: wp.array2d(dtype=float) + search: wp.array2d(dtype=float) + search_dot: wp.array(dtype=float) + gauss: wp.array(dtype=float) + cost: wp.array(dtype=float) + prev_cost: wp.array(dtype=float) + active: wp.array2d(dtype=bool) + gtol: wp.array(dtype=float) + mv: wp.array2d(dtype=float) + jv: wp.array2d(dtype=float) + quad: wp.array2d(dtype=wp.vec3) + quad_gauss: wp.array(dtype=wp.vec3) + h: wp.array3d(dtype=float) + alpha: wp.array(dtype=float) + prev_grad: wp.array2d(dtype=float) + prev_Mgrad: wp.array2d(dtype=float) + beta: wp.array(dtype=float) + beta_num: wp.array(dtype=float) + beta_den: wp.array(dtype=float) + done: wp.array(dtype=bool) + # linesearch + ls_done: wp.array(dtype=bool) + p0: wp.array(dtype=wp.vec3) + lo: wp.array(dtype=wp.vec3) + lo_alpha: wp.array(dtype=float) + hi: wp.array(dtype=wp.vec3) + hi_alpha: wp.array(dtype=float) + lo_next: wp.array(dtype=wp.vec3) + lo_next_alpha: wp.array(dtype=float) + hi_next: wp.array(dtype=wp.vec3) + hi_next_alpha: wp.array(dtype=float) + mid: wp.array(dtype=wp.vec3) + mid_alpha: wp.array(dtype=float) + cost_candidate: wp.array2d(dtype=float) + # elliptic cone + u: wp.array(dtype=vec6) + uu: wp.array(dtype=float) + uv: wp.array(dtype=float) + vv: wp.array(dtype=float) + condim: wp.array2d(dtype=int) + + +@dataclasses.dataclass +class TileSet: + """Tiling configuration for decomposable block diagonal matrix. + + For non-square, non-block-diagonal tiles, use two tilesets. + + Attributes: + adr: address of each tile in the set + size: size of all the tiles in this set + """ + + adr: wp.array(dtype=int) + size: int + + +# TODO(team): make Model/Data fields sort order match mujoco + + +@dataclasses.dataclass +class Model: + """Model definition and parameters. + + Attributes: + nq: number of generalized coordinates + nv: number of degrees of freedom + nu: number of actuators/controls + na: number of activation states + nbody: number of bodies + njnt: number of joints + ngeom: number of geoms + nsite: number of sites + ncam: number of cameras + nlight: number of lights + nexclude: number of excluded geom pairs + neq: number of equality constraints + nmocap: number of mocap bodies + ngravcomp: number of bodies with nonzero gravcomp + nM: number of non-zeros in sparse inertia matrix + nC: number of non-zeros in sparse reduced dof-dof matrix + ntendon: number of tendons + nwrap: number of wrap objects in all tendon paths + nsensor: number of sensors + nsensordata: number of elements in sensor data vector + nmeshvert: number of vertices for all meshes + nmeshface: number of faces for all meshes + nmeshgraph: number of ints in mesh auxiliary data + nmeshpoly: number of polygons in all meshes + nmeshpolyvert: number of vertices in all polygons + nmeshpolymap: number of polygons in vertex map + nlsp: number of step sizes for parallel linsearch + npair: number of predefined geom pairs + nhfield: number of heightfields + nhfielddata: size of elevation data + opt: physics options + stat: model statistics + qpos0: qpos values at default pose (nworld, nq) + qpos_spring: reference pose for springs (nworld, nq) + qM_fullm_i: sparse mass matrix addressing + qM_fullm_j: sparse mass matrix addressing + qM_mulm_i: sparse mass matrix addressing + qM_mulm_j: sparse mass matrix addressing + qM_madr_ij: sparse mass matrix addressing + qLD_update_tree: dof tree ordering for qLD updates + qLD_update_treeadr: index of each dof tree level + M_rownnz: number of non-zeros in each row of qM (nv,) + M_rowadr: index of each row in qM (nv,) + M_colind: column indices of non-zeros in qM (nM,) + mapM2M: index mapping from M (legacy) to M (CSR) (nC) + qM_tiles: tiling configuration + body_tree: list of body ids by tree level + body_parentid: id of body's parent (nbody,) + body_rootid: id of root above body (nbody,) + body_weldid: id of body that this body is welded to (nbody,) + body_mocapid: id of mocap data; -1: none (nbody,) + body_jntnum: number of joints for this body (nbody,) + body_jntadr: start addr of joints; -1: no joints (nbody,) + body_dofnum: number of motion degrees of freedom (nbody,) + body_dofadr: start addr of dofs; -1: no dofs (nbody,) + body_geomnum: number of geoms (nbody,) + body_geomadr: start addr of geoms; -1: no geoms (nbody,) + body_pos: position offset rel. to parent body (nworld, nbody, 3) + body_quat: orientation offset rel. to parent body (nworld, nbody, 4) + body_ipos: local position of center of mass (nworld, nbody, 3) + body_iquat: local orientation of inertia ellipsoid (nworld, nbody, 4) + body_mass: mass (nworld, nbody,) + body_subtreemass: mass of subtree starting at this body (nworld, nbody,) + subtree_mass: mass of subtree (nworld, nbody,) + body_inertia: diagonal inertia in ipos/iquat frame (nworld, nbody, 3) + body_invweight0: mean inv inert in qpos0 (trn, rot) (nworld, nbody, 2) + body_contype: OR over all geom contypes (nbody,) + body_conaffinity: OR over all geom conaffinities (nbody,) + body_gravcomp: antigravity force, units of body weight (nworld, nbody) + jnt_type: type of joint (mjtJoint) (njnt,) + jnt_qposadr: start addr in 'qpos' for joint's data (njnt,) + jnt_dofadr: start addr in 'qvel' for joint's data (njnt,) + jnt_bodyid: id of joint's body (njnt,) + jnt_limited: does joint have limits (njnt,) + jnt_actfrclimited: does joint have actuator force limits (njnt,) + jnt_solref: constraint solver reference: limit (nworld, njnt, mjNREF) + jnt_solimp: constraint solver impedance: limit (nworld, njnt, mjNIMP) + jnt_pos: local anchor position (nworld, njnt, 3) + jnt_axis: local joint axis (nworld, njnt, 3) + jnt_stiffness: stiffness coefficient (nworld, njnt) + jnt_range: joint limits (nworld, njnt, 2) + jnt_actfrcrange: range of total actuator force (nworld, njnt, 2) + jnt_margin: min distance for limit detection (nworld, njnt) + jnt_limited_slide_hinge_adr: limited/slide/hinge jntadr + jnt_limited_ball_adr: limited/ball jntadr + jnt_actgravcomp: is gravcomp force applied via actuators (njnt,) + dof_bodyid: id of dof's body (nv,) + dof_jntid: id of dof's joint (nv,) + dof_parentid: id of dof's parent; -1: none (nv,) + dof_Madr: dof address in M-diagonal (nv,) + dof_armature: dof armature inertia/mass (nworld, nv) + dof_damping: damping coefficient (nworld, nv) + dof_invweight0: diag. inverse inertia in qpos0 (nworld, nv) + dof_frictionloss: dof friction loss (nworld, nv) + dof_solimp: constraint solver impedance: frictionloss (nworld, nv, NIMP) + dof_solref: constraint solver reference: frictionloss (nworld, nv, NREF) + dof_tri_row: np.tril_indices (mjm.nv)[0] + dof_tri_col: np.tril_indices (mjm.nv)[1] + geom_type: geometric type (mjtGeom) (ngeom,) + geom_contype: geom contact type (ngeom,) + geom_conaffinity: geom contact affinity (ngeom,) + geom_condim: contact dimensionality (1, 3, 4, 6) (ngeom,) + geom_bodyid: id of geom's body (ngeom,) + geom_dataid: id of geom's mesh/hfield; -1: none (ngeom,) + geom_group: geom group inclusion/exclusion mask (ngeom,) + geom_matid: material id for rendering (nworld, ngeom,) + geom_priority: geom contact priority (ngeom,) + geom_solmix: mixing coef for solref/imp in geom pair (nworld, ngeom,) + geom_solref: constraint solver reference: contact (nworld, ngeom, mjNREF) + geom_solimp: constraint solver impedance: contact (nworld, ngeom, mjNIMP) + geom_size: geom-specific size parameters (ngeom, 3) + geom_aabb: bounding box, (center, size) (ngeom, 6) + geom_rbound: radius of bounding sphere (nworld, ngeom,) + geom_pos: local position offset rel. to body (nworld, ngeom, 3) + geom_quat: local orientation offset rel. to body (nworld, ngeom, 4) + geom_friction: friction for (slide, spin, roll) (nworld, ngeom, 3) + geom_margin: detect contact if dist bool: + """Check for intersection of two 2D line segments. + + Args: + p1: 2D point from segment 1 + p2: 2D point from segment 1 + p3: 2D point from segment 2 + p4: 2D point from segment 2 + + Returns: + intersection status of line segments + """ + # compute determinant, check + det = (p4[1] - p3[1]) * (p2[0] - p1[0]) - (p4[0] - p3[0]) * (p2[1] - p1[1]) + + if wp.abs(det) < MJ_MINVAL: + return False + + # compute intersection point on each line + a = ((p4[0] - p3[0]) * (p1[1] - p3[1]) - (p4[1] - p3[1]) * (p1[0] - p3[0])) / det + b = ((p2[0] - p1[0]) * (p1[1] - p3[1]) - (p2[1] - p1[1]) * (p1[0] - p3[0])) / det + + if a >= 0 and a <= 1.0 and b >= 0.0 and b <= 1.0: + return True + else: + return False + + +@wp.func +def halton(index: int, base: int) -> float: + n0 = index + b = float(base) + f = float(1.0) / b + hn = float(0.0) + + while n0 > 0: + n1 = n0 // base + r = n0 - n1 * base + hn += f * float(r) + f /= b + n0 = n1 + + return hn + + +@wp.func +def length_circle(p0: wp.vec2, p1: wp.vec2, ind: int, radius: float) -> float: + """Curve length along circle. + + Args: + p0: 2D point + p1: 2D point + ind: input for flip + radius: circle radius + + Returns: + curve length + """ + # compute angle between 0 and pi + p0n, _ = math.normalize_with_norm(p0) + p1n, _ = math.normalize_with_norm(p1) + + angle = wp.acos(wp.dot(p0n, p1n)) + + # flip if necessary + cross = p0[1] * p1[0] - p0[0] * p1[1] + if (cross > 0.0 and ind != 0) or (cross < 0.0 and ind == 0): + angle = 2.0 * wp.pi - angle + + return radius * angle + + +@wp.func +def wrap_circle(end: wp.vec4, side: wp.vec2, radius: float) -> Tuple[float, wp.vec2, wp.vec2]: + """2D circle wrap. + + Args: + end: two 2D points + side: optional 2D side point, no side point: wp.vec2(wp.inf) + radius: circle radius + + Returns: + length of circular wrap or -1.0 if no wrap, pair of 2D wrap points + """ + valid_side = wp.norm_l2(side) < wp.inf + + end0 = wp.vec2(end[0], end[1]) + end1 = wp.vec2(end[2], end[3]) + + sqlen0 = wp.dot(end0, end0) + sqlen1 = wp.dot(end1, end1) + sqrad = radius * radius + + # either point inside circle or circle too small: no wrap + if (sqlen0 < sqrad) or (sqlen1 < sqrad) or (radius < MJ_MINVAL): + return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf) + + # points too close: no wrap + dif = end1 - end0 + dd = wp.dot(dif, dif) + if dd < MJ_MINVAL: + return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf) + + # find nearest point on line segment to origin: a * dif + d0 + a = -wp.dot(dif, end0) / dd + a = wp.clamp(a, 0.0, 1.0) + + # check for intersection and side + tmp = a * dif + end0 + if (wp.dot(tmp, tmp) > sqrad) and (not valid_side or wp.dot(side, tmp) >= 0.0): + return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf) + + sqrt0 = wp.sqrt(sqlen0 - sqrad) + sqrt1 = wp.sqrt(sqlen1 - sqrad) + + # construct the two solutions, compute goodness + sol00 = wp.vec2( + (end[0] * sqrad + radius * end[1] * sqrt0) / sqlen0, + (end[1] * sqrad - radius * end[0] * sqrt0) / sqlen0, + ) + sol01 = wp.vec2( + (end[2] * sqrad - radius * end[3] * sqrt1) / sqlen1, + (end[3] * sqrad + radius * end[2] * sqrt1) / sqlen1, + ) + + sol10 = wp.vec2( + (end[0] * sqrad - radius * end[1] * sqrt0) / sqlen0, + (end[1] * sqrad + radius * end[0] * sqrt0) / sqlen0, + ) + sol11 = wp.vec2( + (end[2] * sqrad + radius * end[3] * sqrt1) / sqlen1, + (end[3] * sqrad - radius * end[2] * sqrt1) / sqlen1, + ) + + # goodness: close to sd, or shorter path + if valid_side: + tmp0, _ = math.normalize_with_norm(sol00 + sol01) + good0 = wp.dot(tmp0, side) + tmp1, _ = math.normalize_with_norm(sol10 + sol11) + good1 = wp.dot(tmp1, side) + else: + tmp0 = sol00 - sol01 + good0 = -wp.dot(tmp0, tmp0) + tmp1 = sol10 - sol11 + good1 = -wp.dot(tmp1, tmp1) + + # penalize for intersection + if is_intersect(end0, sol00, end1, sol01): + good0 = -10000.0 + if is_intersect(end0, sol10, end1, sol11): + good1 = -10000.0 + + # select the better solution + if good0 > good1: + pnt0 = sol00 + pnt1 = sol01 + ind = 0 + else: + pnt0 = sol10 + pnt1 = sol11 + ind = 1 + + # check for intersection + if is_intersect(end0, pnt0, end1, pnt1): + return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf) + + # return curve length + return length_circle(pnt0, pnt1, ind, radius), pnt0, pnt1 + + +@wp.func +def wrap_inside( + # In: + end: wp.vec4, + radius: float, + # TODO(team): update kernel analyzer to allow defaults + maxiter: int = 20, # kernel_analyzer: ignore + zinit: float = 1.0 - 1.0e-7, # kernel_analyzer: ignore + tolerance: float = 1.0e-6, # kernel_analyzer: ignore +) -> Tuple[float, wp.vec2, wp.vec2]: + """2D inside wrap. + + Args: + end: two 2D points + radius: circle radius + maxiter: maximum number of solver iterations + zinit: initialization for solver + tolerance: solver convergence tolerance + + Returns: + 0.0 if wrap else -1.0, pair of 2D wrap points + """ + + end0 = wp.vec2(end[0], end[1]) + end1 = wp.vec2(end[2], end[3]) + + # constants + len0 = wp.norm_l2(end0) + len1 = wp.norm_l2(end1) + dif = end1 - end0 + dd = wp.dot(dif, dif) + + # either point inside circle or circle too small: no wrap + if (len0 <= radius) or (len1 <= radius) or (radius < MJ_MINVAL) or (len0 < MJ_MINVAL) or (len1 < MJ_MINVAL): + return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf) + + # segment-circle intersection: no wrap + if dd > MJ_MINVAL: + # find nearest point on line segment to origin: d0 + a * dif + a = -wp.dot(dif, end0) / dd + + # in segment + if (a > 0.0) and (a < 1.0): + tmp = end0 + a * dif + if wp.norm_l2(tmp) <= radius: + return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf) + + # prepare default in case of numerical failure: average + pnt = 0.5 * (end0 + end1) + pnt, _ = math.normalize_with_norm(pnt) + pnt *= radius + + # compute function parameters: asin(A * z) + asin(B * z) - 2 * asin(z) + G = 0 + A = radius / len0 + B = radius / len1 + sq_A = A * A + sq_B = B * B + cosG = (len0 * len0 + len1 * len1 - dd) / (2.0 * len0 * len1) + if cosG < -1.0 + MJ_MINVAL: + return -1.0, pnt, pnt + elif cosG > 1.0 - MJ_MINVAL: + return 0.0, pnt, pnt + G = wp.acos(cosG) + + # init + z = zinit + f = wp.asin(A * z) + wp.asin(B * z) - 2.0 * wp.asin(z) + G + + # make sure init is not on the other side + if f > 0.0: + return 0.0, pnt, pnt + + # Newton method + iter = int(0) + + while (iter < maxiter) and (wp.abs(f) > tolerance): + # derivative + sq_z = z * z + df = ( + A / wp.max(MJ_MINVAL, wp.sqrt(1.0 - sq_z * sq_A)) + + B / wp.max(MJ_MINVAL, wp.sqrt(1.0 - sq_z * sq_B)) + - 2.0 / wp.max(MJ_MINVAL, wp.sqrt(1.0 - sq_z)) + ) + + # check sign; SHOULD NOT OCCUR + if df > -MJ_MINVAL: + return 0.0, pnt, pnt + + # new point + z1 = z - f / df + + # make sure we are moving to the left; SHOULD NOT OCCUR + if z1 > z: + return 0.0, pnt, pnt + + # update solution + z = z1 + f = wp.asin(A * z) + wp.asin(B * z) - 2.0 * wp.asin(z) + G + + # exit if positive: SHOULD NOT OCCUR + if f > tolerance: + return 0.0, pnt, pnt + + iter += 1 + + # check convergence + if iter >= maxiter: + return 0.0, pnt, pnt + + # finalize: rotation by ang from vec = a or b, depending on cross(a, b) sign + if end[0] * end[3] - end[1] * end[2] > 0.0: + vec = end0 + ang = wp.asin(z) - wp.asin(A * z) + else: + vec = end1 + ang = wp.asin(z) - wp.asin(B * z) + + vec, _ = math.normalize_with_norm(vec) + pnt = wp.vec2( + radius * (wp.cos(ang) * vec[0] - wp.sin(ang) * vec[1]), + radius * (wp.sin(ang) * vec[0] + wp.cos(ang) * vec[1]), + ) + + return 0.0, pnt, pnt + + +@wp.func +def wrap( + x0: wp.vec3, x1: wp.vec3, pos: wp.vec3, mat: wp.mat33, radius: float, geomtype: int, side: wp.vec3 +) -> Tuple[float, wp.vec3, wp.vec3]: + """Wrap tendons around spheres and cylinders. + + Args: + x0: 3D endpoint + x1: 3D endpoint + pos: position of geom + mat: orientation of geom + radius: geom radius + type: wrap type (mjtWrap) + side: 3D position for sidesite, no side point: wp.vec3(wp.inf) + + Returns: + length of circular wrap else -1.0 if no wrap, pair of 3D wrap points + """ + # check object type + if geomtype != int(WrapType.SPHERE.value) and geomtype != int(WrapType.CYLINDER.value): + return wp.inf, wp.vec3(wp.inf), wp.vec3(wp.inf) + + # map sites to wrap object's local frame + matT = wp.transpose(mat) + p0 = matT @ (x0 - pos) + p1 = matT @ (x1 - pos) + + # too close to origin: return + if (wp.norm_l2(p0) < MJ_MINVAL) or (wp.norm_l2(p1) < MJ_MINVAL): + return -1.0, wp.vec3(wp.inf), wp.vec3(wp.inf) + + # construct 2D frame for circle wrap + if geomtype == int(WrapType.SPHERE.value): + # 1st axis = p0 + axis0, _ = math.normalize_with_norm(p0) + + # normal to p0-0-p1 plane = cross(p0, p1) + normal = wp.cross(p0, p1) + normal, nrm = math.normalize_with_norm(normal) + + # if (p0, p1) parallel: different normal + if nrm < MJ_MINVAL: + # find max component of axis0 + axis0_abs = wp.abs(axis0) + i = int(0) + if (axis0_abs[1] > axis0_abs[0]) and (axis0_abs[1] > axis0_abs[2]): + i = 1 + if (axis0_abs[2] > axis0_abs[0]) and (axis0_abs[2] > axis0_abs[1]): + i = 2 + + # init second axis: 0 at i; 1 elsewhere + axis1 = wp.vec3(1.0) + axis1[i] = 0.0 + + # recompute normal + normal = wp.cross(axis0, axis1) + normal, _ = math.normalize_with_norm(normal) + + # 2nd axis = cross(normal, p0) + axis1 = wp.cross(normal, axis0) + axis1, _ = math.normalize_with_norm(axis1) + else: # WrapType.CYLINDER + # 1st axis = x + axis0 = wp.vec3(1.0, 0.0, 0.0) + + # 2nd axis = y + axis1 = wp.vec3(0.0, 1.0, 0.0) + + # project points in 2D frame: p => end + end = wp.vec4( + wp.dot(p0, axis0), + wp.dot(p0, axis1), + wp.dot(p1, axis0), + wp.dot(p1, axis1), + ) + + # handle sidesite + valid_side = wp.norm_l2(side) < wp.inf + + if valid_side: + # side point: apply same projection as x0, x1 + sidepnt = matT @ (side - pos) + + # side point: project and rescale + sidepnt_proj = wp.vec2( + wp.dot(sidepnt, axis0), + wp.dot(sidepnt, axis1), + ) + + sidepnt_proj, _ = math.normalize_with_norm(sidepnt_proj) + sidepnt_proj *= radius + else: + sidepnt_proj = wp.vec2(wp.inf) + + # apply inside wrap + if valid_side and wp.norm_l2(sidepnt) < radius: + wlen, pnt0, pnt1 = wrap_inside(end, radius) + else: # apply circle wrap + wlen, pnt0, pnt1 = wrap_circle(end, sidepnt_proj, radius) + + # no wrap: return + if wlen < 0.0: + return -1.0, wp.vec3(wp.inf), wp.vec3(wp.inf) + + # reconstruct 3D points in local frame: res + res0 = axis0 * pnt0[0] + axis1 * pnt0[1] + res1 = axis0 * pnt1[0] + axis1 * pnt1[1] + + # cylinder: correct along z + if geomtype == int(WrapType.CYLINDER.value): + # set vertical coordinates + L0 = wp.sqrt((p0[0] - res0[0]) * (p0[0] - res0[0]) + (p0[1] - res0[1]) * (p0[1] - res0[1])) + L1 = wp.sqrt((p1[0] - res1[0]) * (p1[0] - res1[0]) + (p1[1] - res1[1]) * (p1[1] - res1[1])) + res0[2] = p0[2] + (p1[2] - p0[2]) * L0 / (L0 + wlen + L1) + res1[2] = p0[2] + (p1[2] - p0[2]) * (L0 + wlen) / (L0 + wlen + L1) + + # correct wlen for height + height = wp.abs(res1[2] - res0[2]) + wlen = wp.sqrt(wlen * wlen + height * height) + + # map back to global frame: wpnt + wpnt0 = mat @ res0 + pos + wpnt1 = mat @ res1 + pos + + return wlen, wpnt0, wpnt1 + + +@wp.func +def muscle_gain_length(length: float, lmin: float, lmax: float) -> float: + """Normalized muscle length-gain curve.""" + + if (lmin > length) or (length > lmax): + return 0.0 + + # mid-ranges (maximum is at 1.0) + a = 0.5 * (lmin + 1.0) + b = 0.5 * (1.0 + lmax) + + if length <= a: + x = (length - lmin) / wp.max(MJ_MINVAL, a - lmin) + return 0.5 * x * x + elif length <= 1.0: + x = (1.0 - length) / wp.max(MJ_MINVAL, 1.0 - a) + return 1.0 - 0.5 * x * x + elif length <= b: + x = (length - 1.0) / wp.max(MJ_MINVAL, b - 1.0) + return 1.0 - 0.5 * x * x + else: + x = (lmax - length) / wp.max(MJ_MINVAL, lmax - b) + return 0.5 * x * x + + +@wp.func +def muscle_gain(len: float, vel: float, lengthrange: wp.vec2, acc0: float, prm: vec10) -> float: + """Muscle active force, prm = (range[2], force, scale, lmin, lmax, vmax, fpmax, fvmax).""" + + # unpack parameters + range_ = wp.vec2(prm[0], prm[1]) + force = prm[2] + scale = prm[3] + lmin = prm[4] + lmax = prm[5] + vmax = prm[6] + fvmax = prm[8] + + # scale force if negative + if force < 0.0: + force = scale / wp.max(MJ_MINVAL, acc0) + + # optimum length + L0 = (lengthrange[1] - lengthrange[0]) / wp.max(MJ_MINVAL, range_[1] - range_[0]) + + # normalized length and velocity + L = range_[0] + (len - lengthrange[0]) / wp.max(MJ_MINVAL, L0) + V = vel / wp.max(MJ_MINVAL, L0 * vmax) + + # length curve + FL = muscle_gain_length(L, lmin, lmax) + + # velocity curve + y = fvmax - 1.0 + if V <= -1.0: + FV = 0.0 + elif V <= 0.0: + FV = (V + 1.0) * (V + 1.0) + elif V <= y: + FV = fvmax - (y - V) * (y - V) / wp.max(MJ_MINVAL, y) + else: + FV = fvmax + + # compute FVL and scale, make it negative + return -force * FL * FV + + +@wp.func +def muscle_bias(len: float, lengthrange: wp.vec2, acc0: float, prm: vec10) -> float: + """Calculates muscle passive force. + + prm = (range[2], force, scale, lmin, lmax, vmax, fpmax, fvmax).""" + + # unpack parameters + range_ = wp.vec2(prm[0], prm[1]) + force = prm[2] + scale = prm[3] + lmax = prm[5] + fpmax = prm[7] + + # scale force if negative + if force < 0.0: + force = scale / wp.max(MJ_MINVAL, acc0) + + # optimum length + L0 = (lengthrange[1] - lengthrange[0]) / wp.max(MJ_MINVAL, range_[1] - range_[0]) + + # normalized length + L = range_[0] + (len - lengthrange[0]) / wp.max(MJ_MINVAL, L0) + + # half-quadratic to (L0 + lmax) / 2, linear beyond + b = 0.5 * (1.0 + lmax) + if L <= 1.0: + return 0.0 + elif L <= b: + x = (L - 1.0) / wp.max(MJ_MINVAL, b - 1.0) + return -force * fpmax * 0.5 * x * x + else: + x = (L - b) / wp.max(MJ_MINVAL, b - 1.0) + return -force * fpmax * (0.5 + x) + + +@wp.func +def _sigmoid(x: float) -> float: + """Sigmoid function over 0 <= x <= 1 using quintic polynomial.""" + + if x <= 0.0: + return 0.0 + + if x >= 1.0: + return 1.0 + + # sigmoid f(x) = 6 * x^5 - 15 * x^4 + 10 * x^3 + # solution of f(0) = f'(0) = f''(0) = 0, f(1) = 1, f'(1) = f''(1) = 0 + return x * x * x * (3.0 * x * (2.0 * x - 5.0) + 10.0) + + +@wp.func +def muscle_dynamics_timescale(dctrl: float, tau_act: float, tau_deact: float, smooth_width: float) -> float: + """Muscle time constant with optional smoothing.""" + + # hard switching + if smooth_width < MJ_MINVAL: + if dctrl > 0.0: + return tau_act + else: + return tau_deact + else: # smooth switching + # scale by width, center around 0.5 midpoint, rescale to bounds + return tau_deact + (tau_act - tau_deact) * _sigmoid(dctrl / smooth_width + 0.5) + + +@wp.func +def muscle_dynamics(control: float, activation: float, prm: vec10) -> float: + """Muscle activation dynamics, prm = (tau_act, tau_deact, smooth_width).""" + + # clamp control + ctrlclamp = wp.clamp(control, 0.0, 1.0) + + # clamp activation + actclamp = wp.clamp(activation, 0.0, 1.0) + + # compute timescales as in Millard et al. (2013) https://doi.org/10.1115/1.4023390 + tau_act = prm[0] * (0.5 + 1.5 * actclamp) # activation timescale + tau_deact = prm[1] / (0.5 + 1.5 * actclamp) # deactivation timescale + smooth_width = prm[2] # width of smoothing sigmoid + dctrl = ctrlclamp - activation # excess excitation + + tau = muscle_dynamics_timescale(dctrl, tau_act, tau_deact, smooth_width) + + # filter output + return dctrl / wp.max(MJ_MINVAL, tau) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc_test.py new file mode 100644 index 00000000..dbd43212 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc_test.py @@ -0,0 +1,573 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Tests for miscellaneous utilities.""" + +from typing import Tuple + +import numpy as np +import warp as wp +from absl.testing import absltest +from absl.testing import parameterized + +from mujoco.mjx.third_party.mujoco_warp._src import util_misc +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import WrapType +from mujoco.mjx.third_party.mujoco_warp._src.types import vec10 + + +def _assert_eq(a, b, name): + tol = 1e-3 # avoid test noise + err_msg = f"mismatch: {name}" + np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol) + + +def _is_intersect(p1: np.array, p2: np.array, p3: np.array, p4: np.array) -> bool: + intersect = wp.empty(1, dtype=bool) + + @wp.kernel + def is_intersect( + # In: + p1: wp.vec2, + p2: wp.vec2, + p3: wp.vec2, + p4: wp.vec2, + # Out: + intersect_out: wp.array(dtype=bool), + ): + intersect_out[0] = util_misc.is_intersect(p1, p2, p3, p4) + + wp.launch( + is_intersect, + dim=(1,), + inputs=[ + wp.vec2(p1[0], p1[1]), + wp.vec2(p2[0], p2[1]), + wp.vec2(p3[0], p3[1]), + wp.vec2(p4[0], p4[1]), + ], + outputs=[ + intersect, + ], + ) + return intersect.numpy()[0] + + +def _length_circle(p0: np.array, p1: np.array, ind: int, radius: float) -> float: + length = wp.empty(1, dtype=float) + + @wp.kernel + def length_circle( + # In: + p0: wp.vec2, + p1: wp.vec2, + ind: int, + radius: float, + # Out: + length_out: wp.array(dtype=float), + ): + length_out[0] = util_misc.length_circle(p0, p1, ind, radius) + + wp.launch( + length_circle, + dim=(1,), + inputs=[wp.vec2(p0[0], p0[1]), wp.vec2(p1[0], p1[1]), ind, radius], + outputs=[ + length, + ], + ) + return length.numpy()[0] + + +def _wrap_circle(end: np.array, side: np.array, radius: float) -> Tuple[float, np.array, np.array]: + length = wp.empty(1, dtype=float) + wpnt0 = wp.empty(1, dtype=wp.vec2) + wpnt1 = wp.empty(1, dtype=wp.vec2) + + @wp.kernel + def wrap_circle( + # In: + end: wp.vec4, + side: wp.vec2, + radius: float, + # Out: + length_out: wp.array(dtype=float), + wpnt0_out: wp.array(dtype=wp.vec2), + wpnt1_out: wp.array(dtype=wp.vec2), + ): + length_, wpnt0_, wpnt1_ = util_misc.wrap_circle(end, side, radius) + length_out[0] = length_ + wpnt0_out[0] = wpnt0_ + wpnt1_out[0] = wpnt1_ + + wp.launch( + wrap_circle, + dim=(1,), + inputs=[ + wp.vec4(end[0], end[1], end[2], end[3]), + wp.vec2(side[0], side[1]), + radius, + ], + outputs=[ + length, + wpnt0, + wpnt1, + ], + ) + return length.numpy()[0], wpnt0.numpy()[0], wpnt1.numpy()[0] + + +def _wrap_inside(end: np.array, radius: float) -> Tuple[float, np.array, np.array]: + length = wp.empty(1, dtype=float) + wpnt0 = wp.empty(1, dtype=wp.vec2) + wpnt1 = wp.empty(1, dtype=wp.vec2) + + @wp.kernel + def wrap_inside( + # In: + end: wp.vec4, + radius: float, + # Out: + length_out: wp.array(dtype=float), + wpnt0_out: wp.array(dtype=wp.vec2), + wpnt1_out: wp.array(dtype=wp.vec2), + ): + length_, wpnt0_, wpnt1_ = util_misc.wrap_inside(end, radius) + length_out[0] = length_ + wpnt0_out[0] = wpnt0_ + wpnt1_out[0] = wpnt1_ + + wp.launch( + wrap_inside, + dim=(1,), + inputs=[wp.vec4(end[0], end[1], end[2], end[3]), radius], + outputs=[ + length, + wpnt0, + wpnt1, + ], + ) + return length.numpy()[0], wpnt0.numpy()[0], wpnt1.numpy()[0] + + +def _wrap( + x0: np.array, + x1: np.array, + xpos: np.array, + xmat: np.array, + radius: float, + geomtype: int, + side: np.array, +) -> Tuple[float, np.array, np.array]: + length = wp.empty(1, dtype=float) + wpnt0 = wp.empty(1, dtype=wp.vec3) + wpnt1 = wp.empty(1, dtype=wp.vec3) + + @wp.kernel + def wrap( + # In: + x0: wp.vec3, + x1: wp.vec3, + pos: wp.vec3, + mat: wp.mat33, + radius: float, + geomtype: int, + side: wp.vec3, + # Out: + length_out: wp.array(dtype=float), + wpnt0_out: wp.array(dtype=wp.vec3), + wpnt1_out: wp.array(dtype=wp.vec3), + ): + length_, wpnt0_, wpnt1_ = util_misc.wrap(x0, x1, pos, mat, radius, geomtype, side) + length_out[0] = length_ + wpnt0_out[0] = wpnt0_ + wpnt1_out[0] = wpnt1_ + + wp.launch( + wrap, + dim=(1,), + inputs=[ + wp.vec3(x0[0], x0[1], x0[2]), + wp.vec3(x1[0], x1[1], x1[2]), + wp.vec3(xpos[0], xpos[1], xpos[2]), + wp.mat33( + xmat[0, 0], + xmat[0, 1], + xmat[0, 2], + xmat[1, 0], + xmat[1, 1], + xmat[1, 2], + xmat[2, 0], + xmat[2, 1], + xmat[2, 2], + ), + radius, + geomtype, + wp.vec3(side[0], side[1], side[2]), + ], + outputs=[ + length, + wpnt0, + wpnt1, + ], + ) + return length.numpy()[0], wpnt0.numpy()[0], wpnt1.numpy()[0] + + +def _muscle_dynamics_millard(ctrl, act, prm): + """Compute time constant as in Millard et al. (2013) https://doi.org/10.1115/1.4023390.""" + + # clamp control + ctrlclamp = np.clip(ctrl, 0.0, 1.0) + + # clamp activation + actclamp = np.clip(act, 0.0, 1.0) + + if ctrlclamp > act: + tau = prm[0] * (0.5 + 1.5 * actclamp) + else: + tau = prm[1] / (0.5 + 1.5 * actclamp) + + # filter output + return (ctrlclamp - act) / np.maximum(MJ_MINVAL, tau) + + +def _muscle_dynamics(ctrl, act, prm): + @wp.kernel + def muscle_dynamics(control: float, activation: float, prm: vec10, dynamics_out: wp.array(dtype=float)): + dynamics_out[0] = util_misc.muscle_dynamics(control, activation, prm) + + output = wp.empty(1, dtype=float) + wp.launch( + muscle_dynamics, + dim=(1,), + inputs=[ + ctrl, + act, + vec10(prm[0], prm[1], prm[2], 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0), + ], + outputs=[output], + ) + + return output.numpy()[0] + + +def _muscle_gain_length(length, lmin, lmax): + @wp.kernel + def muscle_gain_length(length: float, lmin: float, lmax: float, gain_length_out: wp.array(dtype=float)): + gain_length_out[0] = util_misc.muscle_gain_length(length, lmin, lmax) + + output = wp.empty(1, dtype=float) + wp.launch(muscle_gain_length, dim=(1,), inputs=[length, lmin, lmax], outputs=[output]) + + return output.numpy()[0] + + +def _muscle_dynamics_timescale(dctrl, tau_act, tau_deact, smooth_width): + @wp.kernel + def muscle_gain_length( + dctrl: float, tau_act: float, tau_deact: float, smooth_width: float, dynamics_timescale_out: wp.array(dtype=float) + ): + dynamics_timescale_out[0] = util_misc.muscle_dynamics_timescale(dctrl, tau_act, tau_deact, smooth_width) + + output = wp.empty(1, dtype=float) + wp.launch( + muscle_gain_length, + dim=(1,), + inputs=[dctrl, tau_act, tau_deact, smooth_width], + outputs=[output], + ) + + return output.numpy()[0] + + +class UtilMiscTest(parameterized.TestCase): + def test_is_intersect(self): + self.assertFalse( + _is_intersect( + np.array([0, 0]), + np.array([1, 0]), + np.array([0, 1]), + np.array([1, 1]), + ) + ) + + self.assertTrue( + _is_intersect( + np.array([0, 0]), + np.array([1, 0]), + np.array([0.5, -1]), + np.array([0.5, 1]), + ) + ) + + self.assertFalse( + _is_intersect( + np.array([0, 0]), + np.array([0, 0]), + np.array([0, 0]), + np.array([0, 0]), + ) + ) + + def test_length_circle(self): + _assert_eq( + _length_circle(np.array([0, 1]), np.array([1, 0]), 0, 1.0), + 0.5 * np.pi, + "length_circle", + ) + _assert_eq( + _length_circle(np.array([0, 1]), np.array([1, 0]), 1, 1.0), + 1.5 * np.pi, + "length_circle", + ) + _assert_eq( + _length_circle(np.array([1, 0]), np.array([0, 1]), 0, 1.0), + 1.5 * np.pi, + "length_circle", + ) + _assert_eq( + _length_circle(np.array([1, 0]), np.array([0, 1]), 1, 1.0), + 0.5 * np.pi, + "length_circle", + ) + + def test_wrap_circle(self): + # no wrap + wlen, wpnt0, wpnt1 = _wrap_circle(np.array([1, 0, 0, 1]), np.array([np.inf, np.inf]), 0.1) + _assert_eq(wlen, -1.0, "wlen") + _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0") + _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1") + + # no wrap + wlen, wpnt0, wpnt1 = _wrap_circle(np.array([1, 0, 0, 1]), np.array([0.0, 0.0]), 0.1) + _assert_eq(wlen, -1.0, "wlen") + _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0") + _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1") + + # wrap + wlen, wpnt0, wpnt1 = _wrap_circle(np.array([np.sqrt(2.0), 0, 0, np.sqrt(2.0)]), np.array([np.inf, np.inf]), 1.0 + 5e-4) + _assert_eq(wlen, 0.0, "wlen") + _assert_eq(wpnt0, np.array([np.sqrt(2.0) / 2.0, np.sqrt(2.0) / 2.0]), "wpnt0") + _assert_eq(wpnt1, np.array([np.sqrt(2.0) / 2.0, np.sqrt(2.0) / 2.0]), "wpnt1") + + # wrap + wlen, wpnt0, wpnt1 = _wrap_circle(np.array([np.sqrt(2.0), 0, 0, np.sqrt(2.0)]), np.array([0.0, 0.0]), 1.0 + 5e-4) + _assert_eq(wlen, 0.0, "wlen") + _assert_eq(wpnt0, np.array([np.sqrt(2.0) / 2.0, np.sqrt(2.0) / 2.0]), "wpnt0") + _assert_eq(wpnt1, np.array([np.sqrt(2.0) / 2.0, np.sqrt(2.0) / 2.0]), "wpnt1") + + # wrap + wlen, wpnt0, wpnt1 = _wrap_circle(np.array([1.0, 0, 0, 1.0]), np.array([0.0, 0.0]), 1.0) + _assert_eq(wlen, 0.5 * np.pi, "wlen") + _assert_eq(wpnt0, np.array([1.0, 0.0]), "wpnt0") + _assert_eq(wpnt1, np.array([0.0, 1.0]), "wpnt1") + + # wrap w/ sidesite + wlen, wpnt0, wpnt1 = _wrap_circle(np.array([0, -100, 0, 100]), np.array([0.2, 0.0]), 0.1) + + # wlen, wpnt0[1], wpnt1[1] are ~0 + _assert_eq(wlen, 0.0, "wlen") + _assert_eq(wpnt0, np.array([0.1, 0]), "wpnt0") + _assert_eq(wpnt1, np.array([0.1, 0]), "wpnt1") + + wlen, wpnt0, wpnt1 = _wrap_circle(np.array([0, -100, 0, 100]), np.array([-0.2, 0.0]), 0.1) + + # wlen, wpnt0[1], wpnt1[1] are ~0 + _assert_eq(wlen, 0.0, "wlen") + _assert_eq(wpnt0, np.array([-0.1, 0]), "wpnt0") + _assert_eq(wpnt1, np.array([-0.1, 0]), "wpnt1") + + def test_wrap_inside(self): + wlen, wpnt0, wpnt1 = _wrap_inside(np.array([1, 0, 0, 1]), 0.7071) + _assert_eq(wlen, 0.0, "wlen") + _assert_eq(wpnt0, np.array([0.5, 0.5]), "wpnt0") + _assert_eq(wpnt1, np.array([0.5, 0.5]), "wpnt1") + + wlen, wpnt0, wpnt1 = _wrap_inside(np.array([0, 0, 1, 0]), 1.0) + _assert_eq(wlen, -1.0, "wlen") + _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0") + _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1") + + wlen, wpnt0, wpnt1 = _wrap_inside(np.array([1, 0, 0, 0]), 1.0) + _assert_eq(wlen, -1.0, "wlen") + _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0") + _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1") + + wlen, wpnt0, wpnt1 = _wrap_inside(np.array([0, 0, 0, 0]), 1.0) + _assert_eq(wlen, -1.0, "wlen") + _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0") + _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1") + + wlen, wpnt0, wpnt1 = _wrap_inside(np.array([1, 0, 0, 0]), 2.0) + _assert_eq(wlen, -1.0, "wlen") + _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0") + _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1") + + wlen, wpnt0, wpnt1 = _wrap_inside(np.array([0, 0, 1, 0]), 2.0) + _assert_eq(wlen, -1.0, "wlen") + _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0") + _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1") + + wlen, wpnt0, wpnt1 = _wrap_inside(np.array([1, 1, 1, 1]), 0.1 * MJ_MINVAL) + _assert_eq(wlen, -1.0, "wlen") + _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0") + _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1") + + wlen, wpnt0, wpnt1 = _wrap_inside(np.array([-1, 0, 1, 0]), 0.1) + _assert_eq(wlen, -1.0, "wlen") + _assert_eq(wpnt0, np.array([np.inf, np.inf]), "wpnt0") + _assert_eq(wpnt1, np.array([np.inf, np.inf]), "wpnt1") + + wlen, wpnt0, wpnt1 = _wrap_inside(np.array([-1, 0.2, 1, 0.2]), 0.1) + _assert_eq(wlen, 0.0, "wlen") + _assert_eq(wpnt0, np.array([0, 0.1]), "wpnt0") + _assert_eq(wpnt1, np.array([0, 0.1]), "wpnt1") + + @parameterized.parameters(WrapType.SPHERE, WrapType.CYLINDER) + def test_wrap(self, wraptype): + # no wrap + x0 = np.array([1, 1, 1]) + x1 = np.array([2, 2, 2]) + xpos = np.array([0, 0, 0]) + xmat = np.eye(3) + radius = 0.1 + side = np.array([np.inf, np.inf, np.inf]) + + wlen, wpnt0, wpnt1 = _wrap(x0, x1, xpos, xmat, radius, wraptype, side) + + _assert_eq(wlen, -1.0, "wlen") + _assert_eq(wpnt0, np.array([np.inf, np.inf, np.inf]), "wpnt0") + _assert_eq(wpnt1, np.array([np.inf, np.inf, np.inf]), "wpnt1") + + # wrap + x0 = np.array([0.1, -1.0, 0.0]) + x1 = np.array([0.1, 1.0, 0.0]) + xpos = np.array([0, 0, 0]) + xmat = np.eye(3) + radius = 0.1 + side = np.array([np.inf, np.inf, np.inf]) + + wlen, wpnt0, wpnt1 = _wrap(x0, x1, xpos, xmat, radius, wraptype, side) + + _assert_eq(wlen, 0.0, "wlen") + _assert_eq(wpnt0, np.array([0.1, 0.0, 0.0]), "wpnt0") + _assert_eq(wpnt1, np.array([0.1, 0.0, 0.0]), "wpnt1") + + # outside wrap w/ sidesite + x0 = np.array([MJ_MINVAL, -100.0, 0.0]) + x1 = np.array([MJ_MINVAL, 100.0, 0.0]) + xpos = np.array([0, 0, 0]) + xmat = np.eye(3) + radius = 0.1 + side = np.array([radius + 10 * MJ_MINVAL, 0, 0]) + + wlen, wpnt0, wpnt1 = _wrap(x0, x1, xpos, xmat, radius, wraptype, side) + + # wlen, wpnt0[1], wpnt1[1] are ~0 + _assert_eq(wlen, 0, "wlen") + _assert_eq(wpnt0, np.array([0.1, 0, 0]), "wpnt0") + _assert_eq(wpnt1, np.array([0.1, 0, 0]), "wpnt1") + + # inside no wrap w/ sidesite + x0 = np.array([0.0, -1.0, 0.0]) + x1 = np.array([0.0, 1.0, 0.0]) + xpos = np.array([0, 0, 0]) + xmat = np.eye(3) + radius = 0.1 + wraptype = WrapType.CYLINDER + side = np.array([0, 0, 0]) + + wlen, wpnt0, wpnt1 = _wrap(x0, x1, xpos, xmat, radius, wraptype, side) + + _assert_eq(wlen, -1.0, "wlen") + _assert_eq(wpnt0, np.array([np.inf, np.inf, np.inf]), "wpnt0") + _assert_eq(wpnt1, np.array([np.inf, np.inf, np.inf]), "wpnt1") + + # inside wrap w/ sidesite + x0 = np.array([1.0, -1.0, 0.0]) + x1 = np.array([1.0, 1.0, 0.0]) + xpos = np.array([0, 0, 0]) + xmat = np.eye(3) + radius = 0.1 + side = np.array([0.0, 0.0, 0.0]) + + wlen, wpnt0, wpnt1 = _wrap(x0, x1, xpos, xmat, radius, wraptype, side) + + _assert_eq(wlen, 0.0, "wlen") + _assert_eq(wpnt0, np.array([0.1, 0.0, 0.0]), "wpnt0") + _assert_eq(wpnt1, np.array([0.1, 0.0, 0.0]), "wpnt1") + + @parameterized.product(ctrl=[-0.1, 0.0, 0.4, 0.5, 1.0, 1.1], act=[-0.1, 0.0, 0.4, 0.5, 1.0, 1.1]) + def test_muscle_dynamics_tausmooth0(self, ctrl, act): + # exact equality if tau_smooth = 0 + prm = np.array([0.01, 0.04, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0]) + + actdot_old = _muscle_dynamics_millard(ctrl, act, prm) + actdot_new = _muscle_dynamics(ctrl, act, prm) + + _assert_eq(actdot_new, actdot_old, "actdot") + + def test_muscle_dynamics_tausmooth_positive(self): + # positive tau_smooth + prm = np.array([0.01, 0.04, 0.2, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0]) + act = 0.5 + eps = 1.0e-6 + + ctrl = 0.4 - eps # smaller than act by just over 0.5 * tau_smooth + _assert_eq( + _muscle_dynamics(ctrl, act, prm), + _muscle_dynamics_millard(ctrl, act, prm), + "actdot", + ) + + ctrl = 0.6 + eps # larger than act by just over 0.5 * tau_smooth + _assert_eq( + _muscle_dynamics(ctrl, act, prm), + _muscle_dynamics_millard(ctrl, act, prm), + "actdot", + ) + + @parameterized.parameters(0.0, 0.1, 0.2, 1.0, 1.1) + def test_muscle_dynamics_timescale(self, dctrl): + # right in the middle should give average of time constants + tau_smooth = 0.2 + tau_act = 0.2 + tau_deact = 0.3 + + lower = _muscle_dynamics_timescale(-dctrl, tau_act, tau_deact, tau_smooth) + upper = _muscle_dynamics_timescale(dctrl, tau_act, tau_deact, tau_smooth) + + _assert_eq(0.5 * (lower + upper), 0.5 * (tau_act + tau_deact), "muscle_dynamics_timescale") + + @parameterized.parameters( + (0.0, 0.0), + (0.5, 0.0), + (0.75, 0.5), + (1.0, 1.0), + (1.25, 0.5), + (1.5, 0.0), + (2.0, 0.0), + ) + def test_muscle_gain_length(self, input, output): + _assert_eq(_muscle_gain_length(input, 0.5, 1.5), output, "length-gain") + + # TODO(team): test util_misc.muscle_gain + # TODO(team): test util_misc.muscle_bias + + +if __name__ == "__main__": + wp.init() + absltest.main() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/warp_util.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/warp_util.py new file mode 100644 index 00000000..572d2312 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/warp_util.py @@ -0,0 +1,196 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import functools +from typing import Callable, Optional + +import warp as wp +from warp.context import Module +from warp.context import assert_conditional_graph_support +from warp.context import get_module + +_STACK = None + + +class EventTracer: + """Calculates elapsed times of functions annotated with `event_scope`. + + Use as a context manager like so: + + @event_trace + def my_warp_function(...): + ... + + with EventTracer() as tracer: + my_warp_function(...) + print(tracer.trace()) + """ + + def __init__(self, enabled: bool = True): + global _STACK + if _STACK is not None: + raise ValueError("only one EventTracer can run at a time") + if enabled: + _STACK = {} + + def __enter__(self): + return self + + def trace(self) -> dict: + """Calculates elapsed times for every node of the trace.""" + global _STACK + + if _STACK is None: + return {} + + ret = {} + + for k, v in _STACK.items(): + events, sub_stack = v + # push into next level of stack + saved_stack, _STACK = _STACK, sub_stack + sub_trace = self.trace() + # pop! + _STACK = saved_stack + events = tuple(wp.get_event_elapsed_time(beg, end) for beg, end in events) + ret[k] = (events, sub_trace) + + return ret + + def __exit__(self, type, value, traceback): + global _STACK + _STACK = None + + +def _merge(a: dict, b: dict) -> dict: + """Merges two event trace stacks.""" + ret = {} + if not a or not b: + return dict(**a, **b) + if set(a) != set(b): + raise ValueError("incompatible stacks") + for key in a: + a1_events, a1_substack = a[key] + a2_events, a2_substack = b[key] + ret[key] = (a1_events + a2_events, _merge(a1_substack, a2_substack)) + return ret + + +def event_scope(fn, name: str = ""): + """Wraps a function and records an event before and after the function invocation.""" + name = name or getattr(fn, "__name__") + + @functools.wraps(fn) + def wrapper(*args, **kwargs): + global _STACK + if _STACK is None: + return fn(*args, **kwargs) + # push into next level of stack + saved_stack, _STACK = _STACK, {} + beg = wp.Event(enable_timing=True) + end = wp.Event(enable_timing=True) + wp.record_event(beg) + res = fn(*args, **kwargs) + wp.record_event(end) + # pop back up to current level + sub_stack, _STACK = _STACK, saved_stack + # append events and substack + prev_events, prev_substack = _STACK.get(name, ((), {})) + events = prev_events + ((beg, end),) + sub_stack = _merge(prev_substack, sub_stack) + _STACK[name] = (events, sub_stack) + return res + + return wrapper + + +# @kernel decorator to automatically set up modules based on nested +# function names +def kernel( + f: Optional[Callable] = None, + *, + enable_backward: Optional[bool] = None, + module: Optional[Module] = None, +): + """ + Decorator to register a Warp kernel from a Python function. + The function must be defined with type annotations for all arguments. + The function must not return anything. + + Example:: + + @kernel + def my_kernel(a: wp.array(dtype=float), b: wp.array(dtype=float)): + tid = wp.tid() + b[tid] = a[tid] + 1.0 + + + @kernel(enable_backward=False) + def my_kernel_no_backward(a: wp.array(dtype=float, ndim=2), x: float): + # the backward pass will not be generated + i, j = wp.tid() + a[i, j] = x + + + @kernel(module="unique") + def my_kernel_unique_module(a: wp.array(dtype=float), b: wp.array(dtype=float)): + # the kernel will be registered in new unique module created just for this + # kernel and its dependent functions and structs + tid = wp.tid() + b[tid] = a[tid] + 1.0 + + Args: + f: The function to be registered as a kernel. + enable_backward: If False, the backward pass will not be generated. + module: The :class:`warp.context.Module` to which the kernel belongs. Alternatively, + if a string `"unique"` is provided, the kernel is assigned to a new module + named after the kernel name and hash. If None, the module is inferred from + the function's module. + + Returns: + The registered kernel. + """ + if module is None: + # create a module name based on the name of the nested function + # get the qualified name, e.g. "main..nested_kernel" + qualname = f.__qualname__ + parts = [part for part in qualname.split(".") if part != ""] + outer_functions = parts[:-1] + module = get_module(".".join([f.__module__] + outer_functions)) + + return wp.kernel(f, enable_backward=enable_backward, module=module) + + +_KERNEL_CACHE = {} + + +def cache_kernel(func): + # caching kernels to avoid crashes in graph_conditional code + @functools.wraps(func) + def wrapper(*args): + key = tuple(a.size if hasattr(a, "size") else hash(a) for a in args) + (hash(func.__name__),) + if key not in _KERNEL_CACHE: + _KERNEL_CACHE[key] = func(*args) + return _KERNEL_CACHE[key] + + return wrapper + + +def conditional_graph_supported(): + try: + assert_conditional_graph_support() + except Exception: + return False + return True diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuation.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuation.xml new file mode 100644 index 00000000..2e532171 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuation.xml @@ -0,0 +1,40 @@ + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuators.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuators.xml new file mode 100644 index 00000000..24c86ae9 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/actuators.xml @@ -0,0 +1,28 @@ + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/adhesion.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/adhesion.xml new file mode 100644 index 00000000..fb891b39 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/adhesion.xml @@ -0,0 +1,21 @@ + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/muscle.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/muscle.xml new file mode 100644 index 00000000..b07e7f1a --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/muscle.xml @@ -0,0 +1,22 @@ + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/position.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/position.xml new file mode 100644 index 00000000..16e9d7a6 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/position.xml @@ -0,0 +1,40 @@ + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/site.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/site.xml new file mode 100644 index 00000000..fcff6eea --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/site.xml @@ -0,0 +1,31 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/slidercrank.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/slidercrank.xml new file mode 100644 index 00000000..75a34153 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/slidercrank.xml @@ -0,0 +1,15 @@ + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/tendon_force_limit.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/tendon_force_limit.xml new file mode 100644 index 00000000..c7c0ed8b --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/actuation/tendon_force_limit.xml @@ -0,0 +1,26 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision.xml new file mode 100644 index 00000000..5cc041e2 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision.xml @@ -0,0 +1,43 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/bolt.py b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/bolt.py new file mode 100644 index 00000000..07a99a1e --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/bolt.py @@ -0,0 +1,70 @@ +import warp as wp + + +@wp.func +def Fract(x: float) -> float: + return x - wp.floor(x) + + +@wp.func +def Subtraction(a: float, b: float) -> float: + return wp.max(a, -b) + + +@wp.func +def Union(a: float, b: float) -> float: + return wp.min(a, b) + + +@wp.func +def Intersection(a: float, b: float) -> float: + return wp.max(a, b) + + +@wp.func +def bolt(p: wp.vec3, attr: wp.vec3) -> float: + screw = 12.0 + radius = wp.sqrt(p[0] * p[0] + p[1] * p[1]) - attr[0] + sqrt12 = wp.sqrt(2.0) / 2.0 + + azimuth = wp.atan2(p[1], p[0]) + triangle = wp.abs(Fract(p[2] * screw - azimuth / wp.pi / 2.0) - 0.5) + thread = (radius - triangle / screw) * sqrt12 + + bolt_val = Subtraction(thread, 0.5 - wp.abs(p[2] + 0.5)) + cone = (p[2] - radius) * sqrt12 + + bolt_val = Subtraction(bolt_val, cone + 1.0 * sqrt12) + + point2D = wp.vec2(p[0], p[1]) + k = 6.0 / wp.pi / 2.0 + angle = -wp.floor((wp.atan2(point2D[1], point2D[0])) * k + 0.5) / k + s = wp.vec2(wp.sin(angle), wp.sin(angle + wp.pi * 0.5)) + + res = wp.vec2(s[1] * point2D[0] - s[0] * point2D[1], s[0] * point2D[0] + s[1] * point2D[1]) + point3D = wp.vec3(res[0], res[1], p[2]) + head = point3D[0] - 0.5 + + head = Intersection(head, wp.abs(point3D[2] + 0.25) - 0.25) + head = Intersection(head, (point3D[2] + radius - 0.22) * sqrt12) + + return Union(bolt_val, head) + + +@wp.func +def bolt_sdf_grad(p: wp.vec3, attr: wp.vec3) -> wp.vec3: + grad = wp.vec3() + eps = 1e-6 + f_original = bolt(p, attr) + x_plus = wp.vec3(p[0] + eps, p[1], p[2]) + f_plus = bolt(x_plus, attr) + grad[0] = (f_plus - f_original) / eps + + x_plus = wp.vec3(p[0], p[1] + eps, p[2]) + f_plus = bolt(x_plus, attr) + grad[1] = (f_plus - f_original) / eps + + x_plus = wp.vec3(p[0], p[1], p[2] + eps) + f_plus = bolt(x_plus, attr) + grad[2] = (f_plus - f_original) / eps + return grad diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nut.py b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nut.py new file mode 100644 index 00000000..47b8cc73 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nut.py @@ -0,0 +1,65 @@ +import warp as wp + + +@wp.func +def Fract(x: float) -> float: + return x - wp.floor(x) + + +@wp.func +def Subtraction(a: float, b: float) -> float: + return wp.max(a, -b) + + +@wp.func +def Union(a: float, b: float) -> float: + return wp.min(a, b) + + +@wp.func +def Intersection(a: float, b: float) -> float: + return wp.max(a, b) + + +@wp.func +def nut(p: wp.vec3, attr: wp.vec3) -> float: + screw = 12.0 + radius2 = wp.sqrt(p[0] * p[0] + p[1] * p[1]) - attr[0] + sqrt12 = wp.sqrt(2.0) / 2.0 + azimuth = wp.atan2(p[1], p[0]) + triangle = wp.abs(Fract(p[2] * screw - azimuth / (wp.pi * 2.0)) - 0.5) + thread2 = (radius2 - triangle / screw) * sqrt12 + cone2 = (p[2] - radius2) * sqrt12 + hole = Subtraction(thread2, cone2 + 0.5 * sqrt12) + hole = Union(hole, -cone2 - 0.05 * sqrt12) + k = 6.0 / wp.pi / 2.0 + angle = -wp.floor((wp.atan2(p[1], p[0])) * k + 0.5) / k + s0 = wp.sin(angle) + s1 = wp.sin(angle + wp.pi * 0.5) + res0 = s1 * p[0] - s0 * p[1] + res1 = s0 * p[0] + s1 * p[1] + point3D0 = res0 + point3D2 = p[2] + head = point3D0 - 0.5 + head = Intersection(head, wp.abs(point3D2 + 0.25) - 0.25) + head = Intersection(head, (point3D2 + radius2 - 0.22) * sqrt12) + return Subtraction(head, hole) + + +@wp.func +def nut_sdf_grad(p: wp.vec3, attr: wp.vec3) -> wp.vec3: + grad = wp.vec3() + eps = 1e-6 + f_original = nut(p, attr) + x_plus = wp.vec3(p[0] + eps, p[1], p[2]) + f_plus = nut(x_plus, attr) + grad[0] = (f_plus - f_original) / eps + + x_plus = wp.vec3(p[0], p[1] + eps, p[2]) + f_plus = nut(x_plus, attr) + grad[1] = (f_plus - f_original) / eps + + x_plus = wp.vec3(p[0], p[1], p[2] + eps) + f_plus = nut(x_plus, attr) + grad[2] = (f_plus - f_original) / eps + return grad diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nutbolt.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nutbolt.xml new file mode 100644 index 00000000..d0fa2d26 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/nutbolt.xml @@ -0,0 +1,56 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/scene.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/scene.xml new file mode 100644 index 00000000..1591519e --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/scene.xml @@ -0,0 +1,38 @@ + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/utils.py b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/utils.py new file mode 100644 index 00000000..76a12f89 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/collision_sdf/utils.py @@ -0,0 +1,86 @@ +# Copyright 2025 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Utility functions for SDF collision handling.""" + +import enum +from typing import Dict + +import mujoco +import warp as wp + +from .bolt import bolt +from .bolt import bolt_sdf_grad +from .nut import nut +from .nut import nut_sdf_grad + + +class SDFType(enum.Enum): + """Enum for SDF types.""" + + NUT = "NUT" + BOLT = "BOLT" + + +def register_sdf_plugins(collision_sdf) -> Dict[str, int]: + xml = """ + + + + + + + + + + + + + """ + + try: + m = mujoco.MjModel.from_xml_string(xml) + except Exception as e: + raise ValueError(f"Failed to create MuJoCo model from XML: {e}") + + sdf_types = {} + + for i in range(m.ngeom): + name = mujoco.mj_id2name(m, mujoco.mjtObj.mjOBJ_GEOM, i) + if name == "ng": + sdf_types[SDFType.NUT.value] = int(m.plugin[i]) + elif name == "bg": + sdf_types[SDFType.BOLT.value] = int(m.plugin[i]) + + @wp.func + def user_sdf(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> float: + result = 0.0 + if sdf_type == wp.static(sdf_types[SDFType.NUT.value]): + result = nut(p, attr) + elif sdf_type == wp.static(sdf_types[SDFType.BOLT.value]): + result = bolt(p, attr) + return result + + @wp.func + def user_sdf_grad(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> wp.vec3: + if sdf_type == wp.static(sdf_types[SDFType.NUT.value]): + return nut_sdf_grad(p, attr) + elif sdf_type == wp.static(sdf_types[SDFType.BOLT.value]): + return bolt_sdf_grad(p, attr) + return wp.vec3() + + collision_sdf.user_sdf = user_sdf + collision_sdf.user_sdf_grad = user_sdf_grad + + return sdf_types diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/constraints.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/constraints.xml new file mode 100644 index 00000000..8e5be10f --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/constraints.xml @@ -0,0 +1,114 @@ + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/cloth.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/cloth.xml new file mode 100644 index 00000000..89e0adfa --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/cloth.xml @@ -0,0 +1,45 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/floppy.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/floppy.xml new file mode 100644 index 00000000..6430e70d --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/floppy.xml @@ -0,0 +1,45 @@ + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/mannequin.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/mannequin.xml new file mode 100644 index 00000000..9237b943 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/mannequin.xml @@ -0,0 +1,174 @@ + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/scene.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/scene.xml new file mode 100644 index 00000000..ef314930 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/scene.xml @@ -0,0 +1,41 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/hfield/hfield.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/hfield/hfield.xml new file mode 100644 index 00000000..27d179ee --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/hfield/hfield.xml @@ -0,0 +1,84 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/humanoid/humanoid.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/humanoid/humanoid.xml new file mode 100644 index 00000000..196aa5c9 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/humanoid/humanoid.xml @@ -0,0 +1,252 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/meshes/dodecahedron.stl b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/meshes/dodecahedron.stl new file mode 100644 index 0000000000000000000000000000000000000000..1a3f9f69a16a7e61769680a8489d155d6f456f67 GIT binary patch literal 1884 zcmb7EJ#Q0H5Zo#y9len#LXilOB8rgWGl-H8ku)jNAUjpQ#<$ zRCKiYFOk^U+j*X6B}FV*cX~TBJF|OtufCk0%|FhjqoeuR$>_!L>~uPtZ#>)F-Wjc5 zeEKkY`o!+_dz|R;)xr4vx3%%t@~h+j_u;f#e113q#Lo)P2g*xlydD@$g)E|R9T>aL z*mDLEIQ9E-dk;DSO0|fFikufSI$B{A?I(=RjE=xiUAeE>oLS5;BcS1~)ru@JuRIT? z8D<0q+-XI6XpYVdGXewda8>`Dmns1>IC6KrqQxQHFiQzS=c<&|JJjs~Z`^XQQ~4>TBF4M*tNDow5q`ZvD!KYzb3yjnFA6=-G0bo+PCc>d_e zmR5v;DO9DHZtqzlwV;5YE4`0=L!bOIAV*|L-PP$P!&CS9Wnjvv&`qj@l8?}TKG{CM z42=UibA~9%JIE&yS_MjUmW*IE?;sgbas3Hd(!-VH9gIyEt9)aW#N43Oe=TvD`uP%oPkZG_lr59uPT}S2aFQw A9{>OV literal 0 HcmV?d00001 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/meshes/tetrahedron.stl b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/meshes/tetrahedron.stl new file mode 100644 index 0000000000000000000000000000000000000000..68fed6b855a9e18a2535f742c8fbca17f0288d35 GIT binary patch literal 284 zcmZQzpe|s68mIR&#a{2{lYI;f4f`QNaM~Uy2E`y5m^e_!eoP%8mB^|d4gjoMDmefE literal 0 HcmV?d00001 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/pendula.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/pendula.xml new file mode 100644 index 00000000..b3ea0453 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/pendula.xml @@ -0,0 +1,167 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/ray.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/ray.xml new file mode 100644 index 00000000..46bc3070 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/ray.xml @@ -0,0 +1,21 @@ + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/armature.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/armature.xml new file mode 100644 index 00000000..6ab99fda --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/armature.xml @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/damping.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/damping.xml new file mode 100644 index 00000000..e4a6d380 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/damping.xml @@ -0,0 +1,20 @@ + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed.xml new file mode 100644 index 00000000..42b6a164 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed.xml @@ -0,0 +1,29 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed_site.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed_site.xml new file mode 100644 index 00000000..261b900b --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/fixed_site.xml @@ -0,0 +1,35 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_fixed_site.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_fixed_site.xml new file mode 100644 index 00000000..25b3cab5 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_fixed_site.xml @@ -0,0 +1,36 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site.xml new file mode 100644 index 00000000..0b4b351c --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site.xml @@ -0,0 +1,46 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site_fixed.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site_fixed.xml new file mode 100644 index 00000000..8d7bf020 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_site_fixed.xml @@ -0,0 +1,32 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_wrap.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_wrap.xml new file mode 100644 index 00000000..1dff6600 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/pulley_wrap.xml @@ -0,0 +1,42 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site.xml new file mode 100644 index 00000000..da50fe17 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site.xml @@ -0,0 +1,44 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site_fixed.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site_fixed.xml new file mode 100644 index 00000000..af85b662 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/site_fixed.xml @@ -0,0 +1,31 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/tendon_limit.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/tendon_limit.xml new file mode 100644 index 00000000..e6f243c8 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/tendon_limit.xml @@ -0,0 +1,29 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/wrap.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/wrap.xml new file mode 100644 index 00000000..ae9d19a1 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/tendon/wrap.xml @@ -0,0 +1,40 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py new file mode 100644 index 00000000..89044207 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/__init__.py @@ -0,0 +1,16 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from .custom_call import jax_kernel diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py new file mode 100644 index 00000000..97faaec5 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/custom_call.py @@ -0,0 +1,363 @@ +# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import ctypes + +import warp as wp +from warp.context import type_str +from warp.jax import get_jax_device +from warp.types import array_t, launch_bounds_t, strides_from_shape + +_jax_warp_p = None + +# Holder for the custom callback to keep it alive. +_cc_callback = None +_registered_kernels = [None] +_registered_kernel_to_id = {} + + +def jax_kernel(kernel, launch_dims=None): + """Create a Jax primitive from a Warp kernel. + + NOTE: This is an experimental feature under development. + + Args: + kernel: The Warp kernel to be wrapped. + launch_dims: Optional. Specify the kernel launch dimensions. If None, + dimensions are inferred from the shape of the first argument. + This option when set will specify the output dimensions. + + Limitations: + - All kernel arguments must be contiguous arrays. + - Input arguments are followed by output arguments in the Warp kernel definition. + - There must be at least one input argument and at least one output argument. + - Only the CUDA backend is supported. + """ + + if _jax_warp_p is None: + # Create and register the primitive + _create_jax_warp_primitive() + if kernel not in _registered_kernel_to_id: + id = len(_registered_kernels) + _registered_kernels.append(kernel) + _registered_kernel_to_id[kernel] = id + else: + id = _registered_kernel_to_id[kernel] + + def bind(*args): + return _jax_warp_p.bind(*args, kernel=id, launch_dims=launch_dims) + + return bind + + +def _warp_custom_callback(stream, buffers, opaque, opaque_len): + # The descriptor is the form + # || + # Example: 42|16,32|16,32;100;16,32 + kernel_id_str, dim_str, args_str = opaque.decode().split("|") + + # Get the kernel from the registry. + kernel_id = int(kernel_id_str) + kernel = _registered_kernels[kernel_id] + + # Parse launch dimensions. + dims = [int(d) for d in dim_str.split(",")] + bounds = launch_bounds_t(dims) + + # Parse arguments. + arg_strings = args_str.split(";") + num_args = len(arg_strings) + assert num_args == len(kernel.adj.args), "Incorrect number of arguments" + + # First param is the launch bounds. + kernel_params = (ctypes.c_void_p * (1 + num_args))() + kernel_params[0] = ctypes.addressof(bounds) + + # Parse array descriptors. + args = [] + for i in range(num_args): + dtype = kernel.adj.args[i].type.dtype + shape = [int(d) for d in arg_strings[i].split(",")] + strides = strides_from_shape(shape, dtype) + + arr = array_t(buffers[i], 0, len(shape), shape, strides) + args.append(arr) # keep a reference + arg_ptr = ctypes.addressof(arr) + + kernel_params[i + 1] = arg_ptr + + # Get current device. + device = wp.device_from_jax(get_jax_device()) + + # Get kernel hooks. + # Note: module was loaded during jit lowering. + hooks = kernel.module.get_kernel_hooks(kernel, device) + assert hooks.forward, "Failed to find kernel entry point" + + # Launch the kernel. + wp.context.runtime.core.wp_cuda_launch_kernel( + device.context, hooks.forward, bounds.size, 0, 256, hooks.forward_smem_bytes, kernel_params, stream + ) + + +def _create_jax_warp_primitive(): + from functools import reduce + + import jax + from jax._src.interpreters import batching + from jax.interpreters import mlir + from jax.interpreters.mlir import ir + from jax.jaxlib.hlo_helpers import custom_call + + global _jax_warp_p + global _cc_callback + + # Create and register the primitive. + # TODO add default implementation that calls the kernel via warp. + try: + # newer JAX versions + import jax.extend + + _jax_warp_p = jax.extend.core.Primitive("jax_warp") + except (ImportError, AttributeError): + # older JAX versions + _jax_warp_p = jax.core.Primitive("jax_warp") + _jax_warp_p.multiple_results = True + + # TODO Just launch the kernel directly, but make sure the argument + # shapes are massaged the same way as below so that vmap works. + def impl(*args): + raise Exception("Not implemented") + + _jax_warp_p.def_impl(impl) + + # Auto-batching. Make sure all the arguments are fully broadcasted + # so that Warp is not confused about dimensions. + def vectorized_multi_batcher(args, dims, **params): + # Figure out the number of outputs. + wp_kernel = _registered_kernels[params["kernel"]] + output_count = len(wp_kernel.adj.args) - len(args) + shape, dim = next((a.shape, d) for a, d in zip(args, dims) if d is not None) + size = shape[dim] + args = [batching.bdim_at_front(a, d, size) if len(a.shape) else a for a, d in zip(args, dims)] + # Create the batched primitive. + return _jax_warp_p.bind(*args, **params), [dims[0]] * output_count + + batching.primitive_batchers[_jax_warp_p] = vectorized_multi_batcher + + def get_vecmat_shape(warp_type): + if hasattr(warp_type.dtype, "_shape_"): + return warp_type.dtype._shape_ + return [] + + def strip_vecmat_dimensions(warp_arg, actual_shape): + shape = get_vecmat_shape(warp_arg.type) + for i, s in enumerate(reversed(shape)): + item = actual_shape[-i - 1] + if s != item: + raise Exception(f"The vector/matrix shape for argument {warp_arg.label} does not match") + return actual_shape[: len(actual_shape) - len(shape)] + + def collapse_into_leading_dimension(warp_arg, actual_shape): + if len(actual_shape) < warp_arg.type.ndim: + raise Exception(f"Argument {warp_arg.label} has too few non-matrix/vector dimensions") + index_rest = len(actual_shape) - warp_arg.type.ndim + 1 + leading_size = reduce(lambda x, y: x * y, actual_shape[:index_rest]) + return [leading_size] + actual_shape[index_rest:] + + # Infer array dimensions from input type. + def infer_dimensions(warp_arg, actual_shape): + actual_shape = strip_vecmat_dimensions(warp_arg, actual_shape) + return collapse_into_leading_dimension(warp_arg, actual_shape) + + def base_type_to_jax(warp_dtype): + if hasattr(warp_dtype, "_wp_scalar_type_"): + return wp.dtype_to_jax(warp_dtype._wp_scalar_type_) + return wp.dtype_to_jax(warp_dtype) + + def base_type_to_jax_ir(warp_dtype): + warp_to_jax_dict = { + wp.float16: ir.F16Type.get(), + wp.float32: ir.F32Type.get(), + wp.float64: ir.F64Type.get(), + wp.int8: ir.IntegerType.get_signless(8), + wp.int16: ir.IntegerType.get_signless(16), + wp.int32: ir.IntegerType.get_signless(32), + wp.int64: ir.IntegerType.get_signless(64), + wp.uint8: ir.IntegerType.get_unsigned(8), + wp.uint16: ir.IntegerType.get_unsigned(16), + wp.uint32: ir.IntegerType.get_unsigned(32), + wp.uint64: ir.IntegerType.get_unsigned(64), + } + if hasattr(warp_dtype, "_wp_scalar_type_"): + warp_dtype = warp_dtype._wp_scalar_type_ + jax_dtype = warp_to_jax_dict.get(warp_dtype) + if jax_dtype is None: + raise TypeError(f"Invalid or unsupported data type: {warp_dtype}") + return jax_dtype + + def base_type_is_compatible(warp_type, jax_ir_type): + jax_ir_to_warp = { + "f16": wp.float16, + "f32": wp.float32, + "f64": wp.float64, + "i8": wp.int8, + "i16": wp.int16, + "i32": wp.int32, + "i64": wp.int64, + "ui8": wp.uint8, + "ui16": wp.uint16, + "ui32": wp.uint32, + "ui64": wp.uint64, + } + expected_warp_type = jax_ir_to_warp.get(str(jax_ir_type)) + if expected_warp_type is not None: + if hasattr(warp_type, "_wp_scalar_type_"): + return warp_type._wp_scalar_type_ == expected_warp_type + else: + return warp_type == expected_warp_type + else: + raise TypeError(f"Invalid or unsupported data type: {jax_ir_type}") + + # Abstract evaluation. + def jax_warp_abstract(*args, kernel=None, launch_dims=None): + wp_kernel = _registered_kernels[kernel] + # All the extra arguments to the warp kernel are outputs. + warp_outputs = [o.type for o in wp_kernel.adj.args[len(args) :]] + + if launch_dims is None: + # Use the first input dimension to infer the output's dimensions if launch_dims is not provided + dims = strip_vecmat_dimensions(wp_kernel.adj.args[0], list(args[0].shape)) + else: + dims = launch_dims + + jax_outputs = [] + for o in warp_outputs: + shape = list(dims) + list(get_vecmat_shape(o)) + dtype = base_type_to_jax(o.dtype) + jax_outputs.append(jax.core.ShapedArray(shape, dtype)) + return jax_outputs + + _jax_warp_p.def_abstract_eval(jax_warp_abstract) + + # Lowering to MLIR. + + # Create python-land custom call target. + CCALLFUNC = ctypes.CFUNCTYPE( + ctypes.c_voidp, ctypes.c_void_p, ctypes.POINTER(ctypes.c_void_p), ctypes.c_char_p, ctypes.c_size_t + ) + _cc_callback = CCALLFUNC(_warp_custom_callback) + ccall_address = ctypes.cast(_cc_callback, ctypes.c_void_p) + + # Put the custom call into a capsule, as required by XLA. + PyCapsule_Destructor = ctypes.CFUNCTYPE(None, ctypes.py_object) + PyCapsule_New = ctypes.pythonapi.PyCapsule_New + PyCapsule_New.restype = ctypes.py_object + PyCapsule_New.argtypes = (ctypes.c_void_p, ctypes.c_char_p, PyCapsule_Destructor) + capsule = PyCapsule_New(ccall_address.value, b"xla._CUSTOM_CALL_TARGET", PyCapsule_Destructor(0)) + + # Register the callback in XLA. + try: + # newer JAX versions + jax.ffi.register_ffi_target("warp_call", capsule, platform="gpu", api_version=0) + except AttributeError: + # older JAX versions + jax.lib.xla_client.register_custom_call_target("warp_call", capsule, platform="gpu") + + def default_layout(shape): + return range(len(shape) - 1, -1, -1) + + def warp_call_lowering(ctx, *args, kernel=None, launch_dims=None): + if not kernel: + raise Exception("Unknown kernel id " + str(kernel)) + wp_kernel = _registered_kernels[kernel] + + # TODO This may not be necessary, but it is perhaps better not to be + # mucking with kernel loading while already running the workload. + module = wp_kernel.module + device = wp.device_from_jax(get_jax_device()) + if not module.load(device): + raise Exception("Could not load kernel on device") + + if launch_dims is None: + # Infer dimensions from the first input. + warp_arg0 = wp_kernel.adj.args[0] + actual_shape0 = ir.RankedTensorType(args[0].type).shape + dims = strip_vecmat_dimensions(warp_arg0, actual_shape0) + warp_dims = collapse_into_leading_dimension(warp_arg0, dims) + else: + dims = launch_dims + warp_dims = launch_dims + # Figure out the types and shapes of the input arrays. + arg_strings = [] + operand_layouts = [] + for actual, warg in zip(args, wp_kernel.adj.args): + wtype = warg.type + rtt = ir.RankedTensorType(actual.type) + + if not isinstance(wtype, wp.array): + raise Exception("Only contiguous arrays are supported for Jax kernel arguments") + + if not base_type_is_compatible(wtype.dtype, rtt.element_type): + raise TypeError( + f"Incompatible data type for argument '{warg.label}', expected {type_str(wtype.dtype)}, got {rtt.element_type}" + ) + + # Infer array dimension (by removing the vector/matrix dimensions and + # collapsing the initial dimensions). + shape = infer_dimensions(warg, rtt.shape) + + if len(shape) != wtype.ndim: + raise TypeError(f"Incompatible array dimensionality for argument '{warg.label}'") + + arg_strings.append(",".join([str(d) for d in shape])) + operand_layouts.append(default_layout(rtt.shape)) + + # Figure out the types and shapes of the output arrays. + result_types = [] + result_layouts = [] + for warg in wp_kernel.adj.args[len(args) :]: + wtype = warg.type + + if not isinstance(wtype, wp.array): + raise Exception("Only contiguous arrays are supported for Jax kernel arguments") + + # Infer dimensions from the first input. + arg_strings.append(",".join([str(d) for d in warp_dims])) + + result_shape = list(dims) + list(get_vecmat_shape(wtype)) + result_types.append(ir.RankedTensorType.get(result_shape, base_type_to_jax_ir(wtype.dtype))) + result_layouts.append(default_layout(result_shape)) + + # Build opaque descriptor for callback. + shape_str = ",".join([str(d) for d in warp_dims]) + args_str = ";".join(arg_strings) + descriptor = f"{kernel}|{shape_str}|{args_str}" + + out = custom_call( + b"warp_call", + result_types=result_types, + operands=args, + backend_config=descriptor.encode("utf-8"), + operand_layouts=operand_layouts, + result_layouts=result_layouts, + ).results + return out + + mlir.register_lowering( + _jax_warp_p, + warp_call_lowering, + platform="gpu", + ) diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py new file mode 100644 index 00000000..82b6a8b8 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/ffi.py @@ -0,0 +1,804 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import ctypes +import threading +import traceback +from typing import Callable, Optional + +import jax + +import warp as wp +from warp.codegen import get_full_arg_spec, make_full_qualified_name +from warp.jax import get_jax_device +from warp.types import array_t, launch_bounds_t, strides_from_shape, type_to_warp + +from .xla_ffi import * + + +class FfiArg: + def __init__(self, name, type, in_out=False): + self.name = name + self.type = type + self.in_out = in_out + self.is_array = isinstance(type, wp.array) + + if self.is_array: + if hasattr(type.dtype, "_wp_scalar_type_"): + self.dtype_shape = type.dtype._shape_ + self.dtype_ndim = len(self.dtype_shape) + self.jax_scalar_type = wp.dtype_to_jax(type.dtype._wp_scalar_type_) + self.jax_ndim = type.ndim + self.dtype_ndim + elif type.dtype in wp.types.value_types: + self.dtype_ndim = 0 + self.dtype_shape = () + self.jax_scalar_type = wp.dtype_to_jax(type.dtype) + self.jax_ndim = type.ndim + else: + raise TypeError(f"Invalid data type for array argument '{name}', expected scalar, vector, or matrix") + self.warp_ndim = type.ndim + elif type in wp.types.value_types: + self.dtype_ndim = 0 + self.dtype_shape = () + self.jax_scalar_type = wp.dtype_to_jax(type_to_warp(type)) + self.jax_ndim = 0 + self.warp_ndim = 0 + else: + raise TypeError(f"Invalid type for argument '{name}', expected array or scalar, got {type}") + + +class FfiLaunchDesc: + def __init__(self, static_inputs, launch_dims): + self.static_inputs = static_inputs + self.launch_dims = launch_dims + + +class FfiKernel: + def __init__(self, kernel, num_outputs, vmap_method, launch_dims, output_dims, in_out_argnames): + self.kernel = kernel + self.name = generate_unique_name(kernel.func) + self.num_outputs = num_outputs + self.vmap_method = vmap_method + self.launch_dims = launch_dims + self.output_dims = output_dims + self.first_array_arg = None + self.launch_id = 0 + self.launch_descriptors = {} + + in_out_argnames_list = in_out_argnames or [] + in_out_argnames = set(in_out_argnames_list) + if len(in_out_argnames_list) != len(in_out_argnames): + raise AssertionError("in_out_argnames must not contain duplicate names") + + self.num_kernel_args = len(kernel.adj.args) + self.num_in_out = len(in_out_argnames) + self.num_inputs = self.num_kernel_args - num_outputs + self.num_in_out + if self.num_outputs < 1: + raise ValueError("At least one output is required") + if self.num_outputs > self.num_kernel_args: + raise ValueError("Number of outputs cannot be greater than the number of kernel arguments") + if self.num_outputs < self.num_in_out: + raise ValueError("Number of outputs cannot be smaller than the number of in_out_argnames") + + # process input args + self.input_args = [] + for i in range(self.num_inputs): + arg_name = kernel.adj.args[i].label + arg = FfiArg(arg_name, kernel.adj.args[i].type, arg_name in in_out_argnames) + if arg_name in in_out_argnames: + in_out_argnames.remove(arg_name) + if arg.is_array: + # keep track of the first input array argument + if self.first_array_arg is None: + self.first_array_arg = i + self.input_args.append(arg) + + # process output args + self.output_args = [] + for i in range(self.num_inputs, self.num_kernel_args): + arg_name = kernel.adj.args[i].label + if arg_name in in_out_argnames: + raise AssertionError( + f"Expected an output-only argument for argument {arg_name}." + " in_out arguments should be placed before output-only arguments." + ) + arg = FfiArg(arg_name, kernel.adj.args[i].type, False) + if not arg.is_array: + raise TypeError("All output arguments must be arrays") + self.output_args.append(arg) + + if in_out_argnames: + raise ValueError(f"in_out_argnames: '{in_out_argnames}' did not match any function argument names.") + + # Build input output aliases. + out_id = 0 + input_output_aliases = {} + for in_id, arg in enumerate(self.input_args): + if not arg.in_out: + continue + input_output_aliases[in_id] = out_id + out_id += 1 + self.input_output_aliases = input_output_aliases + + # register the callback + FFI_CCALLFUNC = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)) + self.callback_func = FFI_CCALLFUNC(lambda call_frame: self.ffi_callback(call_frame)) + ffi_ccall_address = ctypes.cast(self.callback_func, ctypes.c_void_p) + ffi_capsule = jax.ffi.pycapsule(ffi_ccall_address.value) + jax.ffi.register_ffi_target(self.name, ffi_capsule, platform="CUDA") + + def __call__(self, *args, output_dims=None, launch_dims=None, vmap_method=None): + num_inputs = len(args) + if num_inputs != self.num_inputs: + raise ValueError(f"Expected {self.num_inputs} inputs, but got {num_inputs}") + + # default argument fallback + if launch_dims is None: + launch_dims = self.launch_dims + if output_dims is None: + output_dims = self.output_dims + if vmap_method is None: + vmap_method = self.vmap_method + + # output types + out_types = [] + + # process inputs + static_inputs = {} + for i in range(num_inputs): + input_arg = self.input_args[i] + input_value = args[i] + if input_arg.is_array: + # check dtype + if input_value.dtype != input_arg.jax_scalar_type: + raise TypeError( + f"Invalid data type for array argument '{input_arg.name}', expected {input_arg.jax_scalar_type}, got {input_value.dtype}" + ) + # check ndim + if input_value.ndim != input_arg.jax_ndim: + raise TypeError( + f"Invalid dimensionality for array argument '{input_arg.name}', expected {input_arg.jax_ndim} dimensions, got {input_value.ndim}" + ) + # check inner dims + for d in range(input_arg.dtype_ndim): + if input_value.shape[input_arg.type.ndim + d] != input_arg.dtype_shape[d]: + raise TypeError( + f"Invalid inner dimensions for array argument '{input_arg.name}', expected {input_arg.dtype_shape}, got {input_value.shape[-input_arg.dtype_ndim :]}" + ) + else: + # make sure scalar is not a traced variable, should be static + if isinstance(input_value, jax.core.Tracer): + raise ValueError(f"Argument '{input_arg.name}' must be a static value") + # stash the value to be retrieved by callback + static_inputs[input_arg.name] = input_arg.type(input_value) + + # append in-out arg to output types + if input_arg.in_out: + out_types.append(get_jax_output_type(input_arg, input_value.shape)) + + # launch dimensions + if launch_dims is None: + # use the shape of the first input array + if self.first_array_arg is not None: + launch_dims = get_warp_shape(self.input_args[self.first_array_arg], args[self.first_array_arg].shape) + else: + raise RuntimeError("Failed to determine launch dimensions") + elif isinstance(launch_dims, int): + launch_dims = (launch_dims,) + else: + launch_dims = tuple(launch_dims) + + # output shapes + if isinstance(output_dims, dict): + # assume a dictionary of shapes keyed on argument name + for output_arg in self.output_args: + dims = output_dims.get(output_arg.name) + if dims is None: + raise ValueError(f"Missing output dimensions for argument '{output_arg.name}'") + out_types.append(get_jax_output_type(output_arg, dims)) + else: + if output_dims is None: + # use launch dimensions + output_dims = launch_dims + elif isinstance(output_dims, int): + output_dims = (output_dims,) + # assume same dimensions for all outputs + for output_arg in self.output_args: + out_types.append(get_jax_output_type(output_arg, output_dims)) + + call = jax.ffi.ffi_call( + self.name, + out_types, + vmap_method=vmap_method, + input_output_aliases=self.input_output_aliases, + ) + + # ensure the kernel module is loaded before the callback, otherwise graph capture may fail + device = wp.device_from_jax(get_jax_device()) + self.kernel.module.load(device) + + # save launch data to be retrieved by callback + launch_id = self.launch_id + self.launch_descriptors[launch_id] = FfiLaunchDesc(static_inputs, launch_dims) + self.launch_id += 1 + + return call(*args, launch_id=launch_id) + + def ffi_callback(self, call_frame): + try: + # On the first call, XLA runtime will query the API version and traits + # metadata using the |extension| field. Let us respond to that query + # if the metadata extension is present. + extension = call_frame.contents.extension_start + if extension: + # Try to set the version metadata. + if extension.contents.type == XLA_FFI_Extension_Type.Metadata: + metadata_ext = ctypes.cast(extension, ctypes.POINTER(XLA_FFI_Metadata_Extension)) + metadata_ext.contents.metadata.contents.api_version.major_version = 0 + metadata_ext.contents.metadata.contents.api_version.minor_version = 1 + # Turn on CUDA graphs for this handler. + metadata_ext.contents.metadata.contents.traits = ( + XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE + ) + return None + + # retrieve call info + attrs = decode_attrs(call_frame.contents.attrs) + launch_id = int(attrs["launch_id"]) + launch_desc = self.launch_descriptors[launch_id] + + num_inputs = call_frame.contents.args.size + inputs = ctypes.cast(call_frame.contents.args.args, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) + + num_outputs = call_frame.contents.rets.size + outputs = ctypes.cast(call_frame.contents.rets.rets, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) + + assert num_inputs == self.num_inputs + assert num_outputs == self.num_outputs + + launch_bounds = launch_bounds_t(launch_desc.launch_dims) + + # first kernel param is the launch bounds + kernel_params = (ctypes.c_void_p * (1 + self.num_kernel_args))() + kernel_params[0] = ctypes.addressof(launch_bounds) + + arg_refs = [] + + # input and in-out args + for i, input_arg in enumerate(self.input_args): + if input_arg.is_array: + buffer = inputs[i].contents + shape = buffer.dims[: input_arg.type.ndim] + strides = strides_from_shape(shape, input_arg.type.dtype) + arg = array_t(buffer.data, 0, input_arg.type.ndim, shape, strides) + kernel_params[i + 1] = ctypes.addressof(arg) + arg_refs.append(arg) # keep a reference + else: + # scalar argument, get stashed value + value = launch_desc.static_inputs[input_arg.name] + arg = input_arg.type._type_(value) + kernel_params[i + 1] = ctypes.addressof(arg) + arg_refs.append(arg) # keep a reference + + # pure output args (skip in-out FFI buffers) + for i, output_arg in enumerate(self.output_args): + buffer = outputs[i + self.num_in_out].contents + shape = buffer.dims[: output_arg.type.ndim] + strides = strides_from_shape(shape, output_arg.type.dtype) + arg = array_t(buffer.data, 0, output_arg.type.ndim, shape, strides) + kernel_params[num_inputs + i + 1] = ctypes.addressof(arg) + arg_refs.append(arg) # keep a reference + + # get device and stream + device = wp.device_from_jax(get_jax_device()) + stream = get_stream_from_callframe(call_frame.contents) + + # get kernel hooks + hooks = self.kernel.module.get_kernel_hooks(self.kernel, device) + assert hooks.forward, "Failed to find kernel entry point" + + # launch the kernel + wp.context.runtime.core.wp_cuda_launch_kernel( + device.context, + hooks.forward, + launch_bounds.size, + 0, + 256, + hooks.forward_smem_bytes, + kernel_params, + stream, + ) + + except Exception as e: + print(traceback.format_exc()) + return create_ffi_error( + call_frame.contents.api, XLA_FFI_Error_Code.UNKNOWN, f"FFI callback error: {type(e).__name__}: {e}" + ) + + +class FfiCallDesc: + def __init__(self, static_inputs): + self.static_inputs = static_inputs + + +class FfiCallable: + def __init__(self, func, num_outputs, graph_compatible, vmap_method, output_dims, in_out_argnames): + self.func = func + self.name = generate_unique_name(func) + self.num_outputs = num_outputs + self.vmap_method = vmap_method + self.graph_compatible = graph_compatible + self.output_dims = output_dims + self.first_array_arg = None + self.call_id = 0 + self.call_descriptors = {} + + in_out_argnames_list = in_out_argnames or [] + in_out_argnames = set(in_out_argnames_list) + if len(in_out_argnames_list) != len(in_out_argnames): + raise AssertionError("in_out_argnames must not contain duplicate names") + + # get arguments and annotations + argspec = get_full_arg_spec(func) + + num_args = len(argspec.args) + self.num_in_out = len(in_out_argnames) + self.num_inputs = num_args - num_outputs + self.num_in_out + if self.num_outputs < 1: + raise ValueError("At least one output is required") + if self.num_outputs > num_args: + raise ValueError("Number of outputs cannot be greater than the number of kernel arguments") + if self.num_outputs < self.num_in_out: + raise ValueError("Number of outputs cannot be smaller than the number of in_out_argnames") + + if len(argspec.annotations) < num_args: + raise RuntimeError(f"Incomplete argument annotations on function {self.name}") + + # parse type annotations + self.args = [] + arg_idx = 0 + for arg_name, arg_type in argspec.annotations.items(): + if arg_name == "return": + if arg_type is not None: + raise TypeError("Function must not return a value") + else: + arg = FfiArg(arg_name, arg_type, arg_name in in_out_argnames) + if arg_name in in_out_argnames: + in_out_argnames.remove(arg_name) + if arg.is_array: + if arg_idx < self.num_inputs and self.first_array_arg is None: + self.first_array_arg = arg_idx + self.args.append(arg) + + if arg.in_out and arg_idx >= self.num_inputs: + raise AssertionError( + f"Expected an output-only argument for argument {arg_name}." + " in_out arguments should be placed before output-only arguments." + ) + + arg_idx += 1 + + if in_out_argnames: + raise ValueError(f"in_out_argnames: '{in_out_argnames}' did not match any function argument names.") + + self.input_args = self.args[: self.num_inputs] # includes in-out args + self.output_args = self.args[self.num_inputs :] # pure output args + + # Build input output aliases. + out_id = 0 + input_output_aliases = {} + for in_id, arg in enumerate(self.input_args): + if not arg.in_out: + continue + input_output_aliases[in_id] = out_id + out_id += 1 + self.input_output_aliases = input_output_aliases + + # register the callback + FFI_CCALLFUNC = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)) + self.callback_func = FFI_CCALLFUNC(lambda call_frame: self.ffi_callback(call_frame)) + ffi_ccall_address = ctypes.cast(self.callback_func, ctypes.c_void_p) + ffi_capsule = jax.ffi.pycapsule(ffi_ccall_address.value) + jax.ffi.register_ffi_target(self.name, ffi_capsule, platform="CUDA") + + def __call__(self, *args, output_dims=None, vmap_method=None): + num_inputs = len(args) + if num_inputs != self.num_inputs: + input_names = ", ".join(arg.name for arg in self.input_args) + s = "" if self.num_inputs == 1 else "s" + raise ValueError(f"Expected {self.num_inputs} input{s} ({input_names}), but got {num_inputs}") + + # default argument fallback + if vmap_method is None: + vmap_method = self.vmap_method + if output_dims is None: + output_dims = self.output_dims + + # output types + out_types = [] + + # process inputs + static_inputs = {} + for i in range(num_inputs): + input_arg = self.input_args[i] + input_value = args[i] + if input_arg.is_array: + # check dtype + if input_value.dtype != input_arg.jax_scalar_type: + raise TypeError( + f"Invalid data type for array argument '{input_arg.name}', expected {input_arg.jax_scalar_type}, got {input_value.dtype}" + ) + # check ndim + if input_value.ndim != input_arg.jax_ndim: + raise TypeError( + f"Invalid dimensionality for array argument '{input_arg.name}', expected {input_arg.jax_ndim} dimensions, got {input_value.ndim}" + ) + # check inner dims + for d in range(input_arg.dtype_ndim): + if input_value.shape[input_arg.type.ndim + d] != input_arg.dtype_shape[d]: + raise TypeError( + f"Invalid inner dimensions for array argument '{input_arg.name}', expected {input_arg.dtype_shape}, got {input_value.shape[-input_arg.dtype_ndim :]}" + ) + else: + # make sure scalar is not a traced variable, should be static + if isinstance(input_value, jax.core.Tracer): + raise ValueError(f"Argument '{input_arg.name}' must be a static value") + # stash the value to be retrieved by callback + static_inputs[input_arg.name] = input_arg.type(input_value) + + # append in-out arg to output types + if input_arg.in_out: + out_types.append(get_jax_output_type(input_arg, input_value.shape)) + + # output shapes + if isinstance(output_dims, dict): + # assume a dictionary of shapes keyed on argument name + for output_arg in self.output_args: + dims = output_dims.get(output_arg.name) + if dims is None: + raise ValueError(f"Missing output dimensions for argument '{output_arg.name}'") + out_types.append(get_jax_output_type(output_arg, dims)) + else: + if output_dims is None: + if self.first_array_arg is None: + raise ValueError("Unable to determine output dimensions") + output_dims = get_warp_shape(self.input_args[self.first_array_arg], args[self.first_array_arg].shape) + elif isinstance(output_dims, int): + output_dims = (output_dims,) + # assume same dimensions for all outputs + for output_arg in self.output_args: + out_types.append(get_jax_output_type(output_arg, output_dims)) + + call = jax.ffi.ffi_call( + self.name, + out_types, + vmap_method=vmap_method, + input_output_aliases=self.input_output_aliases, + # has_side_effect=True, # force this function to execute even if outputs aren't used + ) + + # load the module + # NOTE: if the target function uses kernels from different modules, they will not be loaded here + device = wp.device_from_jax(get_jax_device()) + module = wp.get_module(self.func.__module__) + module.load(device) + + # save call data to be retrieved by callback + call_id = self.call_id + self.call_descriptors[call_id] = FfiCallDesc(static_inputs) + self.call_id += 1 + return call(*args, call_id=call_id) + + def ffi_callback(self, call_frame): + try: + # TODO Try-catch around the body and return XLA_FFI_Error on error. + extension = call_frame.contents.extension_start + # On the first call, XLA runtime will query the API version and traits + # metadata using the |extension| field. Let us respond to that query + # if the metadata extension is present. + if extension: + # Try to set the version metadata. + if extension.contents.type == XLA_FFI_Extension_Type.Metadata: + metadata_ext = ctypes.cast(extension, ctypes.POINTER(XLA_FFI_Metadata_Extension)) + metadata_ext.contents.metadata.contents.api_version.major_version = 0 + metadata_ext.contents.metadata.contents.api_version.minor_version = 1 + # Turn on CUDA graphs for this handler. + if self.graph_compatible: + metadata_ext.contents.metadata.contents.traits = ( + XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE + ) + return None + + # retrieve call info + attrs = decode_attrs(call_frame.contents.attrs) + call_id = int(attrs["call_id"]) + call_desc = self.call_descriptors[call_id] + + num_inputs = call_frame.contents.args.size + inputs = ctypes.cast(call_frame.contents.args.args, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) + + num_outputs = call_frame.contents.rets.size + outputs = ctypes.cast(call_frame.contents.rets.rets, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) + + assert num_inputs == self.num_inputs + assert num_outputs == self.num_outputs + + device = wp.device_from_jax(get_jax_device()) + cuda_stream = get_stream_from_callframe(call_frame.contents) + stream = wp.Stream(device, cuda_stream=cuda_stream) + + # reconstruct the argument list + arg_list = [] + + # input and in-out args + for i, arg in enumerate(self.input_args): + if arg.is_array: + buffer = inputs[i].contents + shape = buffer.dims[: buffer.rank - arg.dtype_ndim] + arr = wp.array(ptr=buffer.data, dtype=arg.type.dtype, shape=shape, device=device) + arg_list.append(arr) + else: + # scalar argument, get stashed value + value = call_desc.static_inputs[arg.name] + arg_list.append(value) + + # pure output args (skip in-out FFI buffers) + for i, arg in enumerate(self.output_args): + buffer = outputs[i + self.num_in_out].contents + shape = buffer.dims[: buffer.rank - arg.dtype_ndim] + arr = wp.array(ptr=buffer.data, dtype=arg.type.dtype, shape=shape, device=device) + arg_list.append(arr) + + # call the Python function with reconstructed arguments + with wp.ScopedStream(stream, sync_enter=False): + if stream.is_capturing: + with wp.ScopedCapture(stream=stream, external=True) as capture: + self.func(*arg_list) + # keep a reference to the capture object to prevent required modules getting unloaded + call_desc.capture = capture + else: + self.func(*arg_list) + + except Exception as e: + print(traceback.format_exc()) + return create_ffi_error( + call_frame.contents.api, XLA_FFI_Error_Code.UNKNOWN, f"FFI callback error: {type(e).__name__}: {e}" + ) + + return None + + +# Holders for the custom callbacks to keep them alive. +_FFI_CALLABLE_REGISTRY: dict[str, FfiCallable] = {} +_FFI_KERNEL_REGISTRY: dict[str, FfiKernel] = {} +_FFI_REGISTRY_LOCK = threading.Lock() + + +def jax_kernel( + kernel, num_outputs=1, vmap_method="broadcast_all", launch_dims=None, output_dims=None, in_out_argnames=None +): + """Create a JAX callback from a Warp kernel. + + NOTE: This is an experimental feature under development. + + Args: + kernel: The Warp kernel to launch. + num_outputs: Optional. Specify the number of output arguments if greater than 1. + This must include the number of ``in_out_arguments``. + vmap_method: Optional. String specifying how the callback transforms under ``vmap()``. + This argument can also be specified for individual calls. + launch_dims: Optional. Specify the default kernel launch dimensions. If None, launch + dimensions are inferred from the shape of the first array argument. + This argument can also be specified for individual calls. + output_dims: Optional. Specify the default dimensions of output arrays. If None, output + dimensions are inferred from the launch dimensions. + This argument can also be specified for individual calls. + in_out_argnames: Optional. Names of input-output arguments. + + Limitations: + - All kernel arguments must be contiguous arrays or scalars. + - Scalars must be static arguments in JAX. + - Input and input-output arguments must precede the output arguments in the ``kernel`` definition. + - There must be at least one output or input-output argument. + - Only the CUDA backend is supported. + """ + key = ( + kernel.func, + num_outputs, + vmap_method, + tuple(launch_dims) if launch_dims else launch_dims, + tuple(sorted(output_dims.items())) if output_dims else output_dims, + ) + + with _FFI_REGISTRY_LOCK: + if key not in _FFI_KERNEL_REGISTRY: + new_kernel = FfiKernel(kernel, num_outputs, vmap_method, launch_dims, output_dims, in_out_argnames) + _FFI_KERNEL_REGISTRY[key] = new_kernel + + return _FFI_KERNEL_REGISTRY[key] + + +def jax_callable( + func: Callable, + num_outputs: int = 1, + graph_compatible: bool = True, + vmap_method: Optional[str] = "broadcast_all", + output_dims=None, + in_out_argnames=None, +): + """Create a JAX callback from an annotated Python function. + + The Python function arguments must have type annotations like Warp kernels. + + NOTE: This is an experimental feature under development. + + Args: + func: The Python function to call. + num_outputs: Optional. Specify the number of output arguments if greater than 1. + This must include the number of ``in_out_arguments``. + graph_compatible: Optional. Whether the function can be called during CUDA graph capture. + vmap_method: Optional. String specifying how the callback transforms under ``vmap()``. + This argument can also be specified for individual calls. + output_dims: Optional. Specify the default dimensions of output arrays. + If ``None``, output dimensions are inferred from the launch dimensions. + This argument can also be specified for individual calls. + in_out_argnames: Optional. Names of input-output arguments. + + Limitations: + - All kernel arguments must be contiguous arrays or scalars. + - Scalars must be static arguments in JAX. + - Input and input-output arguments must precede the output arguments in the ``func`` definition. + - There must be at least one output or input-output argument. + - Only the CUDA backend is supported. + """ + key = ( + func, + num_outputs, + graph_compatible, + vmap_method, + tuple(sorted(output_dims.items())) if output_dims else output_dims, + ) + + with _FFI_REGISTRY_LOCK: + if key not in _FFI_CALLABLE_REGISTRY: + new_callable = FfiCallable(func, num_outputs, graph_compatible, vmap_method, output_dims, in_out_argnames) + _FFI_CALLABLE_REGISTRY[key] = new_callable + + return _FFI_CALLABLE_REGISTRY[key] + + +############################################################################### +# +# Generic FFI callbacks for Python functions of the form +# func(inputs, outputs, attrs, ctx) +# +############################################################################### + + +def register_ffi_callback(name: str, func: Callable, graph_compatible: bool = True) -> None: + """Create a JAX callback from a Python function. + + The Python function must have the form ``func(inputs, outputs, attrs, ctx)``. + + NOTE: This is an experimental feature under development. + + Args: + name: A unique FFI callback name. + func: The Python function to call. + graph_compatible: Optional. Whether the function can be called during CUDA graph capture. + """ + + # TODO check that the name is not already registered + + def ffi_callback(call_frame): + try: + # TODO Try-catch around the body and return XLA_FFI_Error on error. + extension = call_frame.contents.extension_start + # On the first call, XLA runtime will query the API version and traits + # metadata using the |extension| field. Let us respond to that query + # if the metadata extension is present. + if extension: + # Try to set the version metadata. + if extension.contents.type == XLA_FFI_Extension_Type.Metadata: + metadata_ext = ctypes.cast(extension, ctypes.POINTER(XLA_FFI_Metadata_Extension)) + metadata_ext.contents.metadata.contents.api_version.major_version = 0 + metadata_ext.contents.metadata.contents.api_version.minor_version = 1 + if graph_compatible: + # Turn on CUDA graphs for this handler. + metadata_ext.contents.metadata.contents.traits = ( + XLA_FFI_Handler_TraitsBits.COMMAND_BUFFER_COMPATIBLE + ) + return None + + attrs = decode_attrs(call_frame.contents.attrs) + + input_count = call_frame.contents.args.size + inputs = ctypes.cast(call_frame.contents.args.args, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) + inputs = [FfiBuffer(inputs[i].contents) for i in range(input_count)] + + output_count = call_frame.contents.rets.size + outputs = ctypes.cast(call_frame.contents.rets.rets, ctypes.POINTER(ctypes.POINTER(XLA_FFI_Buffer))) + outputs = [FfiBuffer(outputs[i].contents) for i in range(output_count)] + + ctx = ExecutionContext(call_frame.contents) + + func(inputs, outputs, attrs, ctx) + except Exception as e: + print(traceback.format_exc()) + return create_ffi_error( + call_frame.contents.api, XLA_FFI_Error_Code.UNKNOWN, f"FFI callback error: {type(e).__name__}: {e}" + ) + + return None + + FFI_CCALLFUNC = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_CallFrame)) + callback_func = FFI_CCALLFUNC(ffi_callback) + with _FFI_REGISTRY_LOCK: + _FFI_CALLABLE_REGISTRY[name] = callback_func + ffi_ccall_address = ctypes.cast(callback_func, ctypes.c_void_p) + ffi_capsule = jax.ffi.pycapsule(ffi_ccall_address.value) + jax.ffi.register_ffi_target(name, ffi_capsule, platform="CUDA") + + +############################################################################### +# +# Utilities +# +############################################################################### + +# ensure unique FFI callback names +ffi_name_counts = {} + + +def generate_unique_name(func) -> str: + key = make_full_qualified_name(func) + unique_id = ffi_name_counts.get(key, 0) + ffi_name_counts[key] = unique_id + 1 + return f"{key}_{unique_id}" + + +def get_warp_shape(arg, dims): + if arg.dtype_ndim > 0: + # vector/matrix array + return dims[: arg.warp_ndim] + else: + # scalar array + return dims + + +def get_jax_output_type(arg, dims): + if isinstance(dims, int): + dims = (dims,) + + ndim = len(dims) + + if arg.dtype_ndim > 0: + # vector/matrix array + if ndim == arg.warp_ndim: + return jax.ShapeDtypeStruct((*dims, *arg.dtype_shape), arg.jax_scalar_type) + elif ndim == arg.jax_ndim: + # make sure inner dimensions match + inner_dims = dims[-arg.dtype_ndim :] + for i in range(arg.dtype_ndim): + if inner_dims[i] != arg.dtype_shape[i]: + raise ValueError(f"Invalid output dimensions for argument '{arg.name}': {dims}") + return jax.ShapeDtypeStruct(dims, arg.jax_scalar_type) + else: + raise ValueError(f"Invalid output dimensions for argument '{arg.name}': {dims}") + else: + # scalar array + if ndim != arg.warp_ndim: + raise ValueError(f"Invalid output dimensions for argument '{arg.name}': {dims}") + return jax.ShapeDtypeStruct(dims, arg.jax_scalar_type) diff --git a/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py b/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py new file mode 100644 index 00000000..8f33030b --- /dev/null +++ b/mjx/mujoco/mjx/third_party/warp/jax_experimental/xla_ffi.py @@ -0,0 +1,615 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import ctypes +import enum + +import jax.numpy as jnp +import numpy as np + +import warp as wp + +####################################################################### +# ctypes structures and enums for XLA's FFI API: +# https://github.com/openxla/xla/blob/a1a5e62fbffa3a3b6c409d72607456cf5b353a22/xla/ffi/api/c_api.h +####################################################################### + + +# typedef enum { +# XLA_FFI_Extension_Metadata = 1, +# } XLA_FFI_Extension_Type; +class XLA_FFI_Extension_Type(enum.IntEnum): + Metadata = 1 + + +# typedef struct XLA_FFI_Extension_Base { +# size_t struct_size; +# XLA_FFI_Extension_Type type; +# struct XLA_FFI_Extension_Base* next; +# } XLA_FFI_Extension_Base; +class XLA_FFI_Extension_Base(ctypes.Structure): + pass + + +XLA_FFI_Extension_Base._fields_ = [ + ("struct_size", ctypes.c_size_t), + ("type", ctypes.c_int), # XLA_FFI_Extension_Type + ("next", ctypes.POINTER(XLA_FFI_Extension_Base)), +] + + +# typedef enum { +# XLA_FFI_ExecutionStage_INSTANTIATE = 0, +# XLA_FFI_ExecutionStage_PREPARE = 1, +# XLA_FFI_ExecutionStage_INITIALIZE = 2, +# XLA_FFI_ExecutionStage_EXECUTE = 3, +# } XLA_FFI_ExecutionStage; +class XLA_FFI_ExecutionStage(enum.IntEnum): + INSTANTIATE = 0 + PREPARE = 1 + INITIALIZE = 2 + EXECUTE = 3 + + +# typedef enum { +# XLA_FFI_DataType_INVALID = 0, +# XLA_FFI_DataType_PRED = 1, +# XLA_FFI_DataType_S8 = 2, +# XLA_FFI_DataType_S16 = 3, +# XLA_FFI_DataType_S32 = 4, +# XLA_FFI_DataType_S64 = 5, +# XLA_FFI_DataType_U8 = 6, +# XLA_FFI_DataType_U16 = 7, +# XLA_FFI_DataType_U32 = 8, +# XLA_FFI_DataType_U64 = 9, +# XLA_FFI_DataType_F16 = 10, +# XLA_FFI_DataType_F32 = 11, +# XLA_FFI_DataType_F64 = 12, +# XLA_FFI_DataType_BF16 = 16, +# XLA_FFI_DataType_C64 = 15, +# XLA_FFI_DataType_C128 = 18, +# XLA_FFI_DataType_TOKEN = 17, +# XLA_FFI_DataType_F8E5M2 = 19, +# XLA_FFI_DataType_F8E3M4 = 29, +# XLA_FFI_DataType_F8E4M3 = 28, +# XLA_FFI_DataType_F8E4M3FN = 20, +# XLA_FFI_DataType_F8E4M3B11FNUZ = 23, +# XLA_FFI_DataType_F8E5M2FNUZ = 24, +# XLA_FFI_DataType_F8E4M3FNUZ = 25, +# XLA_FFI_DataType_F4E2M1FN = 32, +# XLA_FFI_DataType_F8E8M0FNU = 33, +# } XLA_FFI_DataType; +class XLA_FFI_DataType(enum.IntEnum): + INVALID = 0 + PRED = 1 + S8 = 2 + S16 = 3 + S32 = 4 + S64 = 5 + U8 = 6 + U16 = 7 + U32 = 8 + U64 = 9 + F16 = 10 + F32 = 11 + F64 = 12 + BF16 = 16 + C64 = 15 + C128 = 18 + TOKEN = 17 + F8E5M2 = 19 + F8E3M4 = 29 + F8E4M3 = 28 + F8E4M3FN = 20 + F8E4M3B11FNUZ = 23 + F8E5M2FNUZ = 24 + F8E4M3FNUZ = 25 + F4E2M1FN = 32 + F8E8M0FNU = 33 + + +# struct XLA_FFI_Buffer { +# size_t struct_size; +# XLA_FFI_Extension_Base* extension_start; +# +# XLA_FFI_DataType dtype; +# void* data; +# int64_t rank; +# int64_t* dims; // length == rank +# }; +class XLA_FFI_Buffer(ctypes.Structure): + _fields_ = ( + ("struct_size", ctypes.c_size_t), + ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), + ("dtype", ctypes.c_int), # XLA_FFI_DataType + ("data", ctypes.c_void_p), + ("rank", ctypes.c_int64), + ("dims", ctypes.POINTER(ctypes.c_int64)), + ) + + +# typedef enum { +# XLA_FFI_ArgType_BUFFER = 1, +# } XLA_FFI_ArgType; +class XLA_FFI_ArgType(enum.IntEnum): + BUFFER = 1 + + +# typedef enum { +# XLA_FFI_RetType_BUFFER = 1, +# } XLA_FFI_RetType; +class XLA_FFI_RetType(enum.IntEnum): + BUFFER = 1 + + +# struct XLA_FFI_Args { +# size_t struct_size; +# XLA_FFI_Extension_Base* extension_start; +# int64_t size; +# XLA_FFI_ArgType* types; // length == size +# void** args; // length == size +# }; +class XLA_FFI_Args(ctypes.Structure): + _fields_ = ( + ("struct_size", ctypes.c_size_t), + ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), + ("size", ctypes.c_int64), + ("types", ctypes.POINTER(ctypes.c_int)), # XLA_FFI_ArgType* + ("args", ctypes.POINTER(ctypes.c_void_p)), + ) + + +# struct XLA_FFI_Rets { +# size_t struct_size; +# XLA_FFI_Extension_Base* extension_start; +# int64_t size; +# XLA_FFI_RetType* types; // length == size +# void** rets; // length == size +# }; +class XLA_FFI_Rets(ctypes.Structure): + _fields_ = ( + ("struct_size", ctypes.c_size_t), + ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), + ("size", ctypes.c_int64), + ("types", ctypes.POINTER(ctypes.c_int)), # XLA_FFI_RetType* + ("rets", ctypes.POINTER(ctypes.c_void_p)), + ) + + +# typedef struct XLA_FFI_ByteSpan { +# const char* ptr; +# size_t len; +# } XLA_FFI_ByteSpan; +class XLA_FFI_ByteSpan(ctypes.Structure): + _fields_ = ( + ("ptr", ctypes.POINTER(ctypes.c_char)), + ("len", ctypes.c_size_t), + ) + + +# typedef struct XLA_FFI_Scalar { +# XLA_FFI_DataType dtype; +# void* value; +# } XLA_FFI_Scalar; +class XLA_FFI_Scalar(ctypes.Structure): + _fields_ = ( + ("dtype", ctypes.c_int), + ("value", ctypes.c_void_p), + ) + + +# typedef struct XLA_FFI_Array { +# XLA_FFI_DataType dtype; +# size_t size; +# void* data; +# } XLA_FFI_Array; +class XLA_FFI_Array(ctypes.Structure): + _fields_ = ( + ("dtype", ctypes.c_int), + ("size", ctypes.c_size_t), + ("data", ctypes.c_void_p), + ) + + +# typedef enum { +# XLA_FFI_AttrType_ARRAY = 1, +# XLA_FFI_AttrType_DICTIONARY = 2, +# XLA_FFI_AttrType_SCALAR = 3, +# XLA_FFI_AttrType_STRING = 4, +# } XLA_FFI_AttrType; +class XLA_FFI_AttrType(enum.IntEnum): + ARRAY = 1 + DICTIONARY = 2 + SCALAR = 3 + STRING = 4 + + +# struct XLA_FFI_Attrs { +# size_t struct_size; +# XLA_FFI_Extension_Base* extension_start; +# int64_t size; +# XLA_FFI_AttrType* types; // length == size +# XLA_FFI_ByteSpan** names; // length == size +# void** attrs; // length == size +# }; +class XLA_FFI_Attrs(ctypes.Structure): + _fields_ = ( + ("struct_size", ctypes.c_size_t), + ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), + ("size", ctypes.c_int64), + ("types", ctypes.POINTER(ctypes.c_int)), # XLA_FFI_AttrType* + ("names", ctypes.POINTER(ctypes.POINTER(XLA_FFI_ByteSpan))), + ("attrs", ctypes.POINTER(ctypes.c_void_p)), + ) + + +# struct XLA_FFI_Api_Version { +# size_t struct_size; +# XLA_FFI_Extension_Base* extension_start; +# int major_version; // out +# int minor_version; // out +# }; +class XLA_FFI_Api_Version(ctypes.Structure): + _fields_ = ( + ("struct_size", ctypes.c_size_t), + ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), + ("major_version", ctypes.c_int), + ("minor_version", ctypes.c_int), + ) + + +# enum XLA_FFI_Handler_TraitsBits { +# // Calls to FFI handler are safe to trace into the command buffer. It means +# // that calls to FFI handler always launch exactly the same device operations +# // (can depend on attribute values) that can be captured and then replayed. +# XLA_FFI_HANDLER_TRAITS_COMMAND_BUFFER_COMPATIBLE = 1u << 0, +# }; +class XLA_FFI_Handler_TraitsBits(enum.IntEnum): + COMMAND_BUFFER_COMPATIBLE = 1 << 0 + + +# struct XLA_FFI_Metadata { +# size_t struct_size; +# XLA_FFI_Api_Version api_version; +# XLA_FFI_Handler_Traits traits; +# }; +class XLA_FFI_Metadata(ctypes.Structure): + _fields_ = ( + ("struct_size", ctypes.c_size_t), + ("api_version", XLA_FFI_Api_Version), # XLA_FFI_Extension_Type + ("traits", ctypes.c_uint32), # XLA_FFI_Handler_Traits + ) + + +# struct XLA_FFI_Metadata_Extension { +# XLA_FFI_Extension_Base extension_base; +# XLA_FFI_Metadata* metadata; +# }; +class XLA_FFI_Metadata_Extension(ctypes.Structure): + _fields_ = ( + ("extension_base", XLA_FFI_Extension_Base), + ("metadata", ctypes.POINTER(XLA_FFI_Metadata)), + ) + + +# typedef enum { +# XLA_FFI_Error_Code_OK = 0, +# XLA_FFI_Error_Code_CANCELLED = 1, +# XLA_FFI_Error_Code_UNKNOWN = 2, +# XLA_FFI_Error_Code_INVALID_ARGUMENT = 3, +# XLA_FFI_Error_Code_DEADLINE_EXCEEDED = 4, +# XLA_FFI_Error_Code_NOT_FOUND = 5, +# XLA_FFI_Error_Code_ALREADY_EXISTS = 6, +# XLA_FFI_Error_Code_PERMISSION_DENIED = 7, +# XLA_FFI_Error_Code_RESOURCE_EXHAUSTED = 8, +# XLA_FFI_Error_Code_FAILED_PRECONDITION = 9, +# XLA_FFI_Error_Code_ABORTED = 10, +# XLA_FFI_Error_Code_OUT_OF_RANGE = 11, +# XLA_FFI_Error_Code_UNIMPLEMENTED = 12, +# XLA_FFI_Error_Code_INTERNAL = 13, +# XLA_FFI_Error_Code_UNAVAILABLE = 14, +# XLA_FFI_Error_Code_DATA_LOSS = 15, +# XLA_FFI_Error_Code_UNAUTHENTICATED = 16 +# } XLA_FFI_Error_Code; +class XLA_FFI_Error_Code(enum.IntEnum): + OK = 0 + CANCELLED = 1 + UNKNOWN = 2 + INVALID_ARGUMENT = 3 + DEADLINE_EXCEEDED = 4 + NOT_FOUND = 5 + ALREADY_EXISTS = 6 + PERMISSION_DENIED = 7 + RESOURCE_EXHAUSTED = 8 + FAILED_PRECONDITION = 9 + ABORTED = 10 + OUT_OF_RANGE = 11 + UNIMPLEMENTED = 12 + INTERNAL = 13 + UNAVAILABLE = 14 + DATA_LOSS = 15 + UNAUTHENTICATED = 16 + + +# struct XLA_FFI_Error_Create_Args { +# size_t struct_size; +# XLA_FFI_Extension_Base* extension_start; +# const char* message; +# XLA_FFI_Error_Code errc; +# }; +class XLA_FFI_Error_Create_Args(ctypes.Structure): + _fields_ = ( + ("struct_size", ctypes.c_size_t), + ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), + ("message", ctypes.c_char_p), + ("errc", ctypes.c_int), + ) # XLA_FFI_Error_Code + + +XLA_FFI_Error_Create = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_Error_Create_Args)) + + +# struct XLA_FFI_Stream_Get_Args { +# size_t struct_size; +# XLA_FFI_Extension_Base* extension_start; +# XLA_FFI_ExecutionContext* ctx; +# void* stream; // out +# }; +class XLA_FFI_Stream_Get_Args(ctypes.Structure): + _fields_ = ( + ("struct_size", ctypes.c_size_t), + ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), + ("ctx", ctypes.c_void_p), # XLA_FFI_ExecutionContext* + ("stream", ctypes.c_void_p), + ) # // out + + +XLA_FFI_Stream_Get = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.POINTER(XLA_FFI_Stream_Get_Args)) + + +# struct XLA_FFI_Api { +# size_t struct_size; +# XLA_FFI_Extension_Base* extension_start; +# +# XLA_FFI_Api_Version api_version; +# XLA_FFI_InternalApi* internal_api; +# +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Error_Create); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Error_GetMessage); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Error_Destroy); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Handler_Register); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Stream_Get); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_TypeId_Register); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_ExecutionContext_Get); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_State_Set); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_State_Get); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_DeviceMemory_Allocate); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_DeviceMemory_Free); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_ThreadPool_Schedule); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_ThreadPool_NumThreads); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Future_Create); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Future_SetAvailable); +# _XLA_FFI_API_STRUCT_FIELD(XLA_FFI_Future_SetError); +# }; +class XLA_FFI_Api(ctypes.Structure): + _fields_ = ( + ("struct_size", ctypes.c_size_t), + ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), + ("api_version", XLA_FFI_Api_Version), + ("internal_api", ctypes.c_void_p), # XLA_FFI_InternalApi* + ("XLA_FFI_Error_Create", XLA_FFI_Error_Create), # XLA_FFI_Error_Create + ("XLA_FFI_Error_GetMessage", ctypes.c_void_p), # XLA_FFI_Error_GetMessage + ("XLA_FFI_Error_Destroy", ctypes.c_void_p), # XLA_FFI_Error_Destroy + ("XLA_FFI_Handler_Register", ctypes.c_void_p), # XLA_FFI_Handler_Register + ("XLA_FFI_Stream_Get", XLA_FFI_Stream_Get), # XLA_FFI_Stream_Get + ("XLA_FFI_TypeId_Register", ctypes.c_void_p), # XLA_FFI_TypeId_Register + ("XLA_FFI_ExecutionContext_Get", ctypes.c_void_p), # XLA_FFI_ExecutionContext_Get + ("XLA_FFI_State_Set", ctypes.c_void_p), # XLA_FFI_State_Set + ("XLA_FFI_State_Get", ctypes.c_void_p), # XLA_FFI_State_Get + ("XLA_FFI_DeviceMemory_Allocate", ctypes.c_void_p), # XLA_FFI_DeviceMemory_Allocate + ("XLA_FFI_DeviceMemory_Free", ctypes.c_void_p), # XLA_FFI_DeviceMemory_Free + ("XLA_FFI_ThreadPool_Schedule", ctypes.c_void_p), # XLA_FFI_ThreadPool_Schedule + ("XLA_FFI_ThreadPool_NumThreads", ctypes.c_void_p), # XLA_FFI_ThreadPool_NumThreads + ("XLA_FFI_Future_Create", ctypes.c_void_p), # XLA_FFI_Future_Create + ("XLA_FFI_Future_SetAvailable", ctypes.c_void_p), # XLA_FFI_Future_SetAvailable + ("XLA_FFI_Future_SetError", ctypes.c_void_p), # XLA_FFI_Future_SetError + ) + + +# struct XLA_FFI_CallFrame { +# size_t struct_size; +# XLA_FFI_Extension_Base* extension_start; +# const XLA_FFI_Api* api; +# XLA_FFI_ExecutionContext* ctx; +# XLA_FFI_ExecutionStage stage; +# XLA_FFI_Args args; +# XLA_FFI_Rets rets; +# XLA_FFI_Attrs attrs; +# +# // XLA FFI handler implementation can use `future` to signal a result of +# // asynchronous computation to the XLA runtime. XLA runtime will keep all +# // arguments, results and attributes alive until `future` is completed. +# XLA_FFI_Future* future; // out +# }; +class XLA_FFI_CallFrame(ctypes.Structure): + _fields_ = ( + ("struct_size", ctypes.c_size_t), + ("extension_start", ctypes.POINTER(XLA_FFI_Extension_Base)), + ("api", ctypes.POINTER(XLA_FFI_Api)), + ("ctx", ctypes.c_void_p), # XLA_FFI_ExecutionContext* + ("stage", ctypes.c_int), # XLA_FFI_ExecutionStage + ("args", XLA_FFI_Args), + ("rets", XLA_FFI_Rets), + ("attrs", XLA_FFI_Attrs), + ("future", ctypes.c_void_p), # XLA_FFI_Future* // out + ) + + +_xla_data_type_to_constructor = { + # XLA_FFI_DataType.INVALID + XLA_FFI_DataType.PRED: jnp.bool, + XLA_FFI_DataType.S8: jnp.int8, + XLA_FFI_DataType.S16: jnp.int16, + XLA_FFI_DataType.S32: jnp.int32, + XLA_FFI_DataType.S64: jnp.int64, + XLA_FFI_DataType.U8: jnp.uint8, + XLA_FFI_DataType.U16: jnp.uint16, + XLA_FFI_DataType.U32: jnp.uint32, + XLA_FFI_DataType.U64: jnp.uint64, + XLA_FFI_DataType.F16: jnp.float16, + XLA_FFI_DataType.F32: jnp.float32, + XLA_FFI_DataType.F64: jnp.float64, + XLA_FFI_DataType.BF16: jnp.bfloat16, + XLA_FFI_DataType.C64: jnp.complex64, + XLA_FFI_DataType.C128: jnp.complex128, + # XLA_FFI_DataType.TOKEN + XLA_FFI_DataType.F8E5M2: jnp.float8_e5m2, + XLA_FFI_DataType.F8E3M4: jnp.float8_e3m4, + XLA_FFI_DataType.F8E4M3: jnp.float8_e4m3, + XLA_FFI_DataType.F8E4M3FN: jnp.float8_e4m3fn, + XLA_FFI_DataType.F8E4M3B11FNUZ: jnp.float8_e4m3b11fnuz, + XLA_FFI_DataType.F8E5M2FNUZ: jnp.float8_e5m2fnuz, + XLA_FFI_DataType.F8E4M3FNUZ: jnp.float8_e4m3fnuz, + # XLA_FFI_DataType.F4E2M1FN: jnp.float4_e2m1fn.dtype, + # XLA_FFI_DataType.F8E8M0FNU: jnp.float8_e8m0fnu.dtype, +} + + +######################################################################## +# Helpers for translating between ctypes and python types +####################################################################### + + +def decode_bytespan(span: XLA_FFI_ByteSpan): + len = span.len + chars = ctypes.cast(span.ptr, ctypes.POINTER(ctypes.c_char * len)) + return chars.contents.value.decode("utf-8") + + +def decode_scalar(scalar: XLA_FFI_Scalar): + # TODO validate if dtype supported + dtype = jnp.dtype(_xla_data_type_to_constructor[scalar.dtype]) + bytes = ctypes.string_at(scalar.value, dtype.itemsize) + return np.frombuffer(bytes, dtype=dtype).reshape(()) + + +def decode_array(array: XLA_FFI_Array): + # TODO validate if dtype supported + dtype = jnp.dtype(_xla_data_type_to_constructor[array.dtype]) + bytes = ctypes.string_at(array.data, dtype.itemsize * array.size) + return np.frombuffer(bytes, dtype=dtype) + + +def decode_attrs(attrs: XLA_FFI_Attrs): + result = {} + for i in range(attrs.size): + attr_name = decode_bytespan(attrs.names[i].contents) + attr_type = attrs.types[i] + if attr_type == XLA_FFI_AttrType.STRING: + bytespan = ctypes.cast(attrs.attrs[i], ctypes.POINTER(XLA_FFI_ByteSpan)) + attr_value = decode_bytespan(bytespan.contents) + elif attr_type == XLA_FFI_AttrType.SCALAR: + attr_value = ctypes.cast(attrs.attrs[i], ctypes.POINTER(XLA_FFI_Scalar)) + attr_value = decode_scalar(attr_value.contents) + elif attr_type == XLA_FFI_AttrType.ARRAY: + attr_value = ctypes.cast(attrs.attrs[i], ctypes.POINTER(XLA_FFI_Array)) + attr_value = decode_array(attr_value.contents) + elif attr_type == XLA_FFI_AttrType.DICTIONARY: + attr_value = ctypes.cast(attrs.attrs[i], ctypes.POINTER(XLA_FFI_Attrs)) + attr_value = decode_attrs(attr_value.contents) + else: + raise Exception("Unexpected attr type") + result[attr_name] = attr_value + return result + + +# error-string to XLA_FFI_Error +def create_ffi_error(api, errc, message): + create_args = XLA_FFI_Error_Create_Args( + ctypes.sizeof(XLA_FFI_Error_Create_Args), + ctypes.POINTER(XLA_FFI_Extension_Base)(), + ctypes.c_char_p(message.encode("utf-8")), + errc, + ) + return api.contents.XLA_FFI_Error_Create(create_args) + + +def create_invalid_argument_ffi_error(api, message): + return create_ffi_error(api, XLA_FFI_Error_Code.INVALID_ARGUMENT, message) + + +# Extract CUDA stream from XLA_FFI_CallFrame. +def get_stream_from_callframe(call_frame): + api = call_frame.api + get_stream_args = XLA_FFI_Stream_Get_Args( + ctypes.sizeof(XLA_FFI_Stream_Get_Args), ctypes.POINTER(XLA_FFI_Extension_Base)(), call_frame.ctx, None + ) + api.contents.XLA_FFI_Stream_Get(get_stream_args) + # TODO check result + return get_stream_args.stream + + +_dtype_from_ffi = { + XLA_FFI_DataType.S8: wp.int8, + XLA_FFI_DataType.S16: wp.int16, + XLA_FFI_DataType.S32: wp.int32, + XLA_FFI_DataType.S64: wp.int64, + XLA_FFI_DataType.U8: wp.uint8, + XLA_FFI_DataType.U16: wp.uint16, + XLA_FFI_DataType.U32: wp.uint32, + XLA_FFI_DataType.U64: wp.uint64, + XLA_FFI_DataType.F16: wp.float16, + XLA_FFI_DataType.F32: wp.float32, + XLA_FFI_DataType.F64: wp.float64, +} + + +def dtype_from_ffi(ffi_dtype): + return _dtype_from_ffi.get(ffi_dtype) + + +def jax_dtype_from_ffi(ffi_dtype): + return _xla_data_type_to_constructor.get(ffi_dtype) + + +# Execution context (stream, stage) +class ExecutionContext: + stage: XLA_FFI_ExecutionStage + stream: int + + def __init__(self, callframe: XLA_FFI_CallFrame): + self.stage = XLA_FFI_ExecutionStage(callframe.stage) + self.stream = get_stream_from_callframe(callframe) + + +class FfiBuffer: + dtype: str + data: int + shape: tuple[int] + + def __init__(self, xla_buffer): + # TODO check if valid + self.dtype = jnp.dtype(_xla_data_type_to_constructor[xla_buffer.dtype]) + self.shape = tuple(xla_buffer.dims[i] for i in range(xla_buffer.rank)) + self.data = xla_buffer.data + + @property + def __cuda_array_interface__(self): + return { + "shape": self.shape, + "typestr": self.dtype.char, + "data": (self.data, False), + "version": 2, + } diff --git a/mjx/mujoco/mjx/viewer.py b/mjx/mujoco/mjx/viewer.py index 5cdce300..b8bbff9d 100644 --- a/mjx/mujoco/mjx/viewer.py +++ b/mjx/mujoco/mjx/viewer.py @@ -15,7 +15,10 @@ """An example integration of MJX with the MuJoCo viewer.""" import logging -import time +import os + +os.environ['XLA_FLAGS'] = '--xla_gpu_graph_min_graph_size=1' +import time # pylint: disable=g-import-not-at-top from typing import Sequence from absl import app @@ -25,13 +28,29 @@ from jax import numpy as jp import mujoco from mujoco import mjx import mujoco.viewer +import warp as wp _JIT = flags.DEFINE_bool('jit', True, 'To jit or not to jit.') _MODEL_PATH = flags.DEFINE_string( 'mjcf', None, 'Path to a MuJoCo MJCF file.', required=True ) - +_IMPL = flags.DEFINE_string('impl', 'jax', 'MJX implementation.') +_WP_KERNEL_CACHE_DIR = flags.DEFINE_string( + 'wp_kernel_cache_dir', + None, + 'Path to the Warp kernel cache directory.', +) +_NCONMAX = flags.DEFINE_integer( + 'nconmax', + None, + 'Maximum number of contacts to simulate, warp only.', +) +_NJMAX = flags.DEFINE_integer( + 'njmax', + None, + 'Maximum number of constraints to simulate, warp only.', +) _VIEWER_GLOBAL_STATE = { 'running': True, @@ -49,6 +68,9 @@ def _main(argv: Sequence[str]) -> None: if len(argv) > 1: raise app.UsageError('Too many command-line arguments.') + if _WP_KERNEL_CACHE_DIR.value: + wp.config.kernel_cache_dir = _WP_KERNEL_CACHE_DIR.value + jax.config.update('jax_debug_nans', True) print(f'Loading model from: {_MODEL_PATH.value}.') @@ -57,8 +79,16 @@ def _main(argv: Sequence[str]) -> None: else: m = mujoco.MjModel.from_xml_path(_MODEL_PATH.value) d = mujoco.MjData(m) - mx = mjx.put_model(m) - dx = mjx.put_data(m, d) + mx = mjx.put_model(m, impl=_IMPL.value) + if _IMPL.value == 'warp': + # TODO(btaba): use put_data. + dx = mjx.make_data( + m, impl=_IMPL.value, nconmax=_NCONMAX.value, njmax=_NJMAX.value + ) + else: + dx = mjx.put_data( + m, d, impl=_IMPL.value, nconmax=_NCONMAX.value, njmax=_NJMAX.value + ) print(f'Default backend: {jax.default_backend()}') step_fn = mjx.step diff --git a/mjx/mujoco/mjx/warp/__init__.py b/mjx/mujoco/mjx/warp/__init__.py new file mode 100644 index 00000000..067adf1e --- /dev/null +++ b/mjx/mujoco/mjx/warp/__init__.py @@ -0,0 +1,70 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +import typing +from typing import Any +from mujoco.mjx.warp import types + +if not typing.TYPE_CHECKING: + # Runtime. + warp: Any = None + mujoco_warp: Any = None + mjwp_types: Any = None + WARP_INSTALLED: bool = False + + # pylint: disable=g-import-not-at-top + try: + import warp + WARP_INSTALLED = True + except ImportError: + WARP_INSTALLED = False + try: + import mujoco.mjx.third_party.mujoco_warp as mujoco_warp + from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types + except ImportError as e: + pass + # pylint: enable=g-import-not-at-top +else: + # Only used for type checking. + class _WpStub: + + def ScopedDevice(self, device: str): # pylint: disable=invalid-name + pass + + def types(self): + pass + + class _MjwpStub: + + def put_model(self, *args, **kwargs): + pass + + def make_data(self, *args, **kwargs): + pass + + def step(self, *args, **kwargs): + pass + + class _MjwpTypesStub: + def TileSet(self, *args, **kwargs): # pylint: disable=invalid-name + pass + + def BlockDim(self, *args, **kwargs): # pylint: disable=invalid-name + pass + + WARP_INSTALLED: bool = True + warp: Any = _WpStub() + mujoco_warp: Any = _MjwpStub() + mjwp_types: Any = _MjwpTypesStub() diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py new file mode 100644 index 00000000..0b86774e --- /dev/null +++ b/mjx/mujoco/mjx/warp/collision_driver.py @@ -0,0 +1,498 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""DO NOT EDIT. This file is auto-generated.""" +import dataclasses +import jax +from mujoco.mjx._src import types +from mujoco.mjx.warp import ffi +import mujoco.mjx.third_party.mujoco_warp as mjwarp +from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types +import warp as wp + + +_m = mjwarp.Model( + **{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init} +) +_d = mjwarp.Data( + **{f.name: None for f in dataclasses.fields(mjwarp.Data) if f.init} +) +_o = mjwarp.Option( + **{f.name: None for f in dataclasses.fields(mjwarp.Option) if f.init} +) +_s = mjwarp.Statistic( + **{f.name: None for f in dataclasses.fields(mjwarp.Statistic) if f.init} +) +_c = mjwarp.Contact( + **{f.name: None for f in dataclasses.fields(mjwarp.Contact) if f.init} +) +_e = mjwarp.Constraint( + **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} +) + + +@ffi.format_args_for_warp +def _collision_shim( + # Model + nworld: int, + block_dim: mjwp_types.BlockDim, + geom_aabb: wp.array2d(dtype=wp.vec3), + geom_condim: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_friction: wp.array2d(dtype=wp.vec3), + geom_gap: wp.array2d(dtype=float), + geom_margin: wp.array2d(dtype=float), + geom_pair_type_count: tuple[int, ...], + geom_plugin_index: wp.array(dtype=int), + geom_pos: wp.array2d(dtype=wp.vec3), + geom_priority: wp.array(dtype=int), + geom_quat: wp.array2d(dtype=wp.quat), + geom_rbound: wp.array2d(dtype=float), + geom_size: wp.array2d(dtype=wp.vec3), + geom_solimp: wp.array2d(dtype=mjwp_types.vec5), + geom_solmix: wp.array2d(dtype=float), + geom_solref: wp.array2d(dtype=wp.vec2), + geom_type: wp.array(dtype=int), + geompair2hfgeompair: wp.array(dtype=int), + has_sdf_geom: bool, + hfield_adr: wp.array(dtype=int), + hfield_data: wp.array(dtype=float), + hfield_ncol: wp.array(dtype=int), + hfield_nrow: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + mesh_graph: wp.array(dtype=int), + mesh_graphadr: wp.array(dtype=int), + mesh_polyadr: wp.array(dtype=int), + mesh_polymap: wp.array(dtype=int), + mesh_polymapadr: wp.array(dtype=int), + mesh_polymapnum: wp.array(dtype=int), + mesh_polynormal: wp.array(dtype=wp.vec3), + mesh_polynum: wp.array(dtype=int), + mesh_polyvert: wp.array(dtype=int), + mesh_polyvertadr: wp.array(dtype=int), + mesh_polyvertnum: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_vertadr: wp.array(dtype=int), + mesh_vertnum: wp.array(dtype=int), + ngeom: int, + nhfield: int, + nxn_geom_pair_filtered: wp.array(dtype=wp.vec2i), + nxn_pairid: wp.array(dtype=int), + nxn_pairid_filtered: wp.array(dtype=int), + pair_dim: wp.array(dtype=int), + pair_friction: wp.array2d(dtype=mjwp_types.vec5), + pair_gap: wp.array2d(dtype=float), + pair_margin: wp.array2d(dtype=float), + pair_solimp: wp.array2d(dtype=mjwp_types.vec5), + pair_solref: wp.array2d(dtype=wp.vec2), + pair_solreffriction: wp.array2d(dtype=wp.vec2), + plugin: wp.array(dtype=int), + plugin_attr: wp.array(dtype=wp.vec3f), + opt__broadphase: int, + opt__broadphase_filter: int, + opt__disableflags: int, + opt__epa_iterations: int, + opt__gjk_iterations: int, + opt__graph_conditional: bool, + opt__sdf_initpoints: int, + opt__sdf_iterations: int, + # Data + nconmax: int, + collision_hftri_index: wp.array(dtype=int), + collision_pair: wp.array(dtype=wp.vec2i), + collision_pairid: wp.array(dtype=int), + collision_worldid: wp.array(dtype=int), + epa_face: wp.array2d(dtype=wp.vec3i), + epa_horizon: wp.array2d(dtype=int), + epa_index: wp.array2d(dtype=int), + epa_map: wp.array2d(dtype=int), + epa_norm2: wp.array2d(dtype=float), + epa_pr: wp.array2d(dtype=wp.vec3), + epa_vert: wp.array2d(dtype=wp.vec3), + epa_vert1: wp.array2d(dtype=wp.vec3), + epa_vert2: wp.array2d(dtype=wp.vec3), + epa_vert_index1: wp.array2d(dtype=int), + epa_vert_index2: wp.array2d(dtype=int), + geom_xmat: wp.array2d(dtype=wp.mat33), + geom_xpos: wp.array2d(dtype=wp.vec3), + ncollision: wp.array(dtype=int), + ncon: wp.array(dtype=int), + ncon_hfield: wp.array2d(dtype=int), + sap_cumulative_sum: wp.array2d(dtype=int), + sap_projection_lower: wp.array3d(dtype=float), + sap_projection_upper: wp.array2d(dtype=float), + sap_range: wp.array2d(dtype=int), + sap_segment_index: wp.array2d(dtype=int), + sap_sort_index: wp.array3d(dtype=int), + contact__dim: wp.array(dtype=int), + contact__dist: wp.array(dtype=float), + contact__frame: wp.array(dtype=wp.mat33), + contact__friction: wp.array(dtype=mjwp_types.vec5), + contact__geom: wp.array(dtype=wp.vec2i), + contact__includemargin: wp.array(dtype=float), + contact__pos: wp.array(dtype=wp.vec3), + contact__solimp: wp.array(dtype=mjwp_types.vec5), + contact__solref: wp.array(dtype=wp.vec2), + contact__solreffriction: wp.array(dtype=wp.vec2), + contact__worldid: wp.array(dtype=int), +): + _m.stat = _s + _m.opt = _o + _d.efc = _e + _d.contact = _c + _m.block_dim = block_dim + _m.geom_aabb = geom_aabb + _m.geom_condim = geom_condim + _m.geom_dataid = geom_dataid + _m.geom_friction = geom_friction + _m.geom_gap = geom_gap + _m.geom_margin = geom_margin + _m.geom_pair_type_count = geom_pair_type_count + _m.geom_plugin_index = geom_plugin_index + _m.geom_pos = geom_pos + _m.geom_priority = geom_priority + _m.geom_quat = geom_quat + _m.geom_rbound = geom_rbound + _m.geom_size = geom_size + _m.geom_solimp = geom_solimp + _m.geom_solmix = geom_solmix + _m.geom_solref = geom_solref + _m.geom_type = geom_type + _m.geompair2hfgeompair = geompair2hfgeompair + _m.has_sdf_geom = has_sdf_geom + _m.hfield_adr = hfield_adr + _m.hfield_data = hfield_data + _m.hfield_ncol = hfield_ncol + _m.hfield_nrow = hfield_nrow + _m.hfield_size = hfield_size + _m.mesh_graph = mesh_graph + _m.mesh_graphadr = mesh_graphadr + _m.mesh_polyadr = mesh_polyadr + _m.mesh_polymap = mesh_polymap + _m.mesh_polymapadr = mesh_polymapadr + _m.mesh_polymapnum = mesh_polymapnum + _m.mesh_polynormal = mesh_polynormal + _m.mesh_polynum = mesh_polynum + _m.mesh_polyvert = mesh_polyvert + _m.mesh_polyvertadr = mesh_polyvertadr + _m.mesh_polyvertnum = mesh_polyvertnum + _m.mesh_vert = mesh_vert + _m.mesh_vertadr = mesh_vertadr + _m.mesh_vertnum = mesh_vertnum + _m.ngeom = ngeom + _m.nhfield = nhfield + _m.nxn_geom_pair_filtered = nxn_geom_pair_filtered + _m.nxn_pairid = nxn_pairid + _m.nxn_pairid_filtered = nxn_pairid_filtered + _m.opt.broadphase = opt__broadphase + _m.opt.broadphase_filter = opt__broadphase_filter + _m.opt.disableflags = opt__disableflags + _m.opt.epa_iterations = opt__epa_iterations + _m.opt.gjk_iterations = opt__gjk_iterations + _m.opt.graph_conditional = opt__graph_conditional + _m.opt.sdf_initpoints = opt__sdf_initpoints + _m.opt.sdf_iterations = opt__sdf_iterations + _m.pair_dim = pair_dim + _m.pair_friction = pair_friction + _m.pair_gap = pair_gap + _m.pair_margin = pair_margin + _m.pair_solimp = pair_solimp + _m.pair_solref = pair_solref + _m.pair_solreffriction = pair_solreffriction + _m.plugin = plugin + _m.plugin_attr = plugin_attr + _d.collision_hftri_index = collision_hftri_index + _d.collision_pair = collision_pair + _d.collision_pairid = collision_pairid + _d.collision_worldid = collision_worldid + _d.contact.dim = contact__dim + _d.contact.dist = contact__dist + _d.contact.frame = contact__frame + _d.contact.friction = contact__friction + _d.contact.geom = contact__geom + _d.contact.includemargin = contact__includemargin + _d.contact.pos = contact__pos + _d.contact.solimp = contact__solimp + _d.contact.solref = contact__solref + _d.contact.solreffriction = contact__solreffriction + _d.contact.worldid = contact__worldid + _d.epa_face = epa_face + _d.epa_horizon = epa_horizon + _d.epa_index = epa_index + _d.epa_map = epa_map + _d.epa_norm2 = epa_norm2 + _d.epa_pr = epa_pr + _d.epa_vert = epa_vert + _d.epa_vert1 = epa_vert1 + _d.epa_vert2 = epa_vert2 + _d.epa_vert_index1 = epa_vert_index1 + _d.epa_vert_index2 = epa_vert_index2 + _d.geom_xmat = geom_xmat + _d.geom_xpos = geom_xpos + _d.ncollision = ncollision + _d.ncon = ncon + _d.ncon_hfield = ncon_hfield + _d.nconmax = nconmax + _d.sap_cumulative_sum = sap_cumulative_sum + _d.sap_projection_lower = sap_projection_lower + _d.sap_projection_upper = sap_projection_upper + _d.sap_range = sap_range + _d.sap_segment_index = sap_segment_index + _d.sap_sort_index = sap_sort_index + _d.nworld = nworld + mjwarp.collision(_m, _d) + + +def _collision_jax_impl(m: types.Model, d: types.Data): + output_dims = { + 'collision_hftri_index': d._impl.collision_hftri_index.shape, + 'collision_pair': d._impl.collision_pair.shape, + 'collision_pairid': d._impl.collision_pairid.shape, + 'collision_worldid': d._impl.collision_worldid.shape, + 'epa_face': d._impl.epa_face.shape, + 'epa_horizon': d._impl.epa_horizon.shape, + 'epa_index': d._impl.epa_index.shape, + 'epa_map': d._impl.epa_map.shape, + 'epa_norm2': d._impl.epa_norm2.shape, + 'epa_pr': d._impl.epa_pr.shape, + 'epa_vert': d._impl.epa_vert.shape, + 'epa_vert1': d._impl.epa_vert1.shape, + 'epa_vert2': d._impl.epa_vert2.shape, + 'epa_vert_index1': d._impl.epa_vert_index1.shape, + 'epa_vert_index2': d._impl.epa_vert_index2.shape, + 'geom_xmat': d.geom_xmat.shape, + 'geom_xpos': d.geom_xpos.shape, + 'ncollision': d._impl.ncollision.shape, + 'ncon': d._impl.ncon.shape, + 'ncon_hfield': d._impl.ncon_hfield.shape, + 'sap_cumulative_sum': d._impl.sap_cumulative_sum.shape, + 'sap_projection_lower': d._impl.sap_projection_lower.shape, + 'sap_projection_upper': d._impl.sap_projection_upper.shape, + 'sap_range': d._impl.sap_range.shape, + 'sap_segment_index': d._impl.sap_segment_index.shape, + 'sap_sort_index': d._impl.sap_sort_index.shape, + 'contact__dim': d._impl.contact__dim.shape, + 'contact__dist': d._impl.contact__dist.shape, + 'contact__frame': d._impl.contact__frame.shape, + 'contact__friction': d._impl.contact__friction.shape, + 'contact__geom': d._impl.contact__geom.shape, + 'contact__includemargin': d._impl.contact__includemargin.shape, + 'contact__pos': d._impl.contact__pos.shape, + 'contact__solimp': d._impl.contact__solimp.shape, + 'contact__solref': d._impl.contact__solref.shape, + 'contact__solreffriction': d._impl.contact__solreffriction.shape, + 'contact__worldid': d._impl.contact__worldid.shape, + } + jf = ffi.jax_callable_variadic_tuple( + _collision_shim, + num_outputs=37, + output_dims=output_dims, + vmap_method=None, + graph_compatible=True, + in_out_argnames={ + 'collision_hftri_index', + 'collision_pair', + 'collision_pairid', + 'collision_worldid', + 'epa_face', + 'epa_horizon', + 'epa_index', + 'epa_map', + 'epa_norm2', + 'epa_pr', + 'epa_vert', + 'epa_vert1', + 'epa_vert2', + 'epa_vert_index1', + 'epa_vert_index2', + 'geom_xmat', + 'geom_xpos', + 'ncollision', + 'ncon', + 'ncon_hfield', + 'sap_cumulative_sum', + 'sap_projection_lower', + 'sap_projection_upper', + 'sap_range', + 'sap_segment_index', + 'sap_sort_index', + 'contact__dim', + 'contact__dist', + 'contact__frame', + 'contact__friction', + 'contact__geom', + 'contact__includemargin', + 'contact__pos', + 'contact__solimp', + 'contact__solref', + 'contact__solreffriction', + 'contact__worldid', + }, + ) + out = jf( + d.qpos.shape[0], + m._impl.block_dim, + m.geom_aabb, + m.geom_condim, + m.geom_dataid, + m.geom_friction, + m.geom_gap, + m.geom_margin, + m._impl.geom_pair_type_count, + m._impl.geom_plugin_index, + m.geom_pos, + m.geom_priority, + m.geom_quat, + m.geom_rbound, + m.geom_size, + m.geom_solimp, + m.geom_solmix, + m.geom_solref, + m.geom_type, + m._impl.geompair2hfgeompair, + m._impl.has_sdf_geom, + m.hfield_adr, + m.hfield_data, + m.hfield_ncol, + m.hfield_nrow, + m.hfield_size, + m.mesh_graph, + m.mesh_graphadr, + m._impl.mesh_polyadr, + m._impl.mesh_polymap, + m._impl.mesh_polymapadr, + m._impl.mesh_polymapnum, + m._impl.mesh_polynormal, + m._impl.mesh_polynum, + m._impl.mesh_polyvert, + m._impl.mesh_polyvertadr, + m._impl.mesh_polyvertnum, + m.mesh_vert, + m.mesh_vertadr, + m.mesh_vertnum, + m.ngeom, + m.nhfield, + m._impl.nxn_geom_pair_filtered, + m._impl.nxn_pairid, + m._impl.nxn_pairid_filtered, + m.pair_dim, + m.pair_friction, + m.pair_gap, + m.pair_margin, + m.pair_solimp, + m.pair_solref, + m.pair_solreffriction, + m._impl.plugin, + m._impl.plugin_attr, + m.opt._impl.broadphase, + m.opt._impl.broadphase_filter, + m.opt.disableflags, + m.opt._impl.epa_iterations, + m.opt._impl.gjk_iterations, + m.opt._impl.graph_conditional, + m.opt._impl.sdf_initpoints, + m.opt._impl.sdf_iterations, + d._impl.nconmax, + d._impl.collision_hftri_index, + d._impl.collision_pair, + d._impl.collision_pairid, + d._impl.collision_worldid, + d._impl.epa_face, + d._impl.epa_horizon, + d._impl.epa_index, + d._impl.epa_map, + d._impl.epa_norm2, + d._impl.epa_pr, + d._impl.epa_vert, + d._impl.epa_vert1, + d._impl.epa_vert2, + d._impl.epa_vert_index1, + d._impl.epa_vert_index2, + d.geom_xmat, + d.geom_xpos, + d._impl.ncollision, + d._impl.ncon, + d._impl.ncon_hfield, + d._impl.sap_cumulative_sum, + d._impl.sap_projection_lower, + d._impl.sap_projection_upper, + d._impl.sap_range, + d._impl.sap_segment_index, + d._impl.sap_sort_index, + d._impl.contact__dim, + d._impl.contact__dist, + d._impl.contact__frame, + d._impl.contact__friction, + d._impl.contact__geom, + d._impl.contact__includemargin, + d._impl.contact__pos, + d._impl.contact__solimp, + d._impl.contact__solref, + d._impl.contact__solreffriction, + d._impl.contact__worldid, + ) + d = d.tree_replace({ + '_impl.collision_hftri_index': out[0], + '_impl.collision_pair': out[1], + '_impl.collision_pairid': out[2], + '_impl.collision_worldid': out[3], + '_impl.epa_face': out[4], + '_impl.epa_horizon': out[5], + '_impl.epa_index': out[6], + '_impl.epa_map': out[7], + '_impl.epa_norm2': out[8], + '_impl.epa_pr': out[9], + '_impl.epa_vert': out[10], + '_impl.epa_vert1': out[11], + '_impl.epa_vert2': out[12], + '_impl.epa_vert_index1': out[13], + '_impl.epa_vert_index2': out[14], + 'geom_xmat': out[15], + 'geom_xpos': out[16], + '_impl.ncollision': out[17], + '_impl.ncon': out[18], + '_impl.ncon_hfield': out[19], + '_impl.sap_cumulative_sum': out[20], + '_impl.sap_projection_lower': out[21], + '_impl.sap_projection_upper': out[22], + '_impl.sap_range': out[23], + '_impl.sap_segment_index': out[24], + '_impl.sap_sort_index': out[25], + '_impl.contact__dim': out[26], + '_impl.contact__dist': out[27], + '_impl.contact__frame': out[28], + '_impl.contact__friction': out[29], + '_impl.contact__geom': out[30], + '_impl.contact__includemargin': out[31], + '_impl.contact__pos': out[32], + '_impl.contact__solimp': out[33], + '_impl.contact__solref': out[34], + '_impl.contact__solreffriction': out[35], + '_impl.contact__worldid': out[36], + }) + return d + + +@jax.custom_batching.custom_vmap +@ffi.marshal_jax_warp_callable +def collision(m: types.Model, d: types.Data): + return _collision_jax_impl(m, d) + + +@collision.def_vmap +@ffi.marshal_custom_vmap +def collision_vmap(unused_axis_size, is_batched, m, d): + d = collision(m, d) + return d, is_batched[1] diff --git a/mjx/mujoco/mjx/warp/collision_driver_test.py b/mjx/mujoco/mjx/warp/collision_driver_test.py new file mode 100644 index 00000000..e69227f9 --- /dev/null +++ b/mjx/mujoco/mjx/warp/collision_driver_test.py @@ -0,0 +1,107 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Tests for collision driver.""" + +from absl.testing import absltest +import jax +import mujoco +from mujoco import mjx +from mujoco.mjx._src import io +import mujoco.mjx.warp as mjxw +from mujoco.mjx.warp import test_util as tu +from mujoco.mjx.warp import warp as wp # pylint: disable=g-importing-member +import numpy as np + + +try: + from mujoco.mjx.warp import collision_driver # pylint: disable=g-import-not-at-top + from mujoco.mjx.warp import smooth # pylint: disable=g-import-not-at-top +except ImportError: + collision_driver = None + smooth = None + + +class CollisionTest(absltest.TestCase): + + def setUp(self): + super().setUp() + if mjxw.WARP_INSTALLED: + wp.clear_kernel_cache() + np.random.seed(0) + + _SPHERE_SPHERE = """ + + + + + + + + + + + + + """ + + def test_collision_nested_vmap(self): + """Tests collision with batched data.""" + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp not installed.') + if not io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device available.') + + m = mujoco.MjModel.from_xml_string(self._SPHERE_SPHERE) + d = mujoco.MjData(m) + mx = mjx.put_model(m, impl='warp') + + def make_data(rng): + dx = mjx.make_data(m, impl='warp') + _, key = jax.random.split(rng) + qpos = jax.random.uniform(key, (m.nq,), minval=-0.01, maxval=0.01) + return dx.replace( + qpos=qpos, + ) + + rng = jax.random.split(jax.random.PRNGKey(0), 8) + rng = rng.reshape((2, 4, -1)) + dx_batch = jax.vmap(jax.vmap(make_data))(rng) + + dx_batch = jax.jit( + jax.vmap( + jax.vmap(smooth.kinematics, in_axes=(None, 0)), in_axes=(None, 0) + ) + )(mx, dx_batch) + dx_batch = jax.jit( + jax.vmap( + jax.vmap(collision_driver.collision, in_axes=(None, 0)), + in_axes=(None, 0), + ) + )(mx, dx_batch) + + for i in range(2): + for j in range(4): + dx = dx_batch[i, j] + + d.qpos[:] = dx.qpos + mujoco.mj_forward(m, d) + + if not d.contact.pos.shape[0]: + continue + tu.assert_contact_eq(d, dx, worldid=i * 4 + j) + + +if __name__ == '__main__': + absltest.main() diff --git a/mjx/mujoco/mjx/warp/ffi.py b/mjx/mujoco/mjx/warp/ffi.py new file mode 100644 index 00000000..e6db9595 --- /dev/null +++ b/mjx/mujoco/mjx/warp/ffi.py @@ -0,0 +1,350 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""FFI helper functions for MJX.""" + +import dataclasses +import functools +import inspect +import typing +from typing import Any, Callable, Optional, Sequence, Tuple, Union + +import jax +from jax import numpy as jp +from mujoco.mjx.warp import types as mjx_warp_types +import numpy as np +import warp as wp +from mujoco.mjx.third_party.warp.jax_experimental import ffi + + +def flatten_signature(signature: inspect.Signature, args: Tuple[Any, ...]): + """Flattens a tuple/dataclass signature.""" + + def expand_parameter(parameter, arg_iter): + p = parameter + if p.kind != inspect.Parameter.POSITIONAL_OR_KEYWORD: + raise ValueError(f'Unsupported parameter kind: {p.kind}') + + try: + arg = next(arg_iter) + except StopIteration: + # We ran out input arguments. Let us keep output parameters as is. + assert typing.get_origin(p.annotation) != tuple, p.annotation + return [p] + + # If it is a tuple, we need to duplicate the parameter for each element. + if isinstance(arg, tuple): + assert typing.get_origin(p.annotation) == tuple, p.annotation + type_args = typing.get_args(p.annotation) + if len(type_args) != 2 or type_args[1] != ...: + raise NotImplementedError( + f'Unsupported tuple argument: {type_args} ' + '(currently, only Tuple[t, ...] is supported).' + ) + type_ = type_args[0] + types = [type_] + if dataclasses.is_dataclass(type_): + # If the tuple element is a dataclass, we need to recycle the + # types in the same order as they appear in the dataclass. + fields = list(type_.__dataclass_fields__.values()) + types = [f.type for f in fields] + return [ + inspect.Parameter( + f'{p.name}__{i}', + p.kind, + default=p.default, + annotation=types[i % len(types)], + ) + for i in range(len(arg) * len(types)) + ] + elif dataclasses.is_dataclass(arg): + assert dataclasses.is_dataclass(p.annotation), p.annotation + fields = list(arg.__dataclass_fields__.values()) + return [ + inspect.Parameter( + f'{p.name}__{fields[i].name}', + p.kind, + default=p.default, + annotation=fields[i].type, + ) + for i in range(len(fields)) + ] + + assert typing.get_origin(p.annotation) != tuple, p.annotation + return [p] + + parameters = [] + arg_iter = iter(args) + for p in signature.parameters.values(): + parameters.extend(expand_parameter(p, arg_iter)) + + return inspect.Signature( + parameters=parameters, return_annotation=signature.return_annotation + ) + + +def jax_callable_variadic_tuple( + func: Callable, # pylint: disable=g-bare-generic + num_outputs: int = 1, + graph_compatible: bool = True, + vmap_method: Optional[str] = None, + output_dims: Optional[dict[str, tuple[int, ...]]] = None, + in_out_argnames: Optional[Sequence[str]] = None, +): + """Wraps a JAX callable to support variadic tuples and dataclasses.""" + + def callable_wrapper(*args, **kwargs): + def func_wrapper(*flat_args, **kwargs): + unflat_args = jax.tree.unflatten(in_tree, flat_args) + return func(*unflat_args, **kwargs) + + # Provide a flattened signature for the Warp callable machinery. + func_wrapper.__signature__ = flatten_signature( + inspect.signature(func), args + ) + my_callable = ffi.jax_callable( + func_wrapper, + num_outputs=num_outputs, + graph_compatible=graph_compatible, + vmap_method=vmap_method, + output_dims=output_dims, + in_out_argnames=in_out_argnames, + ) + + flat_args, in_tree = jax.tree.flatten(args) + return my_callable(*flat_args, **kwargs) + + return callable_wrapper + + +def _format_arg(arg: Any, name: str, annotation: Any, verbose: bool): + """Formats a single argument for warp.""" + typ_args = typing.get_args(annotation) + annotation_origin = typing.get_origin(annotation) + + # Handle variadic tuples. + if annotation_origin == tuple and len(typ_args) == 2 and typ_args[1] == ...: + return tuple( + _format_arg(arg[i], name + f'_{i}', typ_args[0], verbose) + for i in range(len(arg)) + ) + + if not isinstance(annotation, wp.types.array): + if verbose: + print(f'Skipping {name}: {arg}') + return arg + + expected_ndim = annotation.ndim + if arg.ndim != expected_ndim: + raise AssertionError( + f'Arg ndim {arg.ndim} does not match expected ndim {expected_ndim}.' + ) + + # Add stride 0 to first axis in case the underlying argument should be + # batched. + # NB: the outer marshalling does an "expand_dims" on Model fields. + is_batch_field = mjx_warp_types.BATCH_DIM['Model'].get(name, False) + if arg.shape[0] == 1 and is_batch_field: + old_strides = arg.strides + arg.strides = (0,) + arg.strides[1:] + if verbose: + print( + f'Leading batch dim of 1, adding stride: {name} {old_strides} =>' + f' {arg.strides}' + ) + return arg + + if verbose: + print(f'Did nothing: {name}: {arg.shape}') + return arg + + +def format_args_for_warp(func, verbose=False): + @functools.wraps(func) + def wrapper(*args): + args = list(args) + annotations = func.__annotations__ + assert len(args) == len(annotations) + for i, (name, annotation) in enumerate(annotations.items()): + args[i] = _format_arg(args[i], name, annotation, verbose) + return func(*args) + + return wrapper + + +def _get_mapping_from_tree_path( + path: jax.tree_util.KeyPath, + mapping: dict[str, int], +) -> Optional[int]: + """Gets the mapped value from a tree path.""" + if not isinstance(path, tuple): + raise NotImplementedError( + f'Parsing for jax tree path {path} not implemented.' + ) + + if any(isinstance(p, jax.tree_util.SequenceKey) for p in path): + # get the path up to the first sequence key, we assume variadic sequences + is_seq_key = [isinstance(p, jax.tree_util.SequenceKey) for p in path] + path = path[: is_seq_key.index(True)] + + assert all(isinstance(p, jax.tree_util.GetAttrKey) for p in path) + path = [p for p in path if p.name != '_impl'] + attr = '__'.join(p.name for p in path) + + # None if the MJX public field is not present in the MJX-Warp mapping. + return mapping.get(attr) + + +def _expand_dim_from_path( + path: jax.tree_util.KeyPath, leaf: Any, ndim_map: dict[str, int] +) -> Any: + """Expands the dimension of a leaf node based on the ndim_map.""" + ndim = _get_mapping_from_tree_path(path, ndim_map) + if ndim is None or ndim < 0: + return leaf + if ndim > leaf.ndim: + leaf = jp.expand_dims(leaf, axis=np.arange(ndim - leaf.ndim)) + if ndim != leaf.ndim: + raise AssertionError( + f'Leaf node ndim ({leaf.ndim}) and expected ndim ({ndim}) do not match' + f' for path {path}.' + ) + return leaf + + +def _squeeze_dim(leaf_expanded: Any, leaf: Any) -> Any: + if leaf_expanded.ndim < leaf.ndim: + raise AssertionError( + f'Expanded leaf ndim {leaf_expanded.ndim} is smaller than original leaf' + f' ndim {leaf.ndim}' + ) + if leaf_expanded.ndim > leaf.ndim: + return jp.squeeze(leaf_expanded, np.arange(leaf_expanded.ndim - leaf.ndim)) + return leaf_expanded + + +def marshal_jax_warp_callable(func): + """Marshal fields into a MuJoCo Warp function.""" + + @functools.wraps(func) + def wrapper(m, d): + # Expand dims for Warp implicit vmap before calling into the FFI wrapped + # function. + m_expanded = jax.tree.map_with_path( + lambda path, x: _expand_dim_from_path( + path, x, mjx_warp_types.NDIM['Model'] + ), + m, + ) + d_expanded = jax.tree.map_with_path( + lambda path, x: _expand_dim_from_path( + path, x, mjx_warp_types.NDIM['Data'] + ), + d, + ) + d_expanded_result = func(m_expanded, d_expanded) + d_result = jax.tree.map(_squeeze_dim, d_expanded_result, d) + return d_result + + return wrapper + + +def _flatten_batch_dim( + path: jax.tree_util.KeyPath, leaf: Any, ndim_map: dict[str, int] +) -> Any: + ndim = _get_mapping_from_tree_path(path, ndim_map) + if ndim is None or ndim < 0: + return leaf + if ndim < leaf.ndim: + assert leaf.ndim - ndim == 1 + batch_dim = np.prod(leaf.shape[: leaf.ndim - ndim + 1]) + leaf = jp.reshape(leaf, (batch_dim,) + leaf.shape[leaf.ndim - ndim + 1 :]) + return leaf + + +def _unflatten_batch_dim(leaf_squeezed: Any, leaf: Any) -> Any: + if leaf_squeezed.ndim > leaf.ndim: + raise AssertionError( + f'Squeezed leaf ndim {leaf_squeezed.ndim} is greater than original leaf' + f' ndim {leaf.ndim}' + ) + if leaf_squeezed.ndim < leaf.ndim: + return leaf_squeezed.reshape(leaf.shape) + return leaf_squeezed + + +def _maybe_broadcast_to( + path: jax.tree_util.KeyPath, + leaf: Any, + is_batched: bool, + axis_size: Union[int, tuple[int, ...]], + cls_str: str, +) -> Any: + """Broadcasts fields that are used in MuJoCo Warp.""" + ndim = _get_mapping_from_tree_path(path, mjx_warp_types.NDIM[cls_str]) + needs_batch_dim = _get_mapping_from_tree_path( + path, mjx_warp_types.BATCH_DIM[cls_str] + ) + needs_batch_dim = bool(needs_batch_dim) and (ndim is not None and ndim > 0) + if needs_batch_dim and not is_batched: + leaf = jp.broadcast_to(leaf, (axis_size,) + leaf.shape) + return leaf + + +def marshal_custom_vmap(vmap_func): + """Marshal fields for a custom vmap into an MuJoCo Warp function.""" + + @functools.wraps(vmap_func) + def wrapper(axis_size, is_batched, m, d): + # Vmappable data fields may not have been broadcasted if vmap_func is called + # within a vmap trace. Since data fields are read/write in warp, we need to + # explicitly broadcast them here. + d_broadcast = jax.tree.map_with_path( + lambda path, x, is_b: _maybe_broadcast_to( + path, x, is_b, axis_size, 'Data' + ), + d, is_batched[1], # fmt: skip + ) + # Flatten batch dims into the first axis if the vmap was nested. + m_flat = jax.tree.map_with_path( + lambda path, x: _flatten_batch_dim( + path, x, mjx_warp_types.NDIM['Model'] + ), + m, + ) + d_broadcast_flat = jax.tree.map_with_path( + lambda path, x: _flatten_batch_dim( + path, x, mjx_warp_types.NDIM['Data'] + ), + d_broadcast, + ) + d_broadcast_flat_result, out_batched = vmap_func( + axis_size, is_batched, m_flat, d_broadcast_flat + ) + # Explicitly mark MuJoCo Warp data fields as batched after vmapping is done. + out_batched = jax.tree.map_with_path( + # NB: if a field is not in MuJoCo Warp, we let JAX do its magic. + lambda path, x: _get_mapping_from_tree_path( + path, mjx_warp_types.BATCH_DIM['Data'] + ) + or x, + out_batched, + ) + # Unflatten batch dimensions but keep the broadcasting. + d_result = jax.tree.map( + _unflatten_batch_dim, d_broadcast_flat_result, d_broadcast + ) + return d_result, out_batched + + return wrapper diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py new file mode 100644 index 00000000..8238c24d --- /dev/null +++ b/mjx/mujoco/mjx/warp/forward.py @@ -0,0 +1,4123 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""DO NOT EDIT. This file is auto-generated.""" +import dataclasses +import jax +from mujoco.mjx._src import types +from mujoco.mjx.warp import ffi +import mujoco.mjx.third_party.mujoco_warp as mjwarp +from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types +import warp as wp + + +_m = mjwarp.Model( + **{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init} +) +_d = mjwarp.Data( + **{f.name: None for f in dataclasses.fields(mjwarp.Data) if f.init} +) +_o = mjwarp.Option( + **{f.name: None for f in dataclasses.fields(mjwarp.Option) if f.init} +) +_s = mjwarp.Statistic( + **{f.name: None for f in dataclasses.fields(mjwarp.Statistic) if f.init} +) +_c = mjwarp.Contact( + **{f.name: None for f in dataclasses.fields(mjwarp.Contact) if f.init} +) +_e = mjwarp.Constraint( + **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} +) + + +@ffi.format_args_for_warp +def _forward_shim( + # Model + nworld: int, + M_rowadr: wp.array(dtype=int), + M_rownnz: wp.array(dtype=int), + actuator_acc0: wp.array(dtype=float), + actuator_actadr: wp.array(dtype=int), + actuator_actearly: wp.array(dtype=bool), + actuator_actlimited: wp.array(dtype=bool), + actuator_actnum: wp.array(dtype=int), + actuator_actrange: wp.array2d(dtype=wp.vec2), + actuator_biasprm: wp.array2d(dtype=mjwp_types.vec10f), + actuator_biastype: wp.array(dtype=int), + actuator_cranklength: wp.array(dtype=float), + actuator_ctrllimited: wp.array(dtype=bool), + actuator_ctrlrange: wp.array2d(dtype=wp.vec2), + actuator_dynprm: wp.array2d(dtype=mjwp_types.vec10f), + actuator_dyntype: wp.array(dtype=int), + actuator_forcelimited: wp.array(dtype=bool), + actuator_forcerange: wp.array2d(dtype=wp.vec2), + actuator_gainprm: wp.array2d(dtype=mjwp_types.vec10f), + actuator_gaintype: wp.array(dtype=int), + actuator_gear: wp.array2d(dtype=wp.spatial_vector), + actuator_lengthrange: wp.array(dtype=wp.vec2), + actuator_trnid: wp.array(dtype=wp.vec2i), + actuator_trntype: wp.array(dtype=int), + actuator_trntype_body_adr: wp.array(dtype=int), + block_dim: mjwp_types.BlockDim, + body_dofadr: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_gravcomp: wp.array2d(dtype=float), + body_inertia: wp.array2d(dtype=wp.vec3), + body_invweight0: wp.array2d(dtype=wp.vec2), + body_ipos: wp.array2d(dtype=wp.vec3), + body_iquat: wp.array2d(dtype=wp.quat), + body_jntadr: wp.array(dtype=int), + body_jntnum: wp.array(dtype=int), + body_mass: wp.array2d(dtype=float), + body_parentid: wp.array(dtype=int), + body_pos: wp.array2d(dtype=wp.vec3), + body_quat: wp.array2d(dtype=wp.quat), + body_rootid: wp.array(dtype=int), + body_subtreemass: wp.array2d(dtype=float), + body_tree: tuple[wp.array(dtype=int), ...], + body_weldid: wp.array(dtype=int), + cam_bodyid: wp.array(dtype=int), + cam_fovy: wp.array(dtype=float), + cam_intrinsic: wp.array(dtype=wp.vec4), + cam_mat0: wp.array2d(dtype=wp.mat33), + cam_mode: wp.array(dtype=int), + cam_pos: wp.array2d(dtype=wp.vec3), + cam_pos0: wp.array2d(dtype=wp.vec3), + cam_poscom0: wp.array2d(dtype=wp.vec3), + cam_quat: wp.array2d(dtype=wp.quat), + cam_resolution: wp.array(dtype=wp.vec2i), + cam_sensorsize: wp.array(dtype=wp.vec2), + cam_targetbodyid: wp.array(dtype=int), + condim_max: int, + dof_Madr: wp.array(dtype=int), + dof_armature: wp.array2d(dtype=float), + dof_bodyid: wp.array(dtype=int), + dof_damping: wp.array2d(dtype=float), + dof_frictionloss: wp.array2d(dtype=float), + dof_invweight0: wp.array2d(dtype=float), + dof_jntid: wp.array(dtype=int), + dof_parentid: wp.array(dtype=int), + dof_solimp: wp.array2d(dtype=mjwp_types.vec5), + dof_solref: wp.array2d(dtype=wp.vec2), + dof_tri_col: wp.array(dtype=int), + dof_tri_row: wp.array(dtype=int), + eq_connect_adr: wp.array(dtype=int), + eq_data: wp.array2d(dtype=mjwp_types.vec11), + eq_jnt_adr: wp.array(dtype=int), + eq_obj1id: wp.array(dtype=int), + eq_obj2id: wp.array(dtype=int), + eq_objtype: wp.array(dtype=int), + eq_solimp: wp.array2d(dtype=mjwp_types.vec5), + eq_solref: wp.array2d(dtype=wp.vec2), + eq_ten_adr: wp.array(dtype=int), + eq_wld_adr: wp.array(dtype=int), + flex_bending: wp.array(dtype=wp.mat44f), + flex_damping: wp.array(dtype=float), + flex_dim: wp.array(dtype=int), + flex_edge: wp.array(dtype=wp.vec2i), + flex_edgeadr: wp.array(dtype=int), + flex_edgeflap: wp.array(dtype=wp.vec2i), + flex_elem: wp.array(dtype=int), + flex_elemedge: wp.array(dtype=int), + flex_elemedgeadr: wp.array(dtype=int), + flex_stiffness: wp.array(dtype=float), + flex_vertadr: wp.array(dtype=int), + flex_vertbodyid: wp.array(dtype=int), + flexedge_length0: wp.array(dtype=float), + geom_aabb: wp.array2d(dtype=wp.vec3), + geom_bodyid: wp.array(dtype=int), + geom_condim: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_friction: wp.array2d(dtype=wp.vec3), + geom_gap: wp.array2d(dtype=float), + geom_group: wp.array(dtype=int), + geom_margin: wp.array2d(dtype=float), + geom_matid: wp.array2d(dtype=int), + geom_pair_type_count: tuple[int, ...], + geom_plugin_index: wp.array(dtype=int), + geom_pos: wp.array2d(dtype=wp.vec3), + geom_priority: wp.array(dtype=int), + geom_quat: wp.array2d(dtype=wp.quat), + geom_rbound: wp.array2d(dtype=float), + geom_rgba: wp.array2d(dtype=wp.vec4), + geom_size: wp.array2d(dtype=wp.vec3), + geom_solimp: wp.array2d(dtype=mjwp_types.vec5), + geom_solmix: wp.array2d(dtype=float), + geom_solref: wp.array2d(dtype=wp.vec2), + geom_type: wp.array(dtype=int), + geompair2hfgeompair: wp.array(dtype=int), + has_sdf_geom: bool, + hfield_adr: wp.array(dtype=int), + hfield_data: wp.array(dtype=float), + hfield_ncol: wp.array(dtype=int), + hfield_nrow: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + jnt_actfrclimited: wp.array(dtype=bool), + jnt_actfrcrange: wp.array2d(dtype=wp.vec2), + jnt_actgravcomp: wp.array(dtype=int), + jnt_axis: wp.array2d(dtype=wp.vec3), + jnt_bodyid: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + jnt_limited_ball_adr: wp.array(dtype=int), + jnt_limited_slide_hinge_adr: wp.array(dtype=int), + jnt_margin: wp.array2d(dtype=float), + jnt_pos: wp.array2d(dtype=wp.vec3), + jnt_qposadr: wp.array(dtype=int), + jnt_range: wp.array2d(dtype=wp.vec2), + jnt_solimp: wp.array2d(dtype=mjwp_types.vec5), + jnt_solref: wp.array2d(dtype=wp.vec2), + jnt_stiffness: wp.array2d(dtype=float), + jnt_type: wp.array(dtype=int), + light_bodyid: wp.array(dtype=int), + light_dir: wp.array2d(dtype=wp.vec3), + light_dir0: wp.array2d(dtype=wp.vec3), + light_mode: wp.array(dtype=int), + light_pos: wp.array2d(dtype=wp.vec3), + light_pos0: wp.array2d(dtype=wp.vec3), + light_poscom0: wp.array2d(dtype=wp.vec3), + light_targetbodyid: wp.array(dtype=int), + mapM2M: wp.array(dtype=int), + mat_rgba: wp.array2d(dtype=wp.vec4), + mesh_face: wp.array(dtype=wp.vec3i), + mesh_faceadr: wp.array(dtype=int), + mesh_graph: wp.array(dtype=int), + mesh_graphadr: wp.array(dtype=int), + mesh_polyadr: wp.array(dtype=int), + mesh_polymap: wp.array(dtype=int), + mesh_polymapadr: wp.array(dtype=int), + mesh_polymapnum: wp.array(dtype=int), + mesh_polynormal: wp.array(dtype=wp.vec3), + mesh_polynum: wp.array(dtype=int), + mesh_polyvert: wp.array(dtype=int), + mesh_polyvertadr: wp.array(dtype=int), + mesh_polyvertnum: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_vertadr: wp.array(dtype=int), + mesh_vertnum: wp.array(dtype=int), + mocap_bodyid: wp.array(dtype=int), + nC: int, + na: int, + nbody: int, + ncam: int, + neq: int, + nflexedge: int, + nflexelem: int, + nflexvert: int, + ngeom: int, + ngravcomp: int, + nhfield: int, + njnt: int, + nlight: int, + nlsp: int, + nmeshface: int, + nmocap: int, + nsite: int, + ntendon: int, + nu: int, + nv: int, + nxn_geom_pair_filtered: wp.array(dtype=wp.vec2i), + nxn_pairid: wp.array(dtype=int), + nxn_pairid_filtered: wp.array(dtype=int), + pair_dim: wp.array(dtype=int), + pair_friction: wp.array2d(dtype=mjwp_types.vec5), + pair_gap: wp.array2d(dtype=float), + pair_margin: wp.array2d(dtype=float), + pair_solimp: wp.array2d(dtype=mjwp_types.vec5), + pair_solref: wp.array2d(dtype=wp.vec2), + pair_solreffriction: wp.array2d(dtype=wp.vec2), + plugin: wp.array(dtype=int), + plugin_attr: wp.array(dtype=wp.vec3f), + qLD_updates: tuple[wp.array(dtype=wp.vec3i), ...], + qM_fullm_i: wp.array(dtype=int), + qM_fullm_j: wp.array(dtype=int), + qM_madr_ij: wp.array(dtype=int), + qM_mulm_i: wp.array(dtype=int), + qM_mulm_j: wp.array(dtype=int), + qM_tiles: tuple[mjwp_types.TileSet, ...], + qpos0: wp.array2d(dtype=float), + qpos_spring: wp.array2d(dtype=float), + rangefinder_sensor_adr: wp.array(dtype=int), + sensor_acc_adr: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_cutoff: wp.array(dtype=float), + sensor_datatype: wp.array(dtype=int), + sensor_e_kinetic: bool, + sensor_e_potential: bool, + sensor_limitfrc_adr: wp.array(dtype=int), + sensor_limitpos_adr: wp.array(dtype=int), + sensor_limitvel_adr: wp.array(dtype=int), + sensor_objid: wp.array(dtype=int), + sensor_objtype: wp.array(dtype=int), + sensor_pos_adr: wp.array(dtype=int), + sensor_rangefinder_adr: wp.array(dtype=int), + sensor_rangefinder_bodyid: wp.array(dtype=int), + sensor_refid: wp.array(dtype=int), + sensor_reftype: wp.array(dtype=int), + sensor_rne_postconstraint: bool, + sensor_subtree_vel: bool, + sensor_tendonactfrc_adr: wp.array(dtype=int), + sensor_touch_adr: wp.array(dtype=int), + sensor_type: wp.array(dtype=int), + sensor_vel_adr: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + site_pos: wp.array2d(dtype=wp.vec3), + site_quat: wp.array2d(dtype=wp.quat), + site_size: wp.array(dtype=wp.vec3), + site_type: wp.array(dtype=int), + subtree_mass: wp.array2d(dtype=float), + tendon_actfrclimited: wp.array(dtype=bool), + tendon_actfrcrange: wp.array2d(dtype=wp.vec2), + tendon_adr: wp.array(dtype=int), + tendon_armature: wp.array2d(dtype=float), + tendon_damping: wp.array2d(dtype=float), + tendon_frictionloss: wp.array2d(dtype=float), + tendon_geom_adr: wp.array(dtype=int), + tendon_invweight0: wp.array2d(dtype=float), + tendon_jnt_adr: wp.array(dtype=int), + tendon_length0: wp.array2d(dtype=float), + tendon_lengthspring: wp.array2d(dtype=wp.vec2), + tendon_limited_adr: wp.array(dtype=int), + tendon_margin: wp.array2d(dtype=float), + tendon_num: wp.array(dtype=int), + tendon_range: wp.array2d(dtype=wp.vec2), + tendon_site_pair_adr: wp.array(dtype=int), + tendon_solimp_fri: wp.array2d(dtype=mjwp_types.vec5), + tendon_solimp_lim: wp.array2d(dtype=mjwp_types.vec5), + tendon_solref_fri: wp.array2d(dtype=wp.vec2), + tendon_solref_lim: wp.array2d(dtype=wp.vec2), + tendon_stiffness: wp.array2d(dtype=float), + wrap_geom_adr: wp.array(dtype=int), + wrap_jnt_adr: wp.array(dtype=int), + wrap_objid: wp.array(dtype=int), + wrap_prm: wp.array(dtype=float), + wrap_pulley_scale: wp.array(dtype=float), + wrap_site_pair_adr: wp.array(dtype=int), + wrap_type: wp.array(dtype=int), + opt__broadphase: int, + opt__broadphase_filter: int, + opt__cone: int, + opt__density: wp.array(dtype=float), + opt__disableflags: int, + opt__enableflags: int, + opt__epa_iterations: int, + opt__gjk_iterations: int, + opt__graph_conditional: bool, + opt__gravity: wp.array(dtype=wp.vec3), + opt__has_fluid: bool, + opt__impratio: wp.array(dtype=float), + opt__is_sparse: bool, + opt__iterations: int, + opt__ls_iterations: int, + opt__ls_parallel: bool, + opt__ls_tolerance: wp.array(dtype=float), + opt__magnetic: wp.array(dtype=wp.vec3), + opt__run_collision_detection: bool, + opt__sdf_initpoints: int, + opt__sdf_iterations: int, + opt__solver: int, + opt__timestep: wp.array(dtype=float), + opt__tolerance: wp.array(dtype=float), + opt__viscosity: wp.array(dtype=float), + opt__wind: wp.array(dtype=wp.vec3), + stat__meaninertia: float, + # Data + nconmax: int, + njmax: int, + act: wp.array2d(dtype=float), + act_dot: wp.array2d(dtype=float), + actuator_force: wp.array2d(dtype=float), + actuator_length: wp.array2d(dtype=float), + actuator_moment: wp.array3d(dtype=float), + actuator_trntype_body_ncon: wp.array2d(dtype=int), + actuator_velocity: wp.array2d(dtype=float), + cacc: wp.array2d(dtype=wp.spatial_vector), + cam_xmat: wp.array2d(dtype=wp.mat33), + cam_xpos: wp.array2d(dtype=wp.vec3), + cdof: wp.array2d(dtype=wp.spatial_vector), + cdof_dot: wp.array2d(dtype=wp.spatial_vector), + cfrc_ext: wp.array2d(dtype=wp.spatial_vector), + cfrc_int: wp.array2d(dtype=wp.spatial_vector), + cinert: wp.array2d(dtype=mjwp_types.vec10), + collision_hftri_index: wp.array(dtype=int), + collision_pair: wp.array(dtype=wp.vec2i), + collision_pairid: wp.array(dtype=int), + collision_worldid: wp.array(dtype=int), + crb: wp.array2d(dtype=mjwp_types.vec10), + ctrl: wp.array2d(dtype=float), + cvel: wp.array2d(dtype=wp.spatial_vector), + energy: wp.array(dtype=wp.vec2), + epa_face: wp.array2d(dtype=wp.vec3i), + epa_horizon: wp.array2d(dtype=int), + epa_index: wp.array2d(dtype=int), + epa_map: wp.array2d(dtype=int), + epa_norm2: wp.array2d(dtype=float), + epa_pr: wp.array2d(dtype=wp.vec3), + epa_vert: wp.array2d(dtype=wp.vec3), + epa_vert1: wp.array2d(dtype=wp.vec3), + epa_vert2: wp.array2d(dtype=wp.vec3), + epa_vert_index1: wp.array2d(dtype=int), + epa_vert_index2: wp.array2d(dtype=int), + eq_active: wp.array2d(dtype=bool), + flexedge_length: wp.array2d(dtype=float), + flexedge_velocity: wp.array2d(dtype=float), + flexvert_xpos: wp.array2d(dtype=wp.vec3), + fluid_applied: wp.array2d(dtype=wp.spatial_vector), + geom_skip: wp.array(dtype=bool), + geom_xmat: wp.array2d(dtype=wp.mat33), + geom_xpos: wp.array2d(dtype=wp.vec3), + light_xdir: wp.array2d(dtype=wp.vec3), + light_xpos: wp.array2d(dtype=wp.vec3), + mocap_pos: wp.array2d(dtype=wp.vec3), + mocap_quat: wp.array2d(dtype=wp.quat), + ncollision: wp.array(dtype=int), + ncon: wp.array(dtype=int), + ncon_hfield: wp.array2d(dtype=int), + ne: wp.array(dtype=int), + ne_connect: wp.array(dtype=int), + ne_jnt: wp.array(dtype=int), + ne_ten: wp.array(dtype=int), + ne_weld: wp.array(dtype=int), + nefc: wp.array(dtype=int), + nf: wp.array(dtype=int), + nl: wp.array(dtype=int), + nsolving: wp.array(dtype=int), + qLD: wp.array3d(dtype=float), + qLDiagInv: wp.array2d(dtype=float), + qM: wp.array3d(dtype=float), + qacc: wp.array2d(dtype=float), + qacc_smooth: wp.array2d(dtype=float), + qacc_warmstart: wp.array2d(dtype=float), + qfrc_actuator: wp.array2d(dtype=float), + qfrc_applied: wp.array2d(dtype=float), + qfrc_bias: wp.array2d(dtype=float), + qfrc_constraint: wp.array2d(dtype=float), + qfrc_damper: wp.array2d(dtype=float), + qfrc_fluid: wp.array2d(dtype=float), + qfrc_gravcomp: wp.array2d(dtype=float), + qfrc_passive: wp.array2d(dtype=float), + qfrc_smooth: wp.array2d(dtype=float), + qfrc_spring: wp.array2d(dtype=float), + qpos: wp.array2d(dtype=float), + qvel: wp.array2d(dtype=float), + sap_cumulative_sum: wp.array2d(dtype=int), + sap_projection_lower: wp.array3d(dtype=float), + sap_projection_upper: wp.array2d(dtype=float), + sap_range: wp.array2d(dtype=int), + sap_segment_index: wp.array2d(dtype=int), + sap_sort_index: wp.array3d(dtype=int), + sensor_rangefinder_dist: wp.array2d(dtype=float), + sensor_rangefinder_geomid: wp.array2d(dtype=int), + sensor_rangefinder_pnt: wp.array2d(dtype=wp.vec3), + sensor_rangefinder_vec: wp.array2d(dtype=wp.vec3), + sensordata: wp.array2d(dtype=float), + site_xmat: wp.array2d(dtype=wp.mat33), + site_xpos: wp.array2d(dtype=wp.vec3), + solver_niter: wp.array(dtype=int), + subtree_angmom: wp.array2d(dtype=wp.vec3), + subtree_bodyvel: wp.array2d(dtype=wp.spatial_vector), + subtree_com: wp.array2d(dtype=wp.vec3), + subtree_linvel: wp.array2d(dtype=wp.vec3), + ten_J: wp.array3d(dtype=float), + ten_Jdot: wp.array3d(dtype=float), + ten_actfrc: wp.array2d(dtype=float), + ten_bias_coef: wp.array2d(dtype=float), + ten_length: wp.array2d(dtype=float), + ten_velocity: wp.array2d(dtype=float), + ten_wrapadr: wp.array2d(dtype=int), + ten_wrapnum: wp.array2d(dtype=int), + time: wp.array(dtype=float), + wrap_geom_xpos: wp.array2d(dtype=wp.spatial_vector), + wrap_obj: wp.array2d(dtype=wp.vec2i), + wrap_xpos: wp.array2d(dtype=wp.spatial_vector), + xanchor: wp.array2d(dtype=wp.vec3), + xaxis: wp.array2d(dtype=wp.vec3), + xfrc_applied: wp.array2d(dtype=wp.spatial_vector), + ximat: wp.array2d(dtype=wp.mat33), + xipos: wp.array2d(dtype=wp.vec3), + xmat: wp.array2d(dtype=wp.mat33), + xpos: wp.array2d(dtype=wp.vec3), + xquat: wp.array2d(dtype=wp.quat), + contact__dim: wp.array(dtype=int), + contact__dist: wp.array(dtype=float), + contact__efc_address: wp.array2d(dtype=int), + contact__frame: wp.array(dtype=wp.mat33), + contact__friction: wp.array(dtype=mjwp_types.vec5), + contact__geom: wp.array(dtype=wp.vec2i), + contact__includemargin: wp.array(dtype=float), + contact__pos: wp.array(dtype=wp.vec3), + contact__solimp: wp.array(dtype=mjwp_types.vec5), + contact__solref: wp.array(dtype=wp.vec2), + contact__solreffriction: wp.array(dtype=wp.vec2), + contact__worldid: wp.array(dtype=int), + efc__D: wp.array2d(dtype=float), + efc__J: wp.array3d(dtype=float), + efc__Jaref: wp.array2d(dtype=float), + efc__Ma: wp.array2d(dtype=float), + efc__Mgrad: wp.array2d(dtype=float), + efc__active: wp.array2d(dtype=bool), + efc__alpha: wp.array(dtype=float), + efc__aref: wp.array2d(dtype=float), + efc__beta: wp.array(dtype=float), + efc__beta_den: wp.array(dtype=float), + efc__beta_num: wp.array(dtype=float), + efc__cholesky_L_tmp: wp.array3d(dtype=float), + efc__cholesky_y_tmp: wp.array2d(dtype=float), + efc__condim: wp.array2d(dtype=int), + efc__cost: wp.array(dtype=float), + efc__cost_candidate: wp.array2d(dtype=float), + efc__done: wp.array(dtype=bool), + efc__force: wp.array2d(dtype=float), + efc__frictionloss: wp.array2d(dtype=float), + efc__gauss: wp.array(dtype=float), + efc__grad: wp.array2d(dtype=float), + efc__grad_dot: wp.array(dtype=float), + efc__gtol: wp.array(dtype=float), + efc__h: wp.array3d(dtype=float), + efc__hi: wp.array(dtype=wp.vec3), + efc__hi_alpha: wp.array(dtype=float), + efc__hi_next: wp.array(dtype=wp.vec3), + efc__hi_next_alpha: wp.array(dtype=float), + efc__id: wp.array2d(dtype=int), + efc__jv: wp.array2d(dtype=float), + efc__lo: wp.array(dtype=wp.vec3), + efc__lo_alpha: wp.array(dtype=float), + efc__lo_next: wp.array(dtype=wp.vec3), + efc__lo_next_alpha: wp.array(dtype=float), + efc__ls_done: wp.array(dtype=bool), + efc__margin: wp.array2d(dtype=float), + efc__mid: wp.array(dtype=wp.vec3), + efc__mid_alpha: wp.array(dtype=float), + efc__mv: wp.array2d(dtype=float), + efc__p0: wp.array(dtype=wp.vec3), + efc__pos: wp.array2d(dtype=float), + efc__prev_Mgrad: wp.array2d(dtype=float), + efc__prev_cost: wp.array(dtype=float), + efc__prev_grad: wp.array2d(dtype=float), + efc__quad: wp.array2d(dtype=wp.vec3), + efc__quad_gauss: wp.array(dtype=wp.vec3), + efc__search: wp.array2d(dtype=float), + efc__search_dot: wp.array(dtype=float), + efc__type: wp.array2d(dtype=int), + efc__u: wp.array(dtype=mjwp_types.vec6), + efc__uu: wp.array(dtype=float), + efc__uv: wp.array(dtype=float), + efc__vel: wp.array2d(dtype=float), + efc__vv: wp.array(dtype=float), +): + _m.stat = _s + _m.opt = _o + _d.efc = _e + _d.contact = _c + _m.M_rowadr = M_rowadr + _m.M_rownnz = M_rownnz + _m.actuator_acc0 = actuator_acc0 + _m.actuator_actadr = actuator_actadr + _m.actuator_actearly = actuator_actearly + _m.actuator_actlimited = actuator_actlimited + _m.actuator_actnum = actuator_actnum + _m.actuator_actrange = actuator_actrange + _m.actuator_biasprm = actuator_biasprm + _m.actuator_biastype = actuator_biastype + _m.actuator_cranklength = actuator_cranklength + _m.actuator_ctrllimited = actuator_ctrllimited + _m.actuator_ctrlrange = actuator_ctrlrange + _m.actuator_dynprm = actuator_dynprm + _m.actuator_dyntype = actuator_dyntype + _m.actuator_forcelimited = actuator_forcelimited + _m.actuator_forcerange = actuator_forcerange + _m.actuator_gainprm = actuator_gainprm + _m.actuator_gaintype = actuator_gaintype + _m.actuator_gear = actuator_gear + _m.actuator_lengthrange = actuator_lengthrange + _m.actuator_trnid = actuator_trnid + _m.actuator_trntype = actuator_trntype + _m.actuator_trntype_body_adr = actuator_trntype_body_adr + _m.block_dim = block_dim + _m.body_dofadr = body_dofadr + _m.body_dofnum = body_dofnum + _m.body_gravcomp = body_gravcomp + _m.body_inertia = body_inertia + _m.body_invweight0 = body_invweight0 + _m.body_ipos = body_ipos + _m.body_iquat = body_iquat + _m.body_jntadr = body_jntadr + _m.body_jntnum = body_jntnum + _m.body_mass = body_mass + _m.body_parentid = body_parentid + _m.body_pos = body_pos + _m.body_quat = body_quat + _m.body_rootid = body_rootid + _m.body_subtreemass = body_subtreemass + _m.body_tree = body_tree + _m.body_weldid = body_weldid + _m.cam_bodyid = cam_bodyid + _m.cam_fovy = cam_fovy + _m.cam_intrinsic = cam_intrinsic + _m.cam_mat0 = cam_mat0 + _m.cam_mode = cam_mode + _m.cam_pos = cam_pos + _m.cam_pos0 = cam_pos0 + _m.cam_poscom0 = cam_poscom0 + _m.cam_quat = cam_quat + _m.cam_resolution = cam_resolution + _m.cam_sensorsize = cam_sensorsize + _m.cam_targetbodyid = cam_targetbodyid + _m.condim_max = condim_max + _m.dof_Madr = dof_Madr + _m.dof_armature = dof_armature + _m.dof_bodyid = dof_bodyid + _m.dof_damping = dof_damping + _m.dof_frictionloss = dof_frictionloss + _m.dof_invweight0 = dof_invweight0 + _m.dof_jntid = dof_jntid + _m.dof_parentid = dof_parentid + _m.dof_solimp = dof_solimp + _m.dof_solref = dof_solref + _m.dof_tri_col = dof_tri_col + _m.dof_tri_row = dof_tri_row + _m.eq_connect_adr = eq_connect_adr + _m.eq_data = eq_data + _m.eq_jnt_adr = eq_jnt_adr + _m.eq_obj1id = eq_obj1id + _m.eq_obj2id = eq_obj2id + _m.eq_objtype = eq_objtype + _m.eq_solimp = eq_solimp + _m.eq_solref = eq_solref + _m.eq_ten_adr = eq_ten_adr + _m.eq_wld_adr = eq_wld_adr + _m.flex_bending = flex_bending + _m.flex_damping = flex_damping + _m.flex_dim = flex_dim + _m.flex_edge = flex_edge + _m.flex_edgeadr = flex_edgeadr + _m.flex_edgeflap = flex_edgeflap + _m.flex_elem = flex_elem + _m.flex_elemedge = flex_elemedge + _m.flex_elemedgeadr = flex_elemedgeadr + _m.flex_stiffness = flex_stiffness + _m.flex_vertadr = flex_vertadr + _m.flex_vertbodyid = flex_vertbodyid + _m.flexedge_length0 = flexedge_length0 + _m.geom_aabb = geom_aabb + _m.geom_bodyid = geom_bodyid + _m.geom_condim = geom_condim + _m.geom_dataid = geom_dataid + _m.geom_friction = geom_friction + _m.geom_gap = geom_gap + _m.geom_group = geom_group + _m.geom_margin = geom_margin + _m.geom_matid = geom_matid + _m.geom_pair_type_count = geom_pair_type_count + _m.geom_plugin_index = geom_plugin_index + _m.geom_pos = geom_pos + _m.geom_priority = geom_priority + _m.geom_quat = geom_quat + _m.geom_rbound = geom_rbound + _m.geom_rgba = geom_rgba + _m.geom_size = geom_size + _m.geom_solimp = geom_solimp + _m.geom_solmix = geom_solmix + _m.geom_solref = geom_solref + _m.geom_type = geom_type + _m.geompair2hfgeompair = geompair2hfgeompair + _m.has_sdf_geom = has_sdf_geom + _m.hfield_adr = hfield_adr + _m.hfield_data = hfield_data + _m.hfield_ncol = hfield_ncol + _m.hfield_nrow = hfield_nrow + _m.hfield_size = hfield_size + _m.jnt_actfrclimited = jnt_actfrclimited + _m.jnt_actfrcrange = jnt_actfrcrange + _m.jnt_actgravcomp = jnt_actgravcomp + _m.jnt_axis = jnt_axis + _m.jnt_bodyid = jnt_bodyid + _m.jnt_dofadr = jnt_dofadr + _m.jnt_limited_ball_adr = jnt_limited_ball_adr + _m.jnt_limited_slide_hinge_adr = jnt_limited_slide_hinge_adr + _m.jnt_margin = jnt_margin + _m.jnt_pos = jnt_pos + _m.jnt_qposadr = jnt_qposadr + _m.jnt_range = jnt_range + _m.jnt_solimp = jnt_solimp + _m.jnt_solref = jnt_solref + _m.jnt_stiffness = jnt_stiffness + _m.jnt_type = jnt_type + _m.light_bodyid = light_bodyid + _m.light_dir = light_dir + _m.light_dir0 = light_dir0 + _m.light_mode = light_mode + _m.light_pos = light_pos + _m.light_pos0 = light_pos0 + _m.light_poscom0 = light_poscom0 + _m.light_targetbodyid = light_targetbodyid + _m.mapM2M = mapM2M + _m.mat_rgba = mat_rgba + _m.mesh_face = mesh_face + _m.mesh_faceadr = mesh_faceadr + _m.mesh_graph = mesh_graph + _m.mesh_graphadr = mesh_graphadr + _m.mesh_polyadr = mesh_polyadr + _m.mesh_polymap = mesh_polymap + _m.mesh_polymapadr = mesh_polymapadr + _m.mesh_polymapnum = mesh_polymapnum + _m.mesh_polynormal = mesh_polynormal + _m.mesh_polynum = mesh_polynum + _m.mesh_polyvert = mesh_polyvert + _m.mesh_polyvertadr = mesh_polyvertadr + _m.mesh_polyvertnum = mesh_polyvertnum + _m.mesh_vert = mesh_vert + _m.mesh_vertadr = mesh_vertadr + _m.mesh_vertnum = mesh_vertnum + _m.mocap_bodyid = mocap_bodyid + _m.nC = nC + _m.na = na + _m.nbody = nbody + _m.ncam = ncam + _m.neq = neq + _m.nflexedge = nflexedge + _m.nflexelem = nflexelem + _m.nflexvert = nflexvert + _m.ngeom = ngeom + _m.ngravcomp = ngravcomp + _m.nhfield = nhfield + _m.njnt = njnt + _m.nlight = nlight + _m.nlsp = nlsp + _m.nmeshface = nmeshface + _m.nmocap = nmocap + _m.nsite = nsite + _m.ntendon = ntendon + _m.nu = nu + _m.nv = nv + _m.nxn_geom_pair_filtered = nxn_geom_pair_filtered + _m.nxn_pairid = nxn_pairid + _m.nxn_pairid_filtered = nxn_pairid_filtered + _m.opt.broadphase = opt__broadphase + _m.opt.broadphase_filter = opt__broadphase_filter + _m.opt.cone = opt__cone + _m.opt.density = opt__density + _m.opt.disableflags = opt__disableflags + _m.opt.enableflags = opt__enableflags + _m.opt.epa_iterations = opt__epa_iterations + _m.opt.gjk_iterations = opt__gjk_iterations + _m.opt.graph_conditional = opt__graph_conditional + _m.opt.gravity = opt__gravity + _m.opt.has_fluid = opt__has_fluid + _m.opt.impratio = opt__impratio + _m.opt.is_sparse = opt__is_sparse + _m.opt.iterations = opt__iterations + _m.opt.ls_iterations = opt__ls_iterations + _m.opt.ls_parallel = opt__ls_parallel + _m.opt.ls_tolerance = opt__ls_tolerance + _m.opt.magnetic = opt__magnetic + _m.opt.run_collision_detection = opt__run_collision_detection + _m.opt.sdf_initpoints = opt__sdf_initpoints + _m.opt.sdf_iterations = opt__sdf_iterations + _m.opt.solver = opt__solver + _m.opt.timestep = opt__timestep + _m.opt.tolerance = opt__tolerance + _m.opt.viscosity = opt__viscosity + _m.opt.wind = opt__wind + _m.pair_dim = pair_dim + _m.pair_friction = pair_friction + _m.pair_gap = pair_gap + _m.pair_margin = pair_margin + _m.pair_solimp = pair_solimp + _m.pair_solref = pair_solref + _m.pair_solreffriction = pair_solreffriction + _m.plugin = plugin + _m.plugin_attr = plugin_attr + _m.qLD_updates = qLD_updates + _m.qM_fullm_i = qM_fullm_i + _m.qM_fullm_j = qM_fullm_j + _m.qM_madr_ij = qM_madr_ij + _m.qM_mulm_i = qM_mulm_i + _m.qM_mulm_j = qM_mulm_j + _m.qM_tiles = qM_tiles + _m.qpos0 = qpos0 + _m.qpos_spring = qpos_spring + _m.rangefinder_sensor_adr = rangefinder_sensor_adr + _m.sensor_acc_adr = sensor_acc_adr + _m.sensor_adr = sensor_adr + _m.sensor_cutoff = sensor_cutoff + _m.sensor_datatype = sensor_datatype + _m.sensor_e_kinetic = sensor_e_kinetic + _m.sensor_e_potential = sensor_e_potential + _m.sensor_limitfrc_adr = sensor_limitfrc_adr + _m.sensor_limitpos_adr = sensor_limitpos_adr + _m.sensor_limitvel_adr = sensor_limitvel_adr + _m.sensor_objid = sensor_objid + _m.sensor_objtype = sensor_objtype + _m.sensor_pos_adr = sensor_pos_adr + _m.sensor_rangefinder_adr = sensor_rangefinder_adr + _m.sensor_rangefinder_bodyid = sensor_rangefinder_bodyid + _m.sensor_refid = sensor_refid + _m.sensor_reftype = sensor_reftype + _m.sensor_rne_postconstraint = sensor_rne_postconstraint + _m.sensor_subtree_vel = sensor_subtree_vel + _m.sensor_tendonactfrc_adr = sensor_tendonactfrc_adr + _m.sensor_touch_adr = sensor_touch_adr + _m.sensor_type = sensor_type + _m.sensor_vel_adr = sensor_vel_adr + _m.site_bodyid = site_bodyid + _m.site_pos = site_pos + _m.site_quat = site_quat + _m.site_size = site_size + _m.site_type = site_type + _m.stat.meaninertia = stat__meaninertia + _m.subtree_mass = subtree_mass + _m.tendon_actfrclimited = tendon_actfrclimited + _m.tendon_actfrcrange = tendon_actfrcrange + _m.tendon_adr = tendon_adr + _m.tendon_armature = tendon_armature + _m.tendon_damping = tendon_damping + _m.tendon_frictionloss = tendon_frictionloss + _m.tendon_geom_adr = tendon_geom_adr + _m.tendon_invweight0 = tendon_invweight0 + _m.tendon_jnt_adr = tendon_jnt_adr + _m.tendon_length0 = tendon_length0 + _m.tendon_lengthspring = tendon_lengthspring + _m.tendon_limited_adr = tendon_limited_adr + _m.tendon_margin = tendon_margin + _m.tendon_num = tendon_num + _m.tendon_range = tendon_range + _m.tendon_site_pair_adr = tendon_site_pair_adr + _m.tendon_solimp_fri = tendon_solimp_fri + _m.tendon_solimp_lim = tendon_solimp_lim + _m.tendon_solref_fri = tendon_solref_fri + _m.tendon_solref_lim = tendon_solref_lim + _m.tendon_stiffness = tendon_stiffness + _m.wrap_geom_adr = wrap_geom_adr + _m.wrap_jnt_adr = wrap_jnt_adr + _m.wrap_objid = wrap_objid + _m.wrap_prm = wrap_prm + _m.wrap_pulley_scale = wrap_pulley_scale + _m.wrap_site_pair_adr = wrap_site_pair_adr + _m.wrap_type = wrap_type + _d.act = act + _d.act_dot = act_dot + _d.actuator_force = actuator_force + _d.actuator_length = actuator_length + _d.actuator_moment = actuator_moment + _d.actuator_trntype_body_ncon = actuator_trntype_body_ncon + _d.actuator_velocity = actuator_velocity + _d.cacc = cacc + _d.cam_xmat = cam_xmat + _d.cam_xpos = cam_xpos + _d.cdof = cdof + _d.cdof_dot = cdof_dot + _d.cfrc_ext = cfrc_ext + _d.cfrc_int = cfrc_int + _d.cinert = cinert + _d.collision_hftri_index = collision_hftri_index + _d.collision_pair = collision_pair + _d.collision_pairid = collision_pairid + _d.collision_worldid = collision_worldid + _d.contact.dim = contact__dim + _d.contact.dist = contact__dist + _d.contact.efc_address = contact__efc_address + _d.contact.frame = contact__frame + _d.contact.friction = contact__friction + _d.contact.geom = contact__geom + _d.contact.includemargin = contact__includemargin + _d.contact.pos = contact__pos + _d.contact.solimp = contact__solimp + _d.contact.solref = contact__solref + _d.contact.solreffriction = contact__solreffriction + _d.contact.worldid = contact__worldid + _d.crb = crb + _d.ctrl = ctrl + _d.cvel = cvel + _d.efc.D = efc__D + _d.efc.J = efc__J + _d.efc.Jaref = efc__Jaref + _d.efc.Ma = efc__Ma + _d.efc.Mgrad = efc__Mgrad + _d.efc.active = efc__active + _d.efc.alpha = efc__alpha + _d.efc.aref = efc__aref + _d.efc.beta = efc__beta + _d.efc.beta_den = efc__beta_den + _d.efc.beta_num = efc__beta_num + _d.efc.cholesky_L_tmp = efc__cholesky_L_tmp + _d.efc.cholesky_y_tmp = efc__cholesky_y_tmp + _d.efc.condim = efc__condim + _d.efc.cost = efc__cost + _d.efc.cost_candidate = efc__cost_candidate + _d.efc.done = efc__done + _d.efc.force = efc__force + _d.efc.frictionloss = efc__frictionloss + _d.efc.gauss = efc__gauss + _d.efc.grad = efc__grad + _d.efc.grad_dot = efc__grad_dot + _d.efc.gtol = efc__gtol + _d.efc.h = efc__h + _d.efc.hi = efc__hi + _d.efc.hi_alpha = efc__hi_alpha + _d.efc.hi_next = efc__hi_next + _d.efc.hi_next_alpha = efc__hi_next_alpha + _d.efc.id = efc__id + _d.efc.jv = efc__jv + _d.efc.lo = efc__lo + _d.efc.lo_alpha = efc__lo_alpha + _d.efc.lo_next = efc__lo_next + _d.efc.lo_next_alpha = efc__lo_next_alpha + _d.efc.ls_done = efc__ls_done + _d.efc.margin = efc__margin + _d.efc.mid = efc__mid + _d.efc.mid_alpha = efc__mid_alpha + _d.efc.mv = efc__mv + _d.efc.p0 = efc__p0 + _d.efc.pos = efc__pos + _d.efc.prev_Mgrad = efc__prev_Mgrad + _d.efc.prev_cost = efc__prev_cost + _d.efc.prev_grad = efc__prev_grad + _d.efc.quad = efc__quad + _d.efc.quad_gauss = efc__quad_gauss + _d.efc.search = efc__search + _d.efc.search_dot = efc__search_dot + _d.efc.type = efc__type + _d.efc.u = efc__u + _d.efc.uu = efc__uu + _d.efc.uv = efc__uv + _d.efc.vel = efc__vel + _d.efc.vv = efc__vv + _d.energy = energy + _d.epa_face = epa_face + _d.epa_horizon = epa_horizon + _d.epa_index = epa_index + _d.epa_map = epa_map + _d.epa_norm2 = epa_norm2 + _d.epa_pr = epa_pr + _d.epa_vert = epa_vert + _d.epa_vert1 = epa_vert1 + _d.epa_vert2 = epa_vert2 + _d.epa_vert_index1 = epa_vert_index1 + _d.epa_vert_index2 = epa_vert_index2 + _d.eq_active = eq_active + _d.flexedge_length = flexedge_length + _d.flexedge_velocity = flexedge_velocity + _d.flexvert_xpos = flexvert_xpos + _d.fluid_applied = fluid_applied + _d.geom_skip = geom_skip + _d.geom_xmat = geom_xmat + _d.geom_xpos = geom_xpos + _d.light_xdir = light_xdir + _d.light_xpos = light_xpos + _d.mocap_pos = mocap_pos + _d.mocap_quat = mocap_quat + _d.ncollision = ncollision + _d.ncon = ncon + _d.ncon_hfield = ncon_hfield + _d.nconmax = nconmax + _d.ne = ne + _d.ne_connect = ne_connect + _d.ne_jnt = ne_jnt + _d.ne_ten = ne_ten + _d.ne_weld = ne_weld + _d.nefc = nefc + _d.nf = nf + _d.njmax = njmax + _d.nl = nl + _d.nsolving = nsolving + _d.qLD = qLD + _d.qLDiagInv = qLDiagInv + _d.qM = qM + _d.qacc = qacc + _d.qacc_smooth = qacc_smooth + _d.qacc_warmstart = qacc_warmstart + _d.qfrc_actuator = qfrc_actuator + _d.qfrc_applied = qfrc_applied + _d.qfrc_bias = qfrc_bias + _d.qfrc_constraint = qfrc_constraint + _d.qfrc_damper = qfrc_damper + _d.qfrc_fluid = qfrc_fluid + _d.qfrc_gravcomp = qfrc_gravcomp + _d.qfrc_passive = qfrc_passive + _d.qfrc_smooth = qfrc_smooth + _d.qfrc_spring = qfrc_spring + _d.qpos = qpos + _d.qvel = qvel + _d.sap_cumulative_sum = sap_cumulative_sum + _d.sap_projection_lower = sap_projection_lower + _d.sap_projection_upper = sap_projection_upper + _d.sap_range = sap_range + _d.sap_segment_index = sap_segment_index + _d.sap_sort_index = sap_sort_index + _d.sensor_rangefinder_dist = sensor_rangefinder_dist + _d.sensor_rangefinder_geomid = sensor_rangefinder_geomid + _d.sensor_rangefinder_pnt = sensor_rangefinder_pnt + _d.sensor_rangefinder_vec = sensor_rangefinder_vec + _d.sensordata = sensordata + _d.site_xmat = site_xmat + _d.site_xpos = site_xpos + _d.solver_niter = solver_niter + _d.subtree_angmom = subtree_angmom + _d.subtree_bodyvel = subtree_bodyvel + _d.subtree_com = subtree_com + _d.subtree_linvel = subtree_linvel + _d.ten_J = ten_J + _d.ten_Jdot = ten_Jdot + _d.ten_actfrc = ten_actfrc + _d.ten_bias_coef = ten_bias_coef + _d.ten_length = ten_length + _d.ten_velocity = ten_velocity + _d.ten_wrapadr = ten_wrapadr + _d.ten_wrapnum = ten_wrapnum + _d.time = time + _d.wrap_geom_xpos = wrap_geom_xpos + _d.wrap_obj = wrap_obj + _d.wrap_xpos = wrap_xpos + _d.xanchor = xanchor + _d.xaxis = xaxis + _d.xfrc_applied = xfrc_applied + _d.ximat = ximat + _d.xipos = xipos + _d.xmat = xmat + _d.xpos = xpos + _d.xquat = xquat + _d.nworld = nworld + mjwarp.forward(_m, _d) + + +def _forward_jax_impl(m: types.Model, d: types.Data): + output_dims = { + 'act': d.act.shape, + 'act_dot': d.act_dot.shape, + 'actuator_force': d.actuator_force.shape, + 'actuator_length': d._impl.actuator_length.shape, + 'actuator_moment': d._impl.actuator_moment.shape, + 'actuator_trntype_body_ncon': d._impl.actuator_trntype_body_ncon.shape, + 'actuator_velocity': d._impl.actuator_velocity.shape, + 'cacc': d._impl.cacc.shape, + 'cam_xmat': d.cam_xmat.shape, + 'cam_xpos': d.cam_xpos.shape, + 'cdof': d._impl.cdof.shape, + 'cdof_dot': d._impl.cdof_dot.shape, + 'cfrc_ext': d._impl.cfrc_ext.shape, + 'cfrc_int': d._impl.cfrc_int.shape, + 'cinert': d._impl.cinert.shape, + 'collision_hftri_index': d._impl.collision_hftri_index.shape, + 'collision_pair': d._impl.collision_pair.shape, + 'collision_pairid': d._impl.collision_pairid.shape, + 'collision_worldid': d._impl.collision_worldid.shape, + 'crb': d._impl.crb.shape, + 'ctrl': d.ctrl.shape, + 'cvel': d.cvel.shape, + 'energy': d._impl.energy.shape, + 'epa_face': d._impl.epa_face.shape, + 'epa_horizon': d._impl.epa_horizon.shape, + 'epa_index': d._impl.epa_index.shape, + 'epa_map': d._impl.epa_map.shape, + 'epa_norm2': d._impl.epa_norm2.shape, + 'epa_pr': d._impl.epa_pr.shape, + 'epa_vert': d._impl.epa_vert.shape, + 'epa_vert1': d._impl.epa_vert1.shape, + 'epa_vert2': d._impl.epa_vert2.shape, + 'epa_vert_index1': d._impl.epa_vert_index1.shape, + 'epa_vert_index2': d._impl.epa_vert_index2.shape, + 'eq_active': d.eq_active.shape, + 'flexedge_length': d._impl.flexedge_length.shape, + 'flexedge_velocity': d._impl.flexedge_velocity.shape, + 'flexvert_xpos': d._impl.flexvert_xpos.shape, + 'fluid_applied': d._impl.fluid_applied.shape, + 'geom_skip': d._impl.geom_skip.shape, + 'geom_xmat': d.geom_xmat.shape, + 'geom_xpos': d.geom_xpos.shape, + 'light_xdir': d._impl.light_xdir.shape, + 'light_xpos': d._impl.light_xpos.shape, + 'mocap_pos': d.mocap_pos.shape, + 'mocap_quat': d.mocap_quat.shape, + 'ncollision': d._impl.ncollision.shape, + 'ncon': d._impl.ncon.shape, + 'ncon_hfield': d._impl.ncon_hfield.shape, + 'ne': d._impl.ne.shape, + 'ne_connect': d._impl.ne_connect.shape, + 'ne_jnt': d._impl.ne_jnt.shape, + 'ne_ten': d._impl.ne_ten.shape, + 'ne_weld': d._impl.ne_weld.shape, + 'nefc': d._impl.nefc.shape, + 'nf': d._impl.nf.shape, + 'nl': d._impl.nl.shape, + 'nsolving': d._impl.nsolving.shape, + 'qLD': d._impl.qLD.shape, + 'qLDiagInv': d._impl.qLDiagInv.shape, + 'qM': d._impl.qM.shape, + 'qacc': d.qacc.shape, + 'qacc_smooth': d.qacc_smooth.shape, + 'qacc_warmstart': d.qacc_warmstart.shape, + 'qfrc_actuator': d.qfrc_actuator.shape, + 'qfrc_applied': d.qfrc_applied.shape, + 'qfrc_bias': d.qfrc_bias.shape, + 'qfrc_constraint': d.qfrc_constraint.shape, + 'qfrc_damper': d._impl.qfrc_damper.shape, + 'qfrc_fluid': d.qfrc_fluid.shape, + 'qfrc_gravcomp': d.qfrc_gravcomp.shape, + 'qfrc_passive': d.qfrc_passive.shape, + 'qfrc_smooth': d.qfrc_smooth.shape, + 'qfrc_spring': d._impl.qfrc_spring.shape, + 'qpos': d.qpos.shape, + 'qvel': d.qvel.shape, + 'sap_cumulative_sum': d._impl.sap_cumulative_sum.shape, + 'sap_projection_lower': d._impl.sap_projection_lower.shape, + 'sap_projection_upper': d._impl.sap_projection_upper.shape, + 'sap_range': d._impl.sap_range.shape, + 'sap_segment_index': d._impl.sap_segment_index.shape, + 'sap_sort_index': d._impl.sap_sort_index.shape, + 'sensor_rangefinder_dist': d._impl.sensor_rangefinder_dist.shape, + 'sensor_rangefinder_geomid': d._impl.sensor_rangefinder_geomid.shape, + 'sensor_rangefinder_pnt': d._impl.sensor_rangefinder_pnt.shape, + 'sensor_rangefinder_vec': d._impl.sensor_rangefinder_vec.shape, + 'sensordata': d.sensordata.shape, + 'site_xmat': d.site_xmat.shape, + 'site_xpos': d.site_xpos.shape, + 'solver_niter': d._impl.solver_niter.shape, + 'subtree_angmom': d._impl.subtree_angmom.shape, + 'subtree_bodyvel': d._impl.subtree_bodyvel.shape, + 'subtree_com': d.subtree_com.shape, + 'subtree_linvel': d._impl.subtree_linvel.shape, + 'ten_J': d._impl.ten_J.shape, + 'ten_Jdot': d._impl.ten_Jdot.shape, + 'ten_actfrc': d._impl.ten_actfrc.shape, + 'ten_bias_coef': d._impl.ten_bias_coef.shape, + 'ten_length': d._impl.ten_length.shape, + 'ten_velocity': d._impl.ten_velocity.shape, + 'ten_wrapadr': d._impl.ten_wrapadr.shape, + 'ten_wrapnum': d._impl.ten_wrapnum.shape, + 'time': d.time.shape, + 'wrap_geom_xpos': d._impl.wrap_geom_xpos.shape, + 'wrap_obj': d._impl.wrap_obj.shape, + 'wrap_xpos': d._impl.wrap_xpos.shape, + 'xanchor': d.xanchor.shape, + 'xaxis': d.xaxis.shape, + 'xfrc_applied': d.xfrc_applied.shape, + 'ximat': d.ximat.shape, + 'xipos': d.xipos.shape, + 'xmat': d.xmat.shape, + 'xpos': d.xpos.shape, + 'xquat': d.xquat.shape, + 'contact__dim': d._impl.contact__dim.shape, + 'contact__dist': d._impl.contact__dist.shape, + 'contact__efc_address': d._impl.contact__efc_address.shape, + 'contact__frame': d._impl.contact__frame.shape, + 'contact__friction': d._impl.contact__friction.shape, + 'contact__geom': d._impl.contact__geom.shape, + 'contact__includemargin': d._impl.contact__includemargin.shape, + 'contact__pos': d._impl.contact__pos.shape, + 'contact__solimp': d._impl.contact__solimp.shape, + 'contact__solref': d._impl.contact__solref.shape, + 'contact__solreffriction': d._impl.contact__solreffriction.shape, + 'contact__worldid': d._impl.contact__worldid.shape, + 'efc__D': d._impl.efc__D.shape, + 'efc__J': d._impl.efc__J.shape, + 'efc__Jaref': d._impl.efc__Jaref.shape, + 'efc__Ma': d._impl.efc__Ma.shape, + 'efc__Mgrad': d._impl.efc__Mgrad.shape, + 'efc__active': d._impl.efc__active.shape, + 'efc__alpha': d._impl.efc__alpha.shape, + 'efc__aref': d._impl.efc__aref.shape, + 'efc__beta': d._impl.efc__beta.shape, + 'efc__beta_den': d._impl.efc__beta_den.shape, + 'efc__beta_num': d._impl.efc__beta_num.shape, + 'efc__cholesky_L_tmp': d._impl.efc__cholesky_L_tmp.shape, + 'efc__cholesky_y_tmp': d._impl.efc__cholesky_y_tmp.shape, + 'efc__condim': d._impl.efc__condim.shape, + 'efc__cost': d._impl.efc__cost.shape, + 'efc__cost_candidate': d._impl.efc__cost_candidate.shape, + 'efc__done': d._impl.efc__done.shape, + 'efc__force': d._impl.efc__force.shape, + 'efc__frictionloss': d._impl.efc__frictionloss.shape, + 'efc__gauss': d._impl.efc__gauss.shape, + 'efc__grad': d._impl.efc__grad.shape, + 'efc__grad_dot': d._impl.efc__grad_dot.shape, + 'efc__gtol': d._impl.efc__gtol.shape, + 'efc__h': d._impl.efc__h.shape, + 'efc__hi': d._impl.efc__hi.shape, + 'efc__hi_alpha': d._impl.efc__hi_alpha.shape, + 'efc__hi_next': d._impl.efc__hi_next.shape, + 'efc__hi_next_alpha': d._impl.efc__hi_next_alpha.shape, + 'efc__id': d._impl.efc__id.shape, + 'efc__jv': d._impl.efc__jv.shape, + 'efc__lo': d._impl.efc__lo.shape, + 'efc__lo_alpha': d._impl.efc__lo_alpha.shape, + 'efc__lo_next': d._impl.efc__lo_next.shape, + 'efc__lo_next_alpha': d._impl.efc__lo_next_alpha.shape, + 'efc__ls_done': d._impl.efc__ls_done.shape, + 'efc__margin': d._impl.efc__margin.shape, + 'efc__mid': d._impl.efc__mid.shape, + 'efc__mid_alpha': d._impl.efc__mid_alpha.shape, + 'efc__mv': d._impl.efc__mv.shape, + 'efc__p0': d._impl.efc__p0.shape, + 'efc__pos': d._impl.efc__pos.shape, + 'efc__prev_Mgrad': d._impl.efc__prev_Mgrad.shape, + 'efc__prev_cost': d._impl.efc__prev_cost.shape, + 'efc__prev_grad': d._impl.efc__prev_grad.shape, + 'efc__quad': d._impl.efc__quad.shape, + 'efc__quad_gauss': d._impl.efc__quad_gauss.shape, + 'efc__search': d._impl.efc__search.shape, + 'efc__search_dot': d._impl.efc__search_dot.shape, + 'efc__type': d._impl.efc__type.shape, + 'efc__u': d._impl.efc__u.shape, + 'efc__uu': d._impl.efc__uu.shape, + 'efc__uv': d._impl.efc__uv.shape, + 'efc__vel': d._impl.efc__vel.shape, + 'efc__vv': d._impl.efc__vv.shape, + } + jf = ffi.jax_callable_variadic_tuple( + _forward_shim, + num_outputs=180, + output_dims=output_dims, + vmap_method=None, + graph_compatible=True, + in_out_argnames={ + 'act', + 'act_dot', + 'actuator_force', + 'actuator_length', + 'actuator_moment', + 'actuator_trntype_body_ncon', + 'actuator_velocity', + 'cacc', + 'cam_xmat', + 'cam_xpos', + 'cdof', + 'cdof_dot', + 'cfrc_ext', + 'cfrc_int', + 'cinert', + 'collision_hftri_index', + 'collision_pair', + 'collision_pairid', + 'collision_worldid', + 'crb', + 'ctrl', + 'cvel', + 'energy', + 'epa_face', + 'epa_horizon', + 'epa_index', + 'epa_map', + 'epa_norm2', + 'epa_pr', + 'epa_vert', + 'epa_vert1', + 'epa_vert2', + 'epa_vert_index1', + 'epa_vert_index2', + 'eq_active', + 'flexedge_length', + 'flexedge_velocity', + 'flexvert_xpos', + 'fluid_applied', + 'geom_skip', + 'geom_xmat', + 'geom_xpos', + 'light_xdir', + 'light_xpos', + 'mocap_pos', + 'mocap_quat', + 'ncollision', + 'ncon', + 'ncon_hfield', + 'ne', + 'ne_connect', + 'ne_jnt', + 'ne_ten', + 'ne_weld', + 'nefc', + 'nf', + 'nl', + 'nsolving', + 'qLD', + 'qLDiagInv', + 'qM', + 'qacc', + 'qacc_smooth', + 'qacc_warmstart', + 'qfrc_actuator', + 'qfrc_applied', + 'qfrc_bias', + 'qfrc_constraint', + 'qfrc_damper', + 'qfrc_fluid', + 'qfrc_gravcomp', + 'qfrc_passive', + 'qfrc_smooth', + 'qfrc_spring', + 'qpos', + 'qvel', + 'sap_cumulative_sum', + 'sap_projection_lower', + 'sap_projection_upper', + 'sap_range', + 'sap_segment_index', + 'sap_sort_index', + 'sensor_rangefinder_dist', + 'sensor_rangefinder_geomid', + 'sensor_rangefinder_pnt', + 'sensor_rangefinder_vec', + 'sensordata', + 'site_xmat', + 'site_xpos', + 'solver_niter', + 'subtree_angmom', + 'subtree_bodyvel', + 'subtree_com', + 'subtree_linvel', + 'ten_J', + 'ten_Jdot', + 'ten_actfrc', + 'ten_bias_coef', + 'ten_length', + 'ten_velocity', + 'ten_wrapadr', + 'ten_wrapnum', + 'time', + 'wrap_geom_xpos', + 'wrap_obj', + 'wrap_xpos', + 'xanchor', + 'xaxis', + 'xfrc_applied', + 'ximat', + 'xipos', + 'xmat', + 'xpos', + 'xquat', + 'contact__dim', + 'contact__dist', + 'contact__efc_address', + 'contact__frame', + 'contact__friction', + 'contact__geom', + 'contact__includemargin', + 'contact__pos', + 'contact__solimp', + 'contact__solref', + 'contact__solreffriction', + 'contact__worldid', + 'efc__D', + 'efc__J', + 'efc__Jaref', + 'efc__Ma', + 'efc__Mgrad', + 'efc__active', + 'efc__alpha', + 'efc__aref', + 'efc__beta', + 'efc__beta_den', + 'efc__beta_num', + 'efc__cholesky_L_tmp', + 'efc__cholesky_y_tmp', + 'efc__condim', + 'efc__cost', + 'efc__cost_candidate', + 'efc__done', + 'efc__force', + 'efc__frictionloss', + 'efc__gauss', + 'efc__grad', + 'efc__grad_dot', + 'efc__gtol', + 'efc__h', + 'efc__hi', + 'efc__hi_alpha', + 'efc__hi_next', + 'efc__hi_next_alpha', + 'efc__id', + 'efc__jv', + 'efc__lo', + 'efc__lo_alpha', + 'efc__lo_next', + 'efc__lo_next_alpha', + 'efc__ls_done', + 'efc__margin', + 'efc__mid', + 'efc__mid_alpha', + 'efc__mv', + 'efc__p0', + 'efc__pos', + 'efc__prev_Mgrad', + 'efc__prev_cost', + 'efc__prev_grad', + 'efc__quad', + 'efc__quad_gauss', + 'efc__search', + 'efc__search_dot', + 'efc__type', + 'efc__u', + 'efc__uu', + 'efc__uv', + 'efc__vel', + 'efc__vv', + }, + ) + out = jf( + d.qpos.shape[0], + m._impl.M_rowadr, + m._impl.M_rownnz, + m.actuator_acc0, + m.actuator_actadr, + m.actuator_actearly, + m.actuator_actlimited, + m.actuator_actnum, + m.actuator_actrange, + m.actuator_biasprm, + m.actuator_biastype, + m.actuator_cranklength, + m.actuator_ctrllimited, + m.actuator_ctrlrange, + m.actuator_dynprm, + m.actuator_dyntype, + m.actuator_forcelimited, + m.actuator_forcerange, + m.actuator_gainprm, + m.actuator_gaintype, + m.actuator_gear, + m.actuator_lengthrange, + m.actuator_trnid, + m.actuator_trntype, + m._impl.actuator_trntype_body_adr, + m._impl.block_dim, + m.body_dofadr, + m.body_dofnum, + m.body_gravcomp, + m.body_inertia, + m.body_invweight0, + m.body_ipos, + m.body_iquat, + m.body_jntadr, + m.body_jntnum, + m.body_mass, + m.body_parentid, + m.body_pos, + m.body_quat, + m.body_rootid, + m.body_subtreemass, + m._impl.body_tree, + m.body_weldid, + m.cam_bodyid, + m.cam_fovy, + m.cam_intrinsic, + m.cam_mat0, + m.cam_mode, + m.cam_pos, + m.cam_pos0, + m.cam_poscom0, + m.cam_quat, + m.cam_resolution, + m.cam_sensorsize, + m.cam_targetbodyid, + m._impl.condim_max, + m.dof_Madr, + m.dof_armature, + m.dof_bodyid, + m.dof_damping, + m.dof_frictionloss, + m.dof_invweight0, + m.dof_jntid, + m.dof_parentid, + m.dof_solimp, + m.dof_solref, + m._impl.dof_tri_col, + m._impl.dof_tri_row, + m._impl.eq_connect_adr, + m.eq_data, + m._impl.eq_jnt_adr, + m.eq_obj1id, + m.eq_obj2id, + m.eq_objtype, + m.eq_solimp, + m.eq_solref, + m._impl.eq_ten_adr, + m._impl.eq_wld_adr, + m._impl.flex_bending, + m._impl.flex_damping, + m._impl.flex_dim, + m._impl.flex_edge, + m._impl.flex_edgeadr, + m._impl.flex_edgeflap, + m._impl.flex_elem, + m._impl.flex_elemedge, + m._impl.flex_elemedgeadr, + m._impl.flex_stiffness, + m._impl.flex_vertadr, + m._impl.flex_vertbodyid, + m._impl.flexedge_length0, + m.geom_aabb, + m.geom_bodyid, + m.geom_condim, + m.geom_dataid, + m.geom_friction, + m.geom_gap, + m.geom_group, + m.geom_margin, + m.geom_matid, + m._impl.geom_pair_type_count, + m._impl.geom_plugin_index, + m.geom_pos, + m.geom_priority, + m.geom_quat, + m.geom_rbound, + m.geom_rgba, + m.geom_size, + m.geom_solimp, + m.geom_solmix, + m.geom_solref, + m.geom_type, + m._impl.geompair2hfgeompair, + m._impl.has_sdf_geom, + m.hfield_adr, + m.hfield_data, + m.hfield_ncol, + m.hfield_nrow, + m.hfield_size, + m.jnt_actfrclimited, + m.jnt_actfrcrange, + m.jnt_actgravcomp, + m.jnt_axis, + m.jnt_bodyid, + m.jnt_dofadr, + m._impl.jnt_limited_ball_adr, + m._impl.jnt_limited_slide_hinge_adr, + m.jnt_margin, + m.jnt_pos, + m.jnt_qposadr, + m.jnt_range, + m.jnt_solimp, + m.jnt_solref, + m.jnt_stiffness, + m.jnt_type, + m._impl.light_bodyid, + m.light_dir, + m.light_dir0, + m.light_mode, + m.light_pos, + m.light_pos0, + m.light_poscom0, + m._impl.light_targetbodyid, + m._impl.mapM2M, + m.mat_rgba, + m.mesh_face, + m.mesh_faceadr, + m.mesh_graph, + m.mesh_graphadr, + m._impl.mesh_polyadr, + m._impl.mesh_polymap, + m._impl.mesh_polymapadr, + m._impl.mesh_polymapnum, + m._impl.mesh_polynormal, + m._impl.mesh_polynum, + m._impl.mesh_polyvert, + m._impl.mesh_polyvertadr, + m._impl.mesh_polyvertnum, + m.mesh_vert, + m.mesh_vertadr, + m.mesh_vertnum, + m._impl.mocap_bodyid, + m.nC, + m.na, + m.nbody, + m.ncam, + m.neq, + m._impl.nflexedge, + m._impl.nflexelem, + m._impl.nflexvert, + m.ngeom, + m.ngravcomp, + m.nhfield, + m.njnt, + m.nlight, + m._impl.nlsp, + m.nmeshface, + m.nmocap, + m.nsite, + m.ntendon, + m.nu, + m.nv, + m._impl.nxn_geom_pair_filtered, + m._impl.nxn_pairid, + m._impl.nxn_pairid_filtered, + m.pair_dim, + m.pair_friction, + m.pair_gap, + m.pair_margin, + m.pair_solimp, + m.pair_solref, + m.pair_solreffriction, + m._impl.plugin, + m._impl.plugin_attr, + m._impl.qLD_updates, + m._impl.qM_fullm_i, + m._impl.qM_fullm_j, + m._impl.qM_madr_ij, + m._impl.qM_mulm_i, + m._impl.qM_mulm_j, + m._impl.qM_tiles, + m.qpos0, + m.qpos_spring, + m._impl.rangefinder_sensor_adr, + m._impl.sensor_acc_adr, + m.sensor_adr, + m.sensor_cutoff, + m.sensor_datatype, + m._impl.sensor_e_kinetic, + m._impl.sensor_e_potential, + m._impl.sensor_limitfrc_adr, + m._impl.sensor_limitpos_adr, + m._impl.sensor_limitvel_adr, + m.sensor_objid, + m.sensor_objtype, + m._impl.sensor_pos_adr, + m._impl.sensor_rangefinder_adr, + m._impl.sensor_rangefinder_bodyid, + m.sensor_refid, + m.sensor_reftype, + m._impl.sensor_rne_postconstraint, + m._impl.sensor_subtree_vel, + m._impl.sensor_tendonactfrc_adr, + m._impl.sensor_touch_adr, + m.sensor_type, + m._impl.sensor_vel_adr, + m.site_bodyid, + m.site_pos, + m.site_quat, + m.site_size, + m.site_type, + m._impl.subtree_mass, + m.tendon_actfrclimited, + m.tendon_actfrcrange, + m.tendon_adr, + m.tendon_armature, + m.tendon_damping, + m.tendon_frictionloss, + m._impl.tendon_geom_adr, + m.tendon_invweight0, + m._impl.tendon_jnt_adr, + m.tendon_length0, + m.tendon_lengthspring, + m._impl.tendon_limited_adr, + m.tendon_margin, + m.tendon_num, + m.tendon_range, + m._impl.tendon_site_pair_adr, + m.tendon_solimp_fri, + m.tendon_solimp_lim, + m.tendon_solref_fri, + m.tendon_solref_lim, + m.tendon_stiffness, + m._impl.wrap_geom_adr, + m._impl.wrap_jnt_adr, + m.wrap_objid, + m.wrap_prm, + m._impl.wrap_pulley_scale, + m._impl.wrap_site_pair_adr, + m.wrap_type, + m.opt._impl.broadphase, + m.opt._impl.broadphase_filter, + m.opt.cone, + m.opt.density, + m.opt.disableflags, + m.opt.enableflags, + m.opt._impl.epa_iterations, + m.opt._impl.gjk_iterations, + m.opt._impl.graph_conditional, + m.opt.gravity, + m.opt._impl.has_fluid, + m.opt.impratio, + m.opt._impl.is_sparse, + m.opt.iterations, + m.opt.ls_iterations, + m.opt._impl.ls_parallel, + m.opt.ls_tolerance, + m.opt.magnetic, + m.opt._impl.run_collision_detection, + m.opt._impl.sdf_initpoints, + m.opt._impl.sdf_iterations, + m.opt.solver, + m.opt.timestep, + m.opt.tolerance, + m.opt.viscosity, + m.opt.wind, + m.stat.meaninertia, + d._impl.nconmax, + d._impl.njmax, + d.act, + d.act_dot, + d.actuator_force, + d._impl.actuator_length, + d._impl.actuator_moment, + d._impl.actuator_trntype_body_ncon, + d._impl.actuator_velocity, + d._impl.cacc, + d.cam_xmat, + d.cam_xpos, + d._impl.cdof, + d._impl.cdof_dot, + d._impl.cfrc_ext, + d._impl.cfrc_int, + d._impl.cinert, + d._impl.collision_hftri_index, + d._impl.collision_pair, + d._impl.collision_pairid, + d._impl.collision_worldid, + d._impl.crb, + d.ctrl, + d.cvel, + d._impl.energy, + d._impl.epa_face, + d._impl.epa_horizon, + d._impl.epa_index, + d._impl.epa_map, + d._impl.epa_norm2, + d._impl.epa_pr, + d._impl.epa_vert, + d._impl.epa_vert1, + d._impl.epa_vert2, + d._impl.epa_vert_index1, + d._impl.epa_vert_index2, + d.eq_active, + d._impl.flexedge_length, + d._impl.flexedge_velocity, + d._impl.flexvert_xpos, + d._impl.fluid_applied, + d._impl.geom_skip, + d.geom_xmat, + d.geom_xpos, + d._impl.light_xdir, + d._impl.light_xpos, + d.mocap_pos, + d.mocap_quat, + d._impl.ncollision, + d._impl.ncon, + d._impl.ncon_hfield, + d._impl.ne, + d._impl.ne_connect, + d._impl.ne_jnt, + d._impl.ne_ten, + d._impl.ne_weld, + d._impl.nefc, + d._impl.nf, + d._impl.nl, + d._impl.nsolving, + d._impl.qLD, + d._impl.qLDiagInv, + d._impl.qM, + d.qacc, + d.qacc_smooth, + d.qacc_warmstart, + d.qfrc_actuator, + d.qfrc_applied, + d.qfrc_bias, + d.qfrc_constraint, + d._impl.qfrc_damper, + d.qfrc_fluid, + d.qfrc_gravcomp, + d.qfrc_passive, + d.qfrc_smooth, + d._impl.qfrc_spring, + d.qpos, + d.qvel, + d._impl.sap_cumulative_sum, + d._impl.sap_projection_lower, + d._impl.sap_projection_upper, + d._impl.sap_range, + d._impl.sap_segment_index, + d._impl.sap_sort_index, + d._impl.sensor_rangefinder_dist, + d._impl.sensor_rangefinder_geomid, + d._impl.sensor_rangefinder_pnt, + d._impl.sensor_rangefinder_vec, + d.sensordata, + d.site_xmat, + d.site_xpos, + d._impl.solver_niter, + d._impl.subtree_angmom, + d._impl.subtree_bodyvel, + d.subtree_com, + d._impl.subtree_linvel, + d._impl.ten_J, + d._impl.ten_Jdot, + d._impl.ten_actfrc, + d._impl.ten_bias_coef, + d._impl.ten_length, + d._impl.ten_velocity, + d._impl.ten_wrapadr, + d._impl.ten_wrapnum, + d.time, + d._impl.wrap_geom_xpos, + d._impl.wrap_obj, + d._impl.wrap_xpos, + d.xanchor, + d.xaxis, + d.xfrc_applied, + d.ximat, + d.xipos, + d.xmat, + d.xpos, + d.xquat, + d._impl.contact__dim, + d._impl.contact__dist, + d._impl.contact__efc_address, + d._impl.contact__frame, + d._impl.contact__friction, + d._impl.contact__geom, + d._impl.contact__includemargin, + d._impl.contact__pos, + d._impl.contact__solimp, + d._impl.contact__solref, + d._impl.contact__solreffriction, + d._impl.contact__worldid, + d._impl.efc__D, + d._impl.efc__J, + d._impl.efc__Jaref, + d._impl.efc__Ma, + d._impl.efc__Mgrad, + d._impl.efc__active, + d._impl.efc__alpha, + d._impl.efc__aref, + d._impl.efc__beta, + d._impl.efc__beta_den, + d._impl.efc__beta_num, + d._impl.efc__cholesky_L_tmp, + d._impl.efc__cholesky_y_tmp, + d._impl.efc__condim, + d._impl.efc__cost, + d._impl.efc__cost_candidate, + d._impl.efc__done, + d._impl.efc__force, + d._impl.efc__frictionloss, + d._impl.efc__gauss, + d._impl.efc__grad, + d._impl.efc__grad_dot, + d._impl.efc__gtol, + d._impl.efc__h, + d._impl.efc__hi, + d._impl.efc__hi_alpha, + d._impl.efc__hi_next, + d._impl.efc__hi_next_alpha, + d._impl.efc__id, + d._impl.efc__jv, + d._impl.efc__lo, + d._impl.efc__lo_alpha, + d._impl.efc__lo_next, + d._impl.efc__lo_next_alpha, + d._impl.efc__ls_done, + d._impl.efc__margin, + d._impl.efc__mid, + d._impl.efc__mid_alpha, + d._impl.efc__mv, + d._impl.efc__p0, + d._impl.efc__pos, + d._impl.efc__prev_Mgrad, + d._impl.efc__prev_cost, + d._impl.efc__prev_grad, + d._impl.efc__quad, + d._impl.efc__quad_gauss, + d._impl.efc__search, + d._impl.efc__search_dot, + d._impl.efc__type, + d._impl.efc__u, + d._impl.efc__uu, + d._impl.efc__uv, + d._impl.efc__vel, + d._impl.efc__vv, + ) + d = d.tree_replace({ + 'act': out[0], + 'act_dot': out[1], + 'actuator_force': out[2], + '_impl.actuator_length': out[3], + '_impl.actuator_moment': out[4], + '_impl.actuator_trntype_body_ncon': out[5], + '_impl.actuator_velocity': out[6], + '_impl.cacc': out[7], + 'cam_xmat': out[8], + 'cam_xpos': out[9], + '_impl.cdof': out[10], + '_impl.cdof_dot': out[11], + '_impl.cfrc_ext': out[12], + '_impl.cfrc_int': out[13], + '_impl.cinert': out[14], + '_impl.collision_hftri_index': out[15], + '_impl.collision_pair': out[16], + '_impl.collision_pairid': out[17], + '_impl.collision_worldid': out[18], + '_impl.crb': out[19], + 'ctrl': out[20], + 'cvel': out[21], + '_impl.energy': out[22], + '_impl.epa_face': out[23], + '_impl.epa_horizon': out[24], + '_impl.epa_index': out[25], + '_impl.epa_map': out[26], + '_impl.epa_norm2': out[27], + '_impl.epa_pr': out[28], + '_impl.epa_vert': out[29], + '_impl.epa_vert1': out[30], + '_impl.epa_vert2': out[31], + '_impl.epa_vert_index1': out[32], + '_impl.epa_vert_index2': out[33], + 'eq_active': out[34], + '_impl.flexedge_length': out[35], + '_impl.flexedge_velocity': out[36], + '_impl.flexvert_xpos': out[37], + '_impl.fluid_applied': out[38], + '_impl.geom_skip': out[39], + 'geom_xmat': out[40], + 'geom_xpos': out[41], + '_impl.light_xdir': out[42], + '_impl.light_xpos': out[43], + 'mocap_pos': out[44], + 'mocap_quat': out[45], + '_impl.ncollision': out[46], + '_impl.ncon': out[47], + '_impl.ncon_hfield': out[48], + '_impl.ne': out[49], + '_impl.ne_connect': out[50], + '_impl.ne_jnt': out[51], + '_impl.ne_ten': out[52], + '_impl.ne_weld': out[53], + '_impl.nefc': out[54], + '_impl.nf': out[55], + '_impl.nl': out[56], + '_impl.nsolving': out[57], + '_impl.qLD': out[58], + '_impl.qLDiagInv': out[59], + '_impl.qM': out[60], + 'qacc': out[61], + 'qacc_smooth': out[62], + 'qacc_warmstart': out[63], + 'qfrc_actuator': out[64], + 'qfrc_applied': out[65], + 'qfrc_bias': out[66], + 'qfrc_constraint': out[67], + '_impl.qfrc_damper': out[68], + 'qfrc_fluid': out[69], + 'qfrc_gravcomp': out[70], + 'qfrc_passive': out[71], + 'qfrc_smooth': out[72], + '_impl.qfrc_spring': out[73], + 'qpos': out[74], + 'qvel': out[75], + '_impl.sap_cumulative_sum': out[76], + '_impl.sap_projection_lower': out[77], + '_impl.sap_projection_upper': out[78], + '_impl.sap_range': out[79], + '_impl.sap_segment_index': out[80], + '_impl.sap_sort_index': out[81], + '_impl.sensor_rangefinder_dist': out[82], + '_impl.sensor_rangefinder_geomid': out[83], + '_impl.sensor_rangefinder_pnt': out[84], + '_impl.sensor_rangefinder_vec': out[85], + 'sensordata': out[86], + 'site_xmat': out[87], + 'site_xpos': out[88], + '_impl.solver_niter': out[89], + '_impl.subtree_angmom': out[90], + '_impl.subtree_bodyvel': out[91], + 'subtree_com': out[92], + '_impl.subtree_linvel': out[93], + '_impl.ten_J': out[94], + '_impl.ten_Jdot': out[95], + '_impl.ten_actfrc': out[96], + '_impl.ten_bias_coef': out[97], + '_impl.ten_length': out[98], + '_impl.ten_velocity': out[99], + '_impl.ten_wrapadr': out[100], + '_impl.ten_wrapnum': out[101], + 'time': out[102], + '_impl.wrap_geom_xpos': out[103], + '_impl.wrap_obj': out[104], + '_impl.wrap_xpos': out[105], + 'xanchor': out[106], + 'xaxis': out[107], + 'xfrc_applied': out[108], + 'ximat': out[109], + 'xipos': out[110], + 'xmat': out[111], + 'xpos': out[112], + 'xquat': out[113], + '_impl.contact__dim': out[114], + '_impl.contact__dist': out[115], + '_impl.contact__efc_address': out[116], + '_impl.contact__frame': out[117], + '_impl.contact__friction': out[118], + '_impl.contact__geom': out[119], + '_impl.contact__includemargin': out[120], + '_impl.contact__pos': out[121], + '_impl.contact__solimp': out[122], + '_impl.contact__solref': out[123], + '_impl.contact__solreffriction': out[124], + '_impl.contact__worldid': out[125], + '_impl.efc__D': out[126], + '_impl.efc__J': out[127], + '_impl.efc__Jaref': out[128], + '_impl.efc__Ma': out[129], + '_impl.efc__Mgrad': out[130], + '_impl.efc__active': out[131], + '_impl.efc__alpha': out[132], + '_impl.efc__aref': out[133], + '_impl.efc__beta': out[134], + '_impl.efc__beta_den': out[135], + '_impl.efc__beta_num': out[136], + '_impl.efc__cholesky_L_tmp': out[137], + '_impl.efc__cholesky_y_tmp': out[138], + '_impl.efc__condim': out[139], + '_impl.efc__cost': out[140], + '_impl.efc__cost_candidate': out[141], + '_impl.efc__done': out[142], + '_impl.efc__force': out[143], + '_impl.efc__frictionloss': out[144], + '_impl.efc__gauss': out[145], + '_impl.efc__grad': out[146], + '_impl.efc__grad_dot': out[147], + '_impl.efc__gtol': out[148], + '_impl.efc__h': out[149], + '_impl.efc__hi': out[150], + '_impl.efc__hi_alpha': out[151], + '_impl.efc__hi_next': out[152], + '_impl.efc__hi_next_alpha': out[153], + '_impl.efc__id': out[154], + '_impl.efc__jv': out[155], + '_impl.efc__lo': out[156], + '_impl.efc__lo_alpha': out[157], + '_impl.efc__lo_next': out[158], + '_impl.efc__lo_next_alpha': out[159], + '_impl.efc__ls_done': out[160], + '_impl.efc__margin': out[161], + '_impl.efc__mid': out[162], + '_impl.efc__mid_alpha': out[163], + '_impl.efc__mv': out[164], + '_impl.efc__p0': out[165], + '_impl.efc__pos': out[166], + '_impl.efc__prev_Mgrad': out[167], + '_impl.efc__prev_cost': out[168], + '_impl.efc__prev_grad': out[169], + '_impl.efc__quad': out[170], + '_impl.efc__quad_gauss': out[171], + '_impl.efc__search': out[172], + '_impl.efc__search_dot': out[173], + '_impl.efc__type': out[174], + '_impl.efc__u': out[175], + '_impl.efc__uu': out[176], + '_impl.efc__uv': out[177], + '_impl.efc__vel': out[178], + '_impl.efc__vv': out[179], + }) + return d + + +@jax.custom_batching.custom_vmap +@ffi.marshal_jax_warp_callable +def forward(m: types.Model, d: types.Data): + return _forward_jax_impl(m, d) + + +@forward.def_vmap +@ffi.marshal_custom_vmap +def forward_vmap(unused_axis_size, is_batched, m, d): + d = forward(m, d) + return d, is_batched[1] + + +_m = mjwarp.Model( + **{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init} +) +_d = mjwarp.Data( + **{f.name: None for f in dataclasses.fields(mjwarp.Data) if f.init} +) +_o = mjwarp.Option( + **{f.name: None for f in dataclasses.fields(mjwarp.Option) if f.init} +) +_s = mjwarp.Statistic( + **{f.name: None for f in dataclasses.fields(mjwarp.Statistic) if f.init} +) +_c = mjwarp.Contact( + **{f.name: None for f in dataclasses.fields(mjwarp.Contact) if f.init} +) +_e = mjwarp.Constraint( + **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} +) + + +@ffi.format_args_for_warp +def _step_shim( + # Model + nworld: int, + M_rowadr: wp.array(dtype=int), + M_rownnz: wp.array(dtype=int), + actuator_acc0: wp.array(dtype=float), + actuator_actadr: wp.array(dtype=int), + actuator_actearly: wp.array(dtype=bool), + actuator_actlimited: wp.array(dtype=bool), + actuator_actnum: wp.array(dtype=int), + actuator_actrange: wp.array2d(dtype=wp.vec2), + actuator_affine_bias_gain: bool, + actuator_biasprm: wp.array2d(dtype=mjwp_types.vec10f), + actuator_biastype: wp.array(dtype=int), + actuator_cranklength: wp.array(dtype=float), + actuator_ctrllimited: wp.array(dtype=bool), + actuator_ctrlrange: wp.array2d(dtype=wp.vec2), + actuator_dynprm: wp.array2d(dtype=mjwp_types.vec10f), + actuator_dyntype: wp.array(dtype=int), + actuator_forcelimited: wp.array(dtype=bool), + actuator_forcerange: wp.array2d(dtype=wp.vec2), + actuator_gainprm: wp.array2d(dtype=mjwp_types.vec10f), + actuator_gaintype: wp.array(dtype=int), + actuator_gear: wp.array2d(dtype=wp.spatial_vector), + actuator_lengthrange: wp.array(dtype=wp.vec2), + actuator_trnid: wp.array(dtype=wp.vec2i), + actuator_trntype: wp.array(dtype=int), + actuator_trntype_body_adr: wp.array(dtype=int), + block_dim: mjwp_types.BlockDim, + body_dofadr: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_gravcomp: wp.array2d(dtype=float), + body_inertia: wp.array2d(dtype=wp.vec3), + body_invweight0: wp.array2d(dtype=wp.vec2), + body_ipos: wp.array2d(dtype=wp.vec3), + body_iquat: wp.array2d(dtype=wp.quat), + body_jntadr: wp.array(dtype=int), + body_jntnum: wp.array(dtype=int), + body_mass: wp.array2d(dtype=float), + body_parentid: wp.array(dtype=int), + body_pos: wp.array2d(dtype=wp.vec3), + body_quat: wp.array2d(dtype=wp.quat), + body_rootid: wp.array(dtype=int), + body_subtreemass: wp.array2d(dtype=float), + body_tree: tuple[wp.array(dtype=int), ...], + body_weldid: wp.array(dtype=int), + cam_bodyid: wp.array(dtype=int), + cam_fovy: wp.array(dtype=float), + cam_intrinsic: wp.array(dtype=wp.vec4), + cam_mat0: wp.array2d(dtype=wp.mat33), + cam_mode: wp.array(dtype=int), + cam_pos: wp.array2d(dtype=wp.vec3), + cam_pos0: wp.array2d(dtype=wp.vec3), + cam_poscom0: wp.array2d(dtype=wp.vec3), + cam_quat: wp.array2d(dtype=wp.quat), + cam_resolution: wp.array(dtype=wp.vec2i), + cam_sensorsize: wp.array(dtype=wp.vec2), + cam_targetbodyid: wp.array(dtype=int), + condim_max: int, + dof_Madr: wp.array(dtype=int), + dof_armature: wp.array2d(dtype=float), + dof_bodyid: wp.array(dtype=int), + dof_damping: wp.array2d(dtype=float), + dof_frictionloss: wp.array2d(dtype=float), + dof_invweight0: wp.array2d(dtype=float), + dof_jntid: wp.array(dtype=int), + dof_parentid: wp.array(dtype=int), + dof_solimp: wp.array2d(dtype=mjwp_types.vec5), + dof_solref: wp.array2d(dtype=wp.vec2), + dof_tri_col: wp.array(dtype=int), + dof_tri_row: wp.array(dtype=int), + eq_connect_adr: wp.array(dtype=int), + eq_data: wp.array2d(dtype=mjwp_types.vec11), + eq_jnt_adr: wp.array(dtype=int), + eq_obj1id: wp.array(dtype=int), + eq_obj2id: wp.array(dtype=int), + eq_objtype: wp.array(dtype=int), + eq_solimp: wp.array2d(dtype=mjwp_types.vec5), + eq_solref: wp.array2d(dtype=wp.vec2), + eq_ten_adr: wp.array(dtype=int), + eq_wld_adr: wp.array(dtype=int), + flex_bending: wp.array(dtype=wp.mat44f), + flex_damping: wp.array(dtype=float), + flex_dim: wp.array(dtype=int), + flex_edge: wp.array(dtype=wp.vec2i), + flex_edgeadr: wp.array(dtype=int), + flex_edgeflap: wp.array(dtype=wp.vec2i), + flex_elem: wp.array(dtype=int), + flex_elemedge: wp.array(dtype=int), + flex_elemedgeadr: wp.array(dtype=int), + flex_stiffness: wp.array(dtype=float), + flex_vertadr: wp.array(dtype=int), + flex_vertbodyid: wp.array(dtype=int), + flexedge_length0: wp.array(dtype=float), + geom_aabb: wp.array2d(dtype=wp.vec3), + geom_bodyid: wp.array(dtype=int), + geom_condim: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_friction: wp.array2d(dtype=wp.vec3), + geom_gap: wp.array2d(dtype=float), + geom_group: wp.array(dtype=int), + geom_margin: wp.array2d(dtype=float), + geom_matid: wp.array2d(dtype=int), + geom_pair_type_count: tuple[int, ...], + geom_plugin_index: wp.array(dtype=int), + geom_pos: wp.array2d(dtype=wp.vec3), + geom_priority: wp.array(dtype=int), + geom_quat: wp.array2d(dtype=wp.quat), + geom_rbound: wp.array2d(dtype=float), + geom_rgba: wp.array2d(dtype=wp.vec4), + geom_size: wp.array2d(dtype=wp.vec3), + geom_solimp: wp.array2d(dtype=mjwp_types.vec5), + geom_solmix: wp.array2d(dtype=float), + geom_solref: wp.array2d(dtype=wp.vec2), + geom_type: wp.array(dtype=int), + geompair2hfgeompair: wp.array(dtype=int), + has_sdf_geom: bool, + hfield_adr: wp.array(dtype=int), + hfield_data: wp.array(dtype=float), + hfield_ncol: wp.array(dtype=int), + hfield_nrow: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + jnt_actfrclimited: wp.array(dtype=bool), + jnt_actfrcrange: wp.array2d(dtype=wp.vec2), + jnt_actgravcomp: wp.array(dtype=int), + jnt_axis: wp.array2d(dtype=wp.vec3), + jnt_bodyid: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + jnt_limited_ball_adr: wp.array(dtype=int), + jnt_limited_slide_hinge_adr: wp.array(dtype=int), + jnt_margin: wp.array2d(dtype=float), + jnt_pos: wp.array2d(dtype=wp.vec3), + jnt_qposadr: wp.array(dtype=int), + jnt_range: wp.array2d(dtype=wp.vec2), + jnt_solimp: wp.array2d(dtype=mjwp_types.vec5), + jnt_solref: wp.array2d(dtype=wp.vec2), + jnt_stiffness: wp.array2d(dtype=float), + jnt_type: wp.array(dtype=int), + light_bodyid: wp.array(dtype=int), + light_dir: wp.array2d(dtype=wp.vec3), + light_dir0: wp.array2d(dtype=wp.vec3), + light_mode: wp.array(dtype=int), + light_pos: wp.array2d(dtype=wp.vec3), + light_pos0: wp.array2d(dtype=wp.vec3), + light_poscom0: wp.array2d(dtype=wp.vec3), + light_targetbodyid: wp.array(dtype=int), + mapM2M: wp.array(dtype=int), + mat_rgba: wp.array2d(dtype=wp.vec4), + mesh_face: wp.array(dtype=wp.vec3i), + mesh_faceadr: wp.array(dtype=int), + mesh_graph: wp.array(dtype=int), + mesh_graphadr: wp.array(dtype=int), + mesh_polyadr: wp.array(dtype=int), + mesh_polymap: wp.array(dtype=int), + mesh_polymapadr: wp.array(dtype=int), + mesh_polymapnum: wp.array(dtype=int), + mesh_polynormal: wp.array(dtype=wp.vec3), + mesh_polynum: wp.array(dtype=int), + mesh_polyvert: wp.array(dtype=int), + mesh_polyvertadr: wp.array(dtype=int), + mesh_polyvertnum: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_vertadr: wp.array(dtype=int), + mesh_vertnum: wp.array(dtype=int), + mocap_bodyid: wp.array(dtype=int), + nC: int, + na: int, + nbody: int, + ncam: int, + neq: int, + nflexedge: int, + nflexelem: int, + nflexvert: int, + ngeom: int, + ngravcomp: int, + nhfield: int, + njnt: int, + nlight: int, + nlsp: int, + nmeshface: int, + nmocap: int, + nsite: int, + ntendon: int, + nu: int, + nv: int, + nxn_geom_pair_filtered: wp.array(dtype=wp.vec2i), + nxn_pairid: wp.array(dtype=int), + nxn_pairid_filtered: wp.array(dtype=int), + pair_dim: wp.array(dtype=int), + pair_friction: wp.array2d(dtype=mjwp_types.vec5), + pair_gap: wp.array2d(dtype=float), + pair_margin: wp.array2d(dtype=float), + pair_solimp: wp.array2d(dtype=mjwp_types.vec5), + pair_solref: wp.array2d(dtype=wp.vec2), + pair_solreffriction: wp.array2d(dtype=wp.vec2), + plugin: wp.array(dtype=int), + plugin_attr: wp.array(dtype=wp.vec3f), + qLD_updates: tuple[wp.array(dtype=wp.vec3i), ...], + qM_fullm_i: wp.array(dtype=int), + qM_fullm_j: wp.array(dtype=int), + qM_madr_ij: wp.array(dtype=int), + qM_mulm_i: wp.array(dtype=int), + qM_mulm_j: wp.array(dtype=int), + qM_tiles: tuple[mjwp_types.TileSet, ...], + qpos0: wp.array2d(dtype=float), + qpos_spring: wp.array2d(dtype=float), + rangefinder_sensor_adr: wp.array(dtype=int), + sensor_acc_adr: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + sensor_cutoff: wp.array(dtype=float), + sensor_datatype: wp.array(dtype=int), + sensor_e_kinetic: bool, + sensor_e_potential: bool, + sensor_limitfrc_adr: wp.array(dtype=int), + sensor_limitpos_adr: wp.array(dtype=int), + sensor_limitvel_adr: wp.array(dtype=int), + sensor_objid: wp.array(dtype=int), + sensor_objtype: wp.array(dtype=int), + sensor_pos_adr: wp.array(dtype=int), + sensor_rangefinder_adr: wp.array(dtype=int), + sensor_rangefinder_bodyid: wp.array(dtype=int), + sensor_refid: wp.array(dtype=int), + sensor_reftype: wp.array(dtype=int), + sensor_rne_postconstraint: bool, + sensor_subtree_vel: bool, + sensor_tendonactfrc_adr: wp.array(dtype=int), + sensor_touch_adr: wp.array(dtype=int), + sensor_type: wp.array(dtype=int), + sensor_vel_adr: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + site_pos: wp.array2d(dtype=wp.vec3), + site_quat: wp.array2d(dtype=wp.quat), + site_size: wp.array(dtype=wp.vec3), + site_type: wp.array(dtype=int), + subtree_mass: wp.array2d(dtype=float), + tendon_actfrclimited: wp.array(dtype=bool), + tendon_actfrcrange: wp.array2d(dtype=wp.vec2), + tendon_adr: wp.array(dtype=int), + tendon_armature: wp.array2d(dtype=float), + tendon_damping: wp.array2d(dtype=float), + tendon_frictionloss: wp.array2d(dtype=float), + tendon_geom_adr: wp.array(dtype=int), + tendon_invweight0: wp.array2d(dtype=float), + tendon_jnt_adr: wp.array(dtype=int), + tendon_length0: wp.array2d(dtype=float), + tendon_lengthspring: wp.array2d(dtype=wp.vec2), + tendon_limited_adr: wp.array(dtype=int), + tendon_margin: wp.array2d(dtype=float), + tendon_num: wp.array(dtype=int), + tendon_range: wp.array2d(dtype=wp.vec2), + tendon_site_pair_adr: wp.array(dtype=int), + tendon_solimp_fri: wp.array2d(dtype=mjwp_types.vec5), + tendon_solimp_lim: wp.array2d(dtype=mjwp_types.vec5), + tendon_solref_fri: wp.array2d(dtype=wp.vec2), + tendon_solref_lim: wp.array2d(dtype=wp.vec2), + tendon_stiffness: wp.array2d(dtype=float), + wrap_geom_adr: wp.array(dtype=int), + wrap_jnt_adr: wp.array(dtype=int), + wrap_objid: wp.array(dtype=int), + wrap_prm: wp.array(dtype=float), + wrap_pulley_scale: wp.array(dtype=float), + wrap_site_pair_adr: wp.array(dtype=int), + wrap_type: wp.array(dtype=int), + opt__broadphase: int, + opt__broadphase_filter: int, + opt__cone: int, + opt__density: wp.array(dtype=float), + opt__disableflags: int, + opt__enableflags: int, + opt__epa_iterations: int, + opt__gjk_iterations: int, + opt__graph_conditional: bool, + opt__gravity: wp.array(dtype=wp.vec3), + opt__has_fluid: bool, + opt__impratio: wp.array(dtype=float), + opt__integrator: int, + opt__is_sparse: bool, + opt__iterations: int, + opt__ls_iterations: int, + opt__ls_parallel: bool, + opt__ls_tolerance: wp.array(dtype=float), + opt__magnetic: wp.array(dtype=wp.vec3), + opt__run_collision_detection: bool, + opt__sdf_initpoints: int, + opt__sdf_iterations: int, + opt__solver: int, + opt__timestep: wp.array(dtype=float), + opt__tolerance: wp.array(dtype=float), + opt__viscosity: wp.array(dtype=float), + opt__wind: wp.array(dtype=wp.vec3), + stat__meaninertia: float, + # Data + nconmax: int, + njmax: int, + act: wp.array2d(dtype=float), + act_dot: wp.array2d(dtype=float), + act_dot_rk: wp.array2d(dtype=float), + act_t0: wp.array2d(dtype=float), + actuator_force: wp.array2d(dtype=float), + actuator_length: wp.array2d(dtype=float), + actuator_moment: wp.array3d(dtype=float), + actuator_trntype_body_ncon: wp.array2d(dtype=int), + actuator_velocity: wp.array2d(dtype=float), + cacc: wp.array2d(dtype=wp.spatial_vector), + cam_xmat: wp.array2d(dtype=wp.mat33), + cam_xpos: wp.array2d(dtype=wp.vec3), + cdof: wp.array2d(dtype=wp.spatial_vector), + cdof_dot: wp.array2d(dtype=wp.spatial_vector), + cfrc_ext: wp.array2d(dtype=wp.spatial_vector), + cfrc_int: wp.array2d(dtype=wp.spatial_vector), + cinert: wp.array2d(dtype=mjwp_types.vec10), + collision_hftri_index: wp.array(dtype=int), + collision_pair: wp.array(dtype=wp.vec2i), + collision_pairid: wp.array(dtype=int), + collision_worldid: wp.array(dtype=int), + crb: wp.array2d(dtype=mjwp_types.vec10), + ctrl: wp.array2d(dtype=float), + cvel: wp.array2d(dtype=wp.spatial_vector), + energy: wp.array(dtype=wp.vec2), + epa_face: wp.array2d(dtype=wp.vec3i), + epa_horizon: wp.array2d(dtype=int), + epa_index: wp.array2d(dtype=int), + epa_map: wp.array2d(dtype=int), + epa_norm2: wp.array2d(dtype=float), + epa_pr: wp.array2d(dtype=wp.vec3), + epa_vert: wp.array2d(dtype=wp.vec3), + epa_vert1: wp.array2d(dtype=wp.vec3), + epa_vert2: wp.array2d(dtype=wp.vec3), + epa_vert_index1: wp.array2d(dtype=int), + epa_vert_index2: wp.array2d(dtype=int), + eq_active: wp.array2d(dtype=bool), + flexedge_length: wp.array2d(dtype=float), + flexedge_velocity: wp.array2d(dtype=float), + flexvert_xpos: wp.array2d(dtype=wp.vec3), + fluid_applied: wp.array2d(dtype=wp.spatial_vector), + geom_skip: wp.array(dtype=bool), + geom_xmat: wp.array2d(dtype=wp.mat33), + geom_xpos: wp.array2d(dtype=wp.vec3), + inverse_mul_m_skip: wp.array(dtype=bool), + light_xdir: wp.array2d(dtype=wp.vec3), + light_xpos: wp.array2d(dtype=wp.vec3), + mocap_pos: wp.array2d(dtype=wp.vec3), + mocap_quat: wp.array2d(dtype=wp.quat), + ncollision: wp.array(dtype=int), + ncon: wp.array(dtype=int), + ncon_hfield: wp.array2d(dtype=int), + ne: wp.array(dtype=int), + ne_connect: wp.array(dtype=int), + ne_jnt: wp.array(dtype=int), + ne_ten: wp.array(dtype=int), + ne_weld: wp.array(dtype=int), + nefc: wp.array(dtype=int), + nf: wp.array(dtype=int), + nl: wp.array(dtype=int), + nsolving: wp.array(dtype=int), + qLD: wp.array3d(dtype=float), + qLD_integration: wp.array3d(dtype=float), + qLDiagInv: wp.array2d(dtype=float), + qLDiagInv_integration: wp.array2d(dtype=float), + qM: wp.array3d(dtype=float), + qM_integration: wp.array3d(dtype=float), + qacc: wp.array2d(dtype=float), + qacc_integration: wp.array2d(dtype=float), + qacc_rk: wp.array2d(dtype=float), + qacc_smooth: wp.array2d(dtype=float), + qacc_warmstart: wp.array2d(dtype=float), + qfrc_actuator: wp.array2d(dtype=float), + qfrc_applied: wp.array2d(dtype=float), + qfrc_bias: wp.array2d(dtype=float), + qfrc_constraint: wp.array2d(dtype=float), + qfrc_damper: wp.array2d(dtype=float), + qfrc_fluid: wp.array2d(dtype=float), + qfrc_gravcomp: wp.array2d(dtype=float), + qfrc_integration: wp.array2d(dtype=float), + qfrc_passive: wp.array2d(dtype=float), + qfrc_smooth: wp.array2d(dtype=float), + qfrc_spring: wp.array2d(dtype=float), + qpos: wp.array2d(dtype=float), + qpos_t0: wp.array2d(dtype=float), + qvel: wp.array2d(dtype=float), + qvel_rk: wp.array2d(dtype=float), + qvel_t0: wp.array2d(dtype=float), + sap_cumulative_sum: wp.array2d(dtype=int), + sap_projection_lower: wp.array3d(dtype=float), + sap_projection_upper: wp.array2d(dtype=float), + sap_range: wp.array2d(dtype=int), + sap_segment_index: wp.array2d(dtype=int), + sap_sort_index: wp.array3d(dtype=int), + sensor_rangefinder_dist: wp.array2d(dtype=float), + sensor_rangefinder_geomid: wp.array2d(dtype=int), + sensor_rangefinder_pnt: wp.array2d(dtype=wp.vec3), + sensor_rangefinder_vec: wp.array2d(dtype=wp.vec3), + sensordata: wp.array2d(dtype=float), + site_xmat: wp.array2d(dtype=wp.mat33), + site_xpos: wp.array2d(dtype=wp.vec3), + solver_niter: wp.array(dtype=int), + subtree_angmom: wp.array2d(dtype=wp.vec3), + subtree_bodyvel: wp.array2d(dtype=wp.spatial_vector), + subtree_com: wp.array2d(dtype=wp.vec3), + subtree_linvel: wp.array2d(dtype=wp.vec3), + ten_J: wp.array3d(dtype=float), + ten_Jdot: wp.array3d(dtype=float), + ten_actfrc: wp.array2d(dtype=float), + ten_bias_coef: wp.array2d(dtype=float), + ten_length: wp.array2d(dtype=float), + ten_velocity: wp.array2d(dtype=float), + ten_wrapadr: wp.array2d(dtype=int), + ten_wrapnum: wp.array2d(dtype=int), + time: wp.array(dtype=float), + wrap_geom_xpos: wp.array2d(dtype=wp.spatial_vector), + wrap_obj: wp.array2d(dtype=wp.vec2i), + wrap_xpos: wp.array2d(dtype=wp.spatial_vector), + xanchor: wp.array2d(dtype=wp.vec3), + xaxis: wp.array2d(dtype=wp.vec3), + xfrc_applied: wp.array2d(dtype=wp.spatial_vector), + ximat: wp.array2d(dtype=wp.mat33), + xipos: wp.array2d(dtype=wp.vec3), + xmat: wp.array2d(dtype=wp.mat33), + xpos: wp.array2d(dtype=wp.vec3), + xquat: wp.array2d(dtype=wp.quat), + contact__dim: wp.array(dtype=int), + contact__dist: wp.array(dtype=float), + contact__efc_address: wp.array2d(dtype=int), + contact__frame: wp.array(dtype=wp.mat33), + contact__friction: wp.array(dtype=mjwp_types.vec5), + contact__geom: wp.array(dtype=wp.vec2i), + contact__includemargin: wp.array(dtype=float), + contact__pos: wp.array(dtype=wp.vec3), + contact__solimp: wp.array(dtype=mjwp_types.vec5), + contact__solref: wp.array(dtype=wp.vec2), + contact__solreffriction: wp.array(dtype=wp.vec2), + contact__worldid: wp.array(dtype=int), + efc__D: wp.array2d(dtype=float), + efc__J: wp.array3d(dtype=float), + efc__Jaref: wp.array2d(dtype=float), + efc__Ma: wp.array2d(dtype=float), + efc__Mgrad: wp.array2d(dtype=float), + efc__active: wp.array2d(dtype=bool), + efc__alpha: wp.array(dtype=float), + efc__aref: wp.array2d(dtype=float), + efc__beta: wp.array(dtype=float), + efc__beta_den: wp.array(dtype=float), + efc__beta_num: wp.array(dtype=float), + efc__cholesky_L_tmp: wp.array3d(dtype=float), + efc__cholesky_y_tmp: wp.array2d(dtype=float), + efc__condim: wp.array2d(dtype=int), + efc__cost: wp.array(dtype=float), + efc__cost_candidate: wp.array2d(dtype=float), + efc__done: wp.array(dtype=bool), + efc__force: wp.array2d(dtype=float), + efc__frictionloss: wp.array2d(dtype=float), + efc__gauss: wp.array(dtype=float), + efc__grad: wp.array2d(dtype=float), + efc__grad_dot: wp.array(dtype=float), + efc__gtol: wp.array(dtype=float), + efc__h: wp.array3d(dtype=float), + efc__hi: wp.array(dtype=wp.vec3), + efc__hi_alpha: wp.array(dtype=float), + efc__hi_next: wp.array(dtype=wp.vec3), + efc__hi_next_alpha: wp.array(dtype=float), + efc__id: wp.array2d(dtype=int), + efc__jv: wp.array2d(dtype=float), + efc__lo: wp.array(dtype=wp.vec3), + efc__lo_alpha: wp.array(dtype=float), + efc__lo_next: wp.array(dtype=wp.vec3), + efc__lo_next_alpha: wp.array(dtype=float), + efc__ls_done: wp.array(dtype=bool), + efc__margin: wp.array2d(dtype=float), + efc__mid: wp.array(dtype=wp.vec3), + efc__mid_alpha: wp.array(dtype=float), + efc__mv: wp.array2d(dtype=float), + efc__p0: wp.array(dtype=wp.vec3), + efc__pos: wp.array2d(dtype=float), + efc__prev_Mgrad: wp.array2d(dtype=float), + efc__prev_cost: wp.array(dtype=float), + efc__prev_grad: wp.array2d(dtype=float), + efc__quad: wp.array2d(dtype=wp.vec3), + efc__quad_gauss: wp.array(dtype=wp.vec3), + efc__search: wp.array2d(dtype=float), + efc__search_dot: wp.array(dtype=float), + efc__type: wp.array2d(dtype=int), + efc__u: wp.array(dtype=mjwp_types.vec6), + efc__uu: wp.array(dtype=float), + efc__uv: wp.array(dtype=float), + efc__vel: wp.array2d(dtype=float), + efc__vv: wp.array(dtype=float), +): + _m.stat = _s + _m.opt = _o + _d.efc = _e + _d.contact = _c + _m.M_rowadr = M_rowadr + _m.M_rownnz = M_rownnz + _m.actuator_acc0 = actuator_acc0 + _m.actuator_actadr = actuator_actadr + _m.actuator_actearly = actuator_actearly + _m.actuator_actlimited = actuator_actlimited + _m.actuator_actnum = actuator_actnum + _m.actuator_actrange = actuator_actrange + _m.actuator_affine_bias_gain = actuator_affine_bias_gain + _m.actuator_biasprm = actuator_biasprm + _m.actuator_biastype = actuator_biastype + _m.actuator_cranklength = actuator_cranklength + _m.actuator_ctrllimited = actuator_ctrllimited + _m.actuator_ctrlrange = actuator_ctrlrange + _m.actuator_dynprm = actuator_dynprm + _m.actuator_dyntype = actuator_dyntype + _m.actuator_forcelimited = actuator_forcelimited + _m.actuator_forcerange = actuator_forcerange + _m.actuator_gainprm = actuator_gainprm + _m.actuator_gaintype = actuator_gaintype + _m.actuator_gear = actuator_gear + _m.actuator_lengthrange = actuator_lengthrange + _m.actuator_trnid = actuator_trnid + _m.actuator_trntype = actuator_trntype + _m.actuator_trntype_body_adr = actuator_trntype_body_adr + _m.block_dim = block_dim + _m.body_dofadr = body_dofadr + _m.body_dofnum = body_dofnum + _m.body_gravcomp = body_gravcomp + _m.body_inertia = body_inertia + _m.body_invweight0 = body_invweight0 + _m.body_ipos = body_ipos + _m.body_iquat = body_iquat + _m.body_jntadr = body_jntadr + _m.body_jntnum = body_jntnum + _m.body_mass = body_mass + _m.body_parentid = body_parentid + _m.body_pos = body_pos + _m.body_quat = body_quat + _m.body_rootid = body_rootid + _m.body_subtreemass = body_subtreemass + _m.body_tree = body_tree + _m.body_weldid = body_weldid + _m.cam_bodyid = cam_bodyid + _m.cam_fovy = cam_fovy + _m.cam_intrinsic = cam_intrinsic + _m.cam_mat0 = cam_mat0 + _m.cam_mode = cam_mode + _m.cam_pos = cam_pos + _m.cam_pos0 = cam_pos0 + _m.cam_poscom0 = cam_poscom0 + _m.cam_quat = cam_quat + _m.cam_resolution = cam_resolution + _m.cam_sensorsize = cam_sensorsize + _m.cam_targetbodyid = cam_targetbodyid + _m.condim_max = condim_max + _m.dof_Madr = dof_Madr + _m.dof_armature = dof_armature + _m.dof_bodyid = dof_bodyid + _m.dof_damping = dof_damping + _m.dof_frictionloss = dof_frictionloss + _m.dof_invweight0 = dof_invweight0 + _m.dof_jntid = dof_jntid + _m.dof_parentid = dof_parentid + _m.dof_solimp = dof_solimp + _m.dof_solref = dof_solref + _m.dof_tri_col = dof_tri_col + _m.dof_tri_row = dof_tri_row + _m.eq_connect_adr = eq_connect_adr + _m.eq_data = eq_data + _m.eq_jnt_adr = eq_jnt_adr + _m.eq_obj1id = eq_obj1id + _m.eq_obj2id = eq_obj2id + _m.eq_objtype = eq_objtype + _m.eq_solimp = eq_solimp + _m.eq_solref = eq_solref + _m.eq_ten_adr = eq_ten_adr + _m.eq_wld_adr = eq_wld_adr + _m.flex_bending = flex_bending + _m.flex_damping = flex_damping + _m.flex_dim = flex_dim + _m.flex_edge = flex_edge + _m.flex_edgeadr = flex_edgeadr + _m.flex_edgeflap = flex_edgeflap + _m.flex_elem = flex_elem + _m.flex_elemedge = flex_elemedge + _m.flex_elemedgeadr = flex_elemedgeadr + _m.flex_stiffness = flex_stiffness + _m.flex_vertadr = flex_vertadr + _m.flex_vertbodyid = flex_vertbodyid + _m.flexedge_length0 = flexedge_length0 + _m.geom_aabb = geom_aabb + _m.geom_bodyid = geom_bodyid + _m.geom_condim = geom_condim + _m.geom_dataid = geom_dataid + _m.geom_friction = geom_friction + _m.geom_gap = geom_gap + _m.geom_group = geom_group + _m.geom_margin = geom_margin + _m.geom_matid = geom_matid + _m.geom_pair_type_count = geom_pair_type_count + _m.geom_plugin_index = geom_plugin_index + _m.geom_pos = geom_pos + _m.geom_priority = geom_priority + _m.geom_quat = geom_quat + _m.geom_rbound = geom_rbound + _m.geom_rgba = geom_rgba + _m.geom_size = geom_size + _m.geom_solimp = geom_solimp + _m.geom_solmix = geom_solmix + _m.geom_solref = geom_solref + _m.geom_type = geom_type + _m.geompair2hfgeompair = geompair2hfgeompair + _m.has_sdf_geom = has_sdf_geom + _m.hfield_adr = hfield_adr + _m.hfield_data = hfield_data + _m.hfield_ncol = hfield_ncol + _m.hfield_nrow = hfield_nrow + _m.hfield_size = hfield_size + _m.jnt_actfrclimited = jnt_actfrclimited + _m.jnt_actfrcrange = jnt_actfrcrange + _m.jnt_actgravcomp = jnt_actgravcomp + _m.jnt_axis = jnt_axis + _m.jnt_bodyid = jnt_bodyid + _m.jnt_dofadr = jnt_dofadr + _m.jnt_limited_ball_adr = jnt_limited_ball_adr + _m.jnt_limited_slide_hinge_adr = jnt_limited_slide_hinge_adr + _m.jnt_margin = jnt_margin + _m.jnt_pos = jnt_pos + _m.jnt_qposadr = jnt_qposadr + _m.jnt_range = jnt_range + _m.jnt_solimp = jnt_solimp + _m.jnt_solref = jnt_solref + _m.jnt_stiffness = jnt_stiffness + _m.jnt_type = jnt_type + _m.light_bodyid = light_bodyid + _m.light_dir = light_dir + _m.light_dir0 = light_dir0 + _m.light_mode = light_mode + _m.light_pos = light_pos + _m.light_pos0 = light_pos0 + _m.light_poscom0 = light_poscom0 + _m.light_targetbodyid = light_targetbodyid + _m.mapM2M = mapM2M + _m.mat_rgba = mat_rgba + _m.mesh_face = mesh_face + _m.mesh_faceadr = mesh_faceadr + _m.mesh_graph = mesh_graph + _m.mesh_graphadr = mesh_graphadr + _m.mesh_polyadr = mesh_polyadr + _m.mesh_polymap = mesh_polymap + _m.mesh_polymapadr = mesh_polymapadr + _m.mesh_polymapnum = mesh_polymapnum + _m.mesh_polynormal = mesh_polynormal + _m.mesh_polynum = mesh_polynum + _m.mesh_polyvert = mesh_polyvert + _m.mesh_polyvertadr = mesh_polyvertadr + _m.mesh_polyvertnum = mesh_polyvertnum + _m.mesh_vert = mesh_vert + _m.mesh_vertadr = mesh_vertadr + _m.mesh_vertnum = mesh_vertnum + _m.mocap_bodyid = mocap_bodyid + _m.nC = nC + _m.na = na + _m.nbody = nbody + _m.ncam = ncam + _m.neq = neq + _m.nflexedge = nflexedge + _m.nflexelem = nflexelem + _m.nflexvert = nflexvert + _m.ngeom = ngeom + _m.ngravcomp = ngravcomp + _m.nhfield = nhfield + _m.njnt = njnt + _m.nlight = nlight + _m.nlsp = nlsp + _m.nmeshface = nmeshface + _m.nmocap = nmocap + _m.nsite = nsite + _m.ntendon = ntendon + _m.nu = nu + _m.nv = nv + _m.nxn_geom_pair_filtered = nxn_geom_pair_filtered + _m.nxn_pairid = nxn_pairid + _m.nxn_pairid_filtered = nxn_pairid_filtered + _m.opt.broadphase = opt__broadphase + _m.opt.broadphase_filter = opt__broadphase_filter + _m.opt.cone = opt__cone + _m.opt.density = opt__density + _m.opt.disableflags = opt__disableflags + _m.opt.enableflags = opt__enableflags + _m.opt.epa_iterations = opt__epa_iterations + _m.opt.gjk_iterations = opt__gjk_iterations + _m.opt.graph_conditional = opt__graph_conditional + _m.opt.gravity = opt__gravity + _m.opt.has_fluid = opt__has_fluid + _m.opt.impratio = opt__impratio + _m.opt.integrator = opt__integrator + _m.opt.is_sparse = opt__is_sparse + _m.opt.iterations = opt__iterations + _m.opt.ls_iterations = opt__ls_iterations + _m.opt.ls_parallel = opt__ls_parallel + _m.opt.ls_tolerance = opt__ls_tolerance + _m.opt.magnetic = opt__magnetic + _m.opt.run_collision_detection = opt__run_collision_detection + _m.opt.sdf_initpoints = opt__sdf_initpoints + _m.opt.sdf_iterations = opt__sdf_iterations + _m.opt.solver = opt__solver + _m.opt.timestep = opt__timestep + _m.opt.tolerance = opt__tolerance + _m.opt.viscosity = opt__viscosity + _m.opt.wind = opt__wind + _m.pair_dim = pair_dim + _m.pair_friction = pair_friction + _m.pair_gap = pair_gap + _m.pair_margin = pair_margin + _m.pair_solimp = pair_solimp + _m.pair_solref = pair_solref + _m.pair_solreffriction = pair_solreffriction + _m.plugin = plugin + _m.plugin_attr = plugin_attr + _m.qLD_updates = qLD_updates + _m.qM_fullm_i = qM_fullm_i + _m.qM_fullm_j = qM_fullm_j + _m.qM_madr_ij = qM_madr_ij + _m.qM_mulm_i = qM_mulm_i + _m.qM_mulm_j = qM_mulm_j + _m.qM_tiles = qM_tiles + _m.qpos0 = qpos0 + _m.qpos_spring = qpos_spring + _m.rangefinder_sensor_adr = rangefinder_sensor_adr + _m.sensor_acc_adr = sensor_acc_adr + _m.sensor_adr = sensor_adr + _m.sensor_cutoff = sensor_cutoff + _m.sensor_datatype = sensor_datatype + _m.sensor_e_kinetic = sensor_e_kinetic + _m.sensor_e_potential = sensor_e_potential + _m.sensor_limitfrc_adr = sensor_limitfrc_adr + _m.sensor_limitpos_adr = sensor_limitpos_adr + _m.sensor_limitvel_adr = sensor_limitvel_adr + _m.sensor_objid = sensor_objid + _m.sensor_objtype = sensor_objtype + _m.sensor_pos_adr = sensor_pos_adr + _m.sensor_rangefinder_adr = sensor_rangefinder_adr + _m.sensor_rangefinder_bodyid = sensor_rangefinder_bodyid + _m.sensor_refid = sensor_refid + _m.sensor_reftype = sensor_reftype + _m.sensor_rne_postconstraint = sensor_rne_postconstraint + _m.sensor_subtree_vel = sensor_subtree_vel + _m.sensor_tendonactfrc_adr = sensor_tendonactfrc_adr + _m.sensor_touch_adr = sensor_touch_adr + _m.sensor_type = sensor_type + _m.sensor_vel_adr = sensor_vel_adr + _m.site_bodyid = site_bodyid + _m.site_pos = site_pos + _m.site_quat = site_quat + _m.site_size = site_size + _m.site_type = site_type + _m.stat.meaninertia = stat__meaninertia + _m.subtree_mass = subtree_mass + _m.tendon_actfrclimited = tendon_actfrclimited + _m.tendon_actfrcrange = tendon_actfrcrange + _m.tendon_adr = tendon_adr + _m.tendon_armature = tendon_armature + _m.tendon_damping = tendon_damping + _m.tendon_frictionloss = tendon_frictionloss + _m.tendon_geom_adr = tendon_geom_adr + _m.tendon_invweight0 = tendon_invweight0 + _m.tendon_jnt_adr = tendon_jnt_adr + _m.tendon_length0 = tendon_length0 + _m.tendon_lengthspring = tendon_lengthspring + _m.tendon_limited_adr = tendon_limited_adr + _m.tendon_margin = tendon_margin + _m.tendon_num = tendon_num + _m.tendon_range = tendon_range + _m.tendon_site_pair_adr = tendon_site_pair_adr + _m.tendon_solimp_fri = tendon_solimp_fri + _m.tendon_solimp_lim = tendon_solimp_lim + _m.tendon_solref_fri = tendon_solref_fri + _m.tendon_solref_lim = tendon_solref_lim + _m.tendon_stiffness = tendon_stiffness + _m.wrap_geom_adr = wrap_geom_adr + _m.wrap_jnt_adr = wrap_jnt_adr + _m.wrap_objid = wrap_objid + _m.wrap_prm = wrap_prm + _m.wrap_pulley_scale = wrap_pulley_scale + _m.wrap_site_pair_adr = wrap_site_pair_adr + _m.wrap_type = wrap_type + _d.act = act + _d.act_dot = act_dot + _d.act_dot_rk = act_dot_rk + _d.act_t0 = act_t0 + _d.actuator_force = actuator_force + _d.actuator_length = actuator_length + _d.actuator_moment = actuator_moment + _d.actuator_trntype_body_ncon = actuator_trntype_body_ncon + _d.actuator_velocity = actuator_velocity + _d.cacc = cacc + _d.cam_xmat = cam_xmat + _d.cam_xpos = cam_xpos + _d.cdof = cdof + _d.cdof_dot = cdof_dot + _d.cfrc_ext = cfrc_ext + _d.cfrc_int = cfrc_int + _d.cinert = cinert + _d.collision_hftri_index = collision_hftri_index + _d.collision_pair = collision_pair + _d.collision_pairid = collision_pairid + _d.collision_worldid = collision_worldid + _d.contact.dim = contact__dim + _d.contact.dist = contact__dist + _d.contact.efc_address = contact__efc_address + _d.contact.frame = contact__frame + _d.contact.friction = contact__friction + _d.contact.geom = contact__geom + _d.contact.includemargin = contact__includemargin + _d.contact.pos = contact__pos + _d.contact.solimp = contact__solimp + _d.contact.solref = contact__solref + _d.contact.solreffriction = contact__solreffriction + _d.contact.worldid = contact__worldid + _d.crb = crb + _d.ctrl = ctrl + _d.cvel = cvel + _d.efc.D = efc__D + _d.efc.J = efc__J + _d.efc.Jaref = efc__Jaref + _d.efc.Ma = efc__Ma + _d.efc.Mgrad = efc__Mgrad + _d.efc.active = efc__active + _d.efc.alpha = efc__alpha + _d.efc.aref = efc__aref + _d.efc.beta = efc__beta + _d.efc.beta_den = efc__beta_den + _d.efc.beta_num = efc__beta_num + _d.efc.cholesky_L_tmp = efc__cholesky_L_tmp + _d.efc.cholesky_y_tmp = efc__cholesky_y_tmp + _d.efc.condim = efc__condim + _d.efc.cost = efc__cost + _d.efc.cost_candidate = efc__cost_candidate + _d.efc.done = efc__done + _d.efc.force = efc__force + _d.efc.frictionloss = efc__frictionloss + _d.efc.gauss = efc__gauss + _d.efc.grad = efc__grad + _d.efc.grad_dot = efc__grad_dot + _d.efc.gtol = efc__gtol + _d.efc.h = efc__h + _d.efc.hi = efc__hi + _d.efc.hi_alpha = efc__hi_alpha + _d.efc.hi_next = efc__hi_next + _d.efc.hi_next_alpha = efc__hi_next_alpha + _d.efc.id = efc__id + _d.efc.jv = efc__jv + _d.efc.lo = efc__lo + _d.efc.lo_alpha = efc__lo_alpha + _d.efc.lo_next = efc__lo_next + _d.efc.lo_next_alpha = efc__lo_next_alpha + _d.efc.ls_done = efc__ls_done + _d.efc.margin = efc__margin + _d.efc.mid = efc__mid + _d.efc.mid_alpha = efc__mid_alpha + _d.efc.mv = efc__mv + _d.efc.p0 = efc__p0 + _d.efc.pos = efc__pos + _d.efc.prev_Mgrad = efc__prev_Mgrad + _d.efc.prev_cost = efc__prev_cost + _d.efc.prev_grad = efc__prev_grad + _d.efc.quad = efc__quad + _d.efc.quad_gauss = efc__quad_gauss + _d.efc.search = efc__search + _d.efc.search_dot = efc__search_dot + _d.efc.type = efc__type + _d.efc.u = efc__u + _d.efc.uu = efc__uu + _d.efc.uv = efc__uv + _d.efc.vel = efc__vel + _d.efc.vv = efc__vv + _d.energy = energy + _d.epa_face = epa_face + _d.epa_horizon = epa_horizon + _d.epa_index = epa_index + _d.epa_map = epa_map + _d.epa_norm2 = epa_norm2 + _d.epa_pr = epa_pr + _d.epa_vert = epa_vert + _d.epa_vert1 = epa_vert1 + _d.epa_vert2 = epa_vert2 + _d.epa_vert_index1 = epa_vert_index1 + _d.epa_vert_index2 = epa_vert_index2 + _d.eq_active = eq_active + _d.flexedge_length = flexedge_length + _d.flexedge_velocity = flexedge_velocity + _d.flexvert_xpos = flexvert_xpos + _d.fluid_applied = fluid_applied + _d.geom_skip = geom_skip + _d.geom_xmat = geom_xmat + _d.geom_xpos = geom_xpos + _d.inverse_mul_m_skip = inverse_mul_m_skip + _d.light_xdir = light_xdir + _d.light_xpos = light_xpos + _d.mocap_pos = mocap_pos + _d.mocap_quat = mocap_quat + _d.ncollision = ncollision + _d.ncon = ncon + _d.ncon_hfield = ncon_hfield + _d.nconmax = nconmax + _d.ne = ne + _d.ne_connect = ne_connect + _d.ne_jnt = ne_jnt + _d.ne_ten = ne_ten + _d.ne_weld = ne_weld + _d.nefc = nefc + _d.nf = nf + _d.njmax = njmax + _d.nl = nl + _d.nsolving = nsolving + _d.qLD = qLD + _d.qLD_integration = qLD_integration + _d.qLDiagInv = qLDiagInv + _d.qLDiagInv_integration = qLDiagInv_integration + _d.qM = qM + _d.qM_integration = qM_integration + _d.qacc = qacc + _d.qacc_integration = qacc_integration + _d.qacc_rk = qacc_rk + _d.qacc_smooth = qacc_smooth + _d.qacc_warmstart = qacc_warmstart + _d.qfrc_actuator = qfrc_actuator + _d.qfrc_applied = qfrc_applied + _d.qfrc_bias = qfrc_bias + _d.qfrc_constraint = qfrc_constraint + _d.qfrc_damper = qfrc_damper + _d.qfrc_fluid = qfrc_fluid + _d.qfrc_gravcomp = qfrc_gravcomp + _d.qfrc_integration = qfrc_integration + _d.qfrc_passive = qfrc_passive + _d.qfrc_smooth = qfrc_smooth + _d.qfrc_spring = qfrc_spring + _d.qpos = qpos + _d.qpos_t0 = qpos_t0 + _d.qvel = qvel + _d.qvel_rk = qvel_rk + _d.qvel_t0 = qvel_t0 + _d.sap_cumulative_sum = sap_cumulative_sum + _d.sap_projection_lower = sap_projection_lower + _d.sap_projection_upper = sap_projection_upper + _d.sap_range = sap_range + _d.sap_segment_index = sap_segment_index + _d.sap_sort_index = sap_sort_index + _d.sensor_rangefinder_dist = sensor_rangefinder_dist + _d.sensor_rangefinder_geomid = sensor_rangefinder_geomid + _d.sensor_rangefinder_pnt = sensor_rangefinder_pnt + _d.sensor_rangefinder_vec = sensor_rangefinder_vec + _d.sensordata = sensordata + _d.site_xmat = site_xmat + _d.site_xpos = site_xpos + _d.solver_niter = solver_niter + _d.subtree_angmom = subtree_angmom + _d.subtree_bodyvel = subtree_bodyvel + _d.subtree_com = subtree_com + _d.subtree_linvel = subtree_linvel + _d.ten_J = ten_J + _d.ten_Jdot = ten_Jdot + _d.ten_actfrc = ten_actfrc + _d.ten_bias_coef = ten_bias_coef + _d.ten_length = ten_length + _d.ten_velocity = ten_velocity + _d.ten_wrapadr = ten_wrapadr + _d.ten_wrapnum = ten_wrapnum + _d.time = time + _d.wrap_geom_xpos = wrap_geom_xpos + _d.wrap_obj = wrap_obj + _d.wrap_xpos = wrap_xpos + _d.xanchor = xanchor + _d.xaxis = xaxis + _d.xfrc_applied = xfrc_applied + _d.ximat = ximat + _d.xipos = xipos + _d.xmat = xmat + _d.xpos = xpos + _d.xquat = xquat + _d.nworld = nworld + mjwarp.step(_m, _d) + + +def _step_jax_impl(m: types.Model, d: types.Data): + output_dims = { + 'act': d.act.shape, + 'act_dot': d.act_dot.shape, + 'act_dot_rk': d._impl.act_dot_rk.shape, + 'act_t0': d._impl.act_t0.shape, + 'actuator_force': d.actuator_force.shape, + 'actuator_length': d._impl.actuator_length.shape, + 'actuator_moment': d._impl.actuator_moment.shape, + 'actuator_trntype_body_ncon': d._impl.actuator_trntype_body_ncon.shape, + 'actuator_velocity': d._impl.actuator_velocity.shape, + 'cacc': d._impl.cacc.shape, + 'cam_xmat': d.cam_xmat.shape, + 'cam_xpos': d.cam_xpos.shape, + 'cdof': d._impl.cdof.shape, + 'cdof_dot': d._impl.cdof_dot.shape, + 'cfrc_ext': d._impl.cfrc_ext.shape, + 'cfrc_int': d._impl.cfrc_int.shape, + 'cinert': d._impl.cinert.shape, + 'collision_hftri_index': d._impl.collision_hftri_index.shape, + 'collision_pair': d._impl.collision_pair.shape, + 'collision_pairid': d._impl.collision_pairid.shape, + 'collision_worldid': d._impl.collision_worldid.shape, + 'crb': d._impl.crb.shape, + 'ctrl': d.ctrl.shape, + 'cvel': d.cvel.shape, + 'energy': d._impl.energy.shape, + 'epa_face': d._impl.epa_face.shape, + 'epa_horizon': d._impl.epa_horizon.shape, + 'epa_index': d._impl.epa_index.shape, + 'epa_map': d._impl.epa_map.shape, + 'epa_norm2': d._impl.epa_norm2.shape, + 'epa_pr': d._impl.epa_pr.shape, + 'epa_vert': d._impl.epa_vert.shape, + 'epa_vert1': d._impl.epa_vert1.shape, + 'epa_vert2': d._impl.epa_vert2.shape, + 'epa_vert_index1': d._impl.epa_vert_index1.shape, + 'epa_vert_index2': d._impl.epa_vert_index2.shape, + 'eq_active': d.eq_active.shape, + 'flexedge_length': d._impl.flexedge_length.shape, + 'flexedge_velocity': d._impl.flexedge_velocity.shape, + 'flexvert_xpos': d._impl.flexvert_xpos.shape, + 'fluid_applied': d._impl.fluid_applied.shape, + 'geom_skip': d._impl.geom_skip.shape, + 'geom_xmat': d.geom_xmat.shape, + 'geom_xpos': d.geom_xpos.shape, + 'inverse_mul_m_skip': d._impl.inverse_mul_m_skip.shape, + 'light_xdir': d._impl.light_xdir.shape, + 'light_xpos': d._impl.light_xpos.shape, + 'mocap_pos': d.mocap_pos.shape, + 'mocap_quat': d.mocap_quat.shape, + 'ncollision': d._impl.ncollision.shape, + 'ncon': d._impl.ncon.shape, + 'ncon_hfield': d._impl.ncon_hfield.shape, + 'ne': d._impl.ne.shape, + 'ne_connect': d._impl.ne_connect.shape, + 'ne_jnt': d._impl.ne_jnt.shape, + 'ne_ten': d._impl.ne_ten.shape, + 'ne_weld': d._impl.ne_weld.shape, + 'nefc': d._impl.nefc.shape, + 'nf': d._impl.nf.shape, + 'nl': d._impl.nl.shape, + 'nsolving': d._impl.nsolving.shape, + 'qLD': d._impl.qLD.shape, + 'qLD_integration': d._impl.qLD_integration.shape, + 'qLDiagInv': d._impl.qLDiagInv.shape, + 'qLDiagInv_integration': d._impl.qLDiagInv_integration.shape, + 'qM': d._impl.qM.shape, + 'qM_integration': d._impl.qM_integration.shape, + 'qacc': d.qacc.shape, + 'qacc_integration': d._impl.qacc_integration.shape, + 'qacc_rk': d._impl.qacc_rk.shape, + 'qacc_smooth': d.qacc_smooth.shape, + 'qacc_warmstart': d.qacc_warmstart.shape, + 'qfrc_actuator': d.qfrc_actuator.shape, + 'qfrc_applied': d.qfrc_applied.shape, + 'qfrc_bias': d.qfrc_bias.shape, + 'qfrc_constraint': d.qfrc_constraint.shape, + 'qfrc_damper': d._impl.qfrc_damper.shape, + 'qfrc_fluid': d.qfrc_fluid.shape, + 'qfrc_gravcomp': d.qfrc_gravcomp.shape, + 'qfrc_integration': d._impl.qfrc_integration.shape, + 'qfrc_passive': d.qfrc_passive.shape, + 'qfrc_smooth': d.qfrc_smooth.shape, + 'qfrc_spring': d._impl.qfrc_spring.shape, + 'qpos': d.qpos.shape, + 'qpos_t0': d._impl.qpos_t0.shape, + 'qvel': d.qvel.shape, + 'qvel_rk': d._impl.qvel_rk.shape, + 'qvel_t0': d._impl.qvel_t0.shape, + 'sap_cumulative_sum': d._impl.sap_cumulative_sum.shape, + 'sap_projection_lower': d._impl.sap_projection_lower.shape, + 'sap_projection_upper': d._impl.sap_projection_upper.shape, + 'sap_range': d._impl.sap_range.shape, + 'sap_segment_index': d._impl.sap_segment_index.shape, + 'sap_sort_index': d._impl.sap_sort_index.shape, + 'sensor_rangefinder_dist': d._impl.sensor_rangefinder_dist.shape, + 'sensor_rangefinder_geomid': d._impl.sensor_rangefinder_geomid.shape, + 'sensor_rangefinder_pnt': d._impl.sensor_rangefinder_pnt.shape, + 'sensor_rangefinder_vec': d._impl.sensor_rangefinder_vec.shape, + 'sensordata': d.sensordata.shape, + 'site_xmat': d.site_xmat.shape, + 'site_xpos': d.site_xpos.shape, + 'solver_niter': d._impl.solver_niter.shape, + 'subtree_angmom': d._impl.subtree_angmom.shape, + 'subtree_bodyvel': d._impl.subtree_bodyvel.shape, + 'subtree_com': d.subtree_com.shape, + 'subtree_linvel': d._impl.subtree_linvel.shape, + 'ten_J': d._impl.ten_J.shape, + 'ten_Jdot': d._impl.ten_Jdot.shape, + 'ten_actfrc': d._impl.ten_actfrc.shape, + 'ten_bias_coef': d._impl.ten_bias_coef.shape, + 'ten_length': d._impl.ten_length.shape, + 'ten_velocity': d._impl.ten_velocity.shape, + 'ten_wrapadr': d._impl.ten_wrapadr.shape, + 'ten_wrapnum': d._impl.ten_wrapnum.shape, + 'time': d.time.shape, + 'wrap_geom_xpos': d._impl.wrap_geom_xpos.shape, + 'wrap_obj': d._impl.wrap_obj.shape, + 'wrap_xpos': d._impl.wrap_xpos.shape, + 'xanchor': d.xanchor.shape, + 'xaxis': d.xaxis.shape, + 'xfrc_applied': d.xfrc_applied.shape, + 'ximat': d.ximat.shape, + 'xipos': d.xipos.shape, + 'xmat': d.xmat.shape, + 'xpos': d.xpos.shape, + 'xquat': d.xquat.shape, + 'contact__dim': d._impl.contact__dim.shape, + 'contact__dist': d._impl.contact__dist.shape, + 'contact__efc_address': d._impl.contact__efc_address.shape, + 'contact__frame': d._impl.contact__frame.shape, + 'contact__friction': d._impl.contact__friction.shape, + 'contact__geom': d._impl.contact__geom.shape, + 'contact__includemargin': d._impl.contact__includemargin.shape, + 'contact__pos': d._impl.contact__pos.shape, + 'contact__solimp': d._impl.contact__solimp.shape, + 'contact__solref': d._impl.contact__solref.shape, + 'contact__solreffriction': d._impl.contact__solreffriction.shape, + 'contact__worldid': d._impl.contact__worldid.shape, + 'efc__D': d._impl.efc__D.shape, + 'efc__J': d._impl.efc__J.shape, + 'efc__Jaref': d._impl.efc__Jaref.shape, + 'efc__Ma': d._impl.efc__Ma.shape, + 'efc__Mgrad': d._impl.efc__Mgrad.shape, + 'efc__active': d._impl.efc__active.shape, + 'efc__alpha': d._impl.efc__alpha.shape, + 'efc__aref': d._impl.efc__aref.shape, + 'efc__beta': d._impl.efc__beta.shape, + 'efc__beta_den': d._impl.efc__beta_den.shape, + 'efc__beta_num': d._impl.efc__beta_num.shape, + 'efc__cholesky_L_tmp': d._impl.efc__cholesky_L_tmp.shape, + 'efc__cholesky_y_tmp': d._impl.efc__cholesky_y_tmp.shape, + 'efc__condim': d._impl.efc__condim.shape, + 'efc__cost': d._impl.efc__cost.shape, + 'efc__cost_candidate': d._impl.efc__cost_candidate.shape, + 'efc__done': d._impl.efc__done.shape, + 'efc__force': d._impl.efc__force.shape, + 'efc__frictionloss': d._impl.efc__frictionloss.shape, + 'efc__gauss': d._impl.efc__gauss.shape, + 'efc__grad': d._impl.efc__grad.shape, + 'efc__grad_dot': d._impl.efc__grad_dot.shape, + 'efc__gtol': d._impl.efc__gtol.shape, + 'efc__h': d._impl.efc__h.shape, + 'efc__hi': d._impl.efc__hi.shape, + 'efc__hi_alpha': d._impl.efc__hi_alpha.shape, + 'efc__hi_next': d._impl.efc__hi_next.shape, + 'efc__hi_next_alpha': d._impl.efc__hi_next_alpha.shape, + 'efc__id': d._impl.efc__id.shape, + 'efc__jv': d._impl.efc__jv.shape, + 'efc__lo': d._impl.efc__lo.shape, + 'efc__lo_alpha': d._impl.efc__lo_alpha.shape, + 'efc__lo_next': d._impl.efc__lo_next.shape, + 'efc__lo_next_alpha': d._impl.efc__lo_next_alpha.shape, + 'efc__ls_done': d._impl.efc__ls_done.shape, + 'efc__margin': d._impl.efc__margin.shape, + 'efc__mid': d._impl.efc__mid.shape, + 'efc__mid_alpha': d._impl.efc__mid_alpha.shape, + 'efc__mv': d._impl.efc__mv.shape, + 'efc__p0': d._impl.efc__p0.shape, + 'efc__pos': d._impl.efc__pos.shape, + 'efc__prev_Mgrad': d._impl.efc__prev_Mgrad.shape, + 'efc__prev_cost': d._impl.efc__prev_cost.shape, + 'efc__prev_grad': d._impl.efc__prev_grad.shape, + 'efc__quad': d._impl.efc__quad.shape, + 'efc__quad_gauss': d._impl.efc__quad_gauss.shape, + 'efc__search': d._impl.efc__search.shape, + 'efc__search_dot': d._impl.efc__search_dot.shape, + 'efc__type': d._impl.efc__type.shape, + 'efc__u': d._impl.efc__u.shape, + 'efc__uu': d._impl.efc__uu.shape, + 'efc__uv': d._impl.efc__uv.shape, + 'efc__vel': d._impl.efc__vel.shape, + 'efc__vv': d._impl.efc__vv.shape, + } + jf = ffi.jax_callable_variadic_tuple( + _step_shim, + num_outputs=192, + output_dims=output_dims, + vmap_method=None, + graph_compatible=True, + in_out_argnames={ + 'act', + 'act_dot', + 'act_dot_rk', + 'act_t0', + 'actuator_force', + 'actuator_length', + 'actuator_moment', + 'actuator_trntype_body_ncon', + 'actuator_velocity', + 'cacc', + 'cam_xmat', + 'cam_xpos', + 'cdof', + 'cdof_dot', + 'cfrc_ext', + 'cfrc_int', + 'cinert', + 'collision_hftri_index', + 'collision_pair', + 'collision_pairid', + 'collision_worldid', + 'crb', + 'ctrl', + 'cvel', + 'energy', + 'epa_face', + 'epa_horizon', + 'epa_index', + 'epa_map', + 'epa_norm2', + 'epa_pr', + 'epa_vert', + 'epa_vert1', + 'epa_vert2', + 'epa_vert_index1', + 'epa_vert_index2', + 'eq_active', + 'flexedge_length', + 'flexedge_velocity', + 'flexvert_xpos', + 'fluid_applied', + 'geom_skip', + 'geom_xmat', + 'geom_xpos', + 'inverse_mul_m_skip', + 'light_xdir', + 'light_xpos', + 'mocap_pos', + 'mocap_quat', + 'ncollision', + 'ncon', + 'ncon_hfield', + 'ne', + 'ne_connect', + 'ne_jnt', + 'ne_ten', + 'ne_weld', + 'nefc', + 'nf', + 'nl', + 'nsolving', + 'qLD', + 'qLD_integration', + 'qLDiagInv', + 'qLDiagInv_integration', + 'qM', + 'qM_integration', + 'qacc', + 'qacc_integration', + 'qacc_rk', + 'qacc_smooth', + 'qacc_warmstart', + 'qfrc_actuator', + 'qfrc_applied', + 'qfrc_bias', + 'qfrc_constraint', + 'qfrc_damper', + 'qfrc_fluid', + 'qfrc_gravcomp', + 'qfrc_integration', + 'qfrc_passive', + 'qfrc_smooth', + 'qfrc_spring', + 'qpos', + 'qpos_t0', + 'qvel', + 'qvel_rk', + 'qvel_t0', + 'sap_cumulative_sum', + 'sap_projection_lower', + 'sap_projection_upper', + 'sap_range', + 'sap_segment_index', + 'sap_sort_index', + 'sensor_rangefinder_dist', + 'sensor_rangefinder_geomid', + 'sensor_rangefinder_pnt', + 'sensor_rangefinder_vec', + 'sensordata', + 'site_xmat', + 'site_xpos', + 'solver_niter', + 'subtree_angmom', + 'subtree_bodyvel', + 'subtree_com', + 'subtree_linvel', + 'ten_J', + 'ten_Jdot', + 'ten_actfrc', + 'ten_bias_coef', + 'ten_length', + 'ten_velocity', + 'ten_wrapadr', + 'ten_wrapnum', + 'time', + 'wrap_geom_xpos', + 'wrap_obj', + 'wrap_xpos', + 'xanchor', + 'xaxis', + 'xfrc_applied', + 'ximat', + 'xipos', + 'xmat', + 'xpos', + 'xquat', + 'contact__dim', + 'contact__dist', + 'contact__efc_address', + 'contact__frame', + 'contact__friction', + 'contact__geom', + 'contact__includemargin', + 'contact__pos', + 'contact__solimp', + 'contact__solref', + 'contact__solreffriction', + 'contact__worldid', + 'efc__D', + 'efc__J', + 'efc__Jaref', + 'efc__Ma', + 'efc__Mgrad', + 'efc__active', + 'efc__alpha', + 'efc__aref', + 'efc__beta', + 'efc__beta_den', + 'efc__beta_num', + 'efc__cholesky_L_tmp', + 'efc__cholesky_y_tmp', + 'efc__condim', + 'efc__cost', + 'efc__cost_candidate', + 'efc__done', + 'efc__force', + 'efc__frictionloss', + 'efc__gauss', + 'efc__grad', + 'efc__grad_dot', + 'efc__gtol', + 'efc__h', + 'efc__hi', + 'efc__hi_alpha', + 'efc__hi_next', + 'efc__hi_next_alpha', + 'efc__id', + 'efc__jv', + 'efc__lo', + 'efc__lo_alpha', + 'efc__lo_next', + 'efc__lo_next_alpha', + 'efc__ls_done', + 'efc__margin', + 'efc__mid', + 'efc__mid_alpha', + 'efc__mv', + 'efc__p0', + 'efc__pos', + 'efc__prev_Mgrad', + 'efc__prev_cost', + 'efc__prev_grad', + 'efc__quad', + 'efc__quad_gauss', + 'efc__search', + 'efc__search_dot', + 'efc__type', + 'efc__u', + 'efc__uu', + 'efc__uv', + 'efc__vel', + 'efc__vv', + }, + ) + out = jf( + d.qpos.shape[0], + m._impl.M_rowadr, + m._impl.M_rownnz, + m.actuator_acc0, + m.actuator_actadr, + m.actuator_actearly, + m.actuator_actlimited, + m.actuator_actnum, + m.actuator_actrange, + m._impl.actuator_affine_bias_gain, + m.actuator_biasprm, + m.actuator_biastype, + m.actuator_cranklength, + m.actuator_ctrllimited, + m.actuator_ctrlrange, + m.actuator_dynprm, + m.actuator_dyntype, + m.actuator_forcelimited, + m.actuator_forcerange, + m.actuator_gainprm, + m.actuator_gaintype, + m.actuator_gear, + m.actuator_lengthrange, + m.actuator_trnid, + m.actuator_trntype, + m._impl.actuator_trntype_body_adr, + m._impl.block_dim, + m.body_dofadr, + m.body_dofnum, + m.body_gravcomp, + m.body_inertia, + m.body_invweight0, + m.body_ipos, + m.body_iquat, + m.body_jntadr, + m.body_jntnum, + m.body_mass, + m.body_parentid, + m.body_pos, + m.body_quat, + m.body_rootid, + m.body_subtreemass, + m._impl.body_tree, + m.body_weldid, + m.cam_bodyid, + m.cam_fovy, + m.cam_intrinsic, + m.cam_mat0, + m.cam_mode, + m.cam_pos, + m.cam_pos0, + m.cam_poscom0, + m.cam_quat, + m.cam_resolution, + m.cam_sensorsize, + m.cam_targetbodyid, + m._impl.condim_max, + m.dof_Madr, + m.dof_armature, + m.dof_bodyid, + m.dof_damping, + m.dof_frictionloss, + m.dof_invweight0, + m.dof_jntid, + m.dof_parentid, + m.dof_solimp, + m.dof_solref, + m._impl.dof_tri_col, + m._impl.dof_tri_row, + m._impl.eq_connect_adr, + m.eq_data, + m._impl.eq_jnt_adr, + m.eq_obj1id, + m.eq_obj2id, + m.eq_objtype, + m.eq_solimp, + m.eq_solref, + m._impl.eq_ten_adr, + m._impl.eq_wld_adr, + m._impl.flex_bending, + m._impl.flex_damping, + m._impl.flex_dim, + m._impl.flex_edge, + m._impl.flex_edgeadr, + m._impl.flex_edgeflap, + m._impl.flex_elem, + m._impl.flex_elemedge, + m._impl.flex_elemedgeadr, + m._impl.flex_stiffness, + m._impl.flex_vertadr, + m._impl.flex_vertbodyid, + m._impl.flexedge_length0, + m.geom_aabb, + m.geom_bodyid, + m.geom_condim, + m.geom_dataid, + m.geom_friction, + m.geom_gap, + m.geom_group, + m.geom_margin, + m.geom_matid, + m._impl.geom_pair_type_count, + m._impl.geom_plugin_index, + m.geom_pos, + m.geom_priority, + m.geom_quat, + m.geom_rbound, + m.geom_rgba, + m.geom_size, + m.geom_solimp, + m.geom_solmix, + m.geom_solref, + m.geom_type, + m._impl.geompair2hfgeompair, + m._impl.has_sdf_geom, + m.hfield_adr, + m.hfield_data, + m.hfield_ncol, + m.hfield_nrow, + m.hfield_size, + m.jnt_actfrclimited, + m.jnt_actfrcrange, + m.jnt_actgravcomp, + m.jnt_axis, + m.jnt_bodyid, + m.jnt_dofadr, + m._impl.jnt_limited_ball_adr, + m._impl.jnt_limited_slide_hinge_adr, + m.jnt_margin, + m.jnt_pos, + m.jnt_qposadr, + m.jnt_range, + m.jnt_solimp, + m.jnt_solref, + m.jnt_stiffness, + m.jnt_type, + m._impl.light_bodyid, + m.light_dir, + m.light_dir0, + m.light_mode, + m.light_pos, + m.light_pos0, + m.light_poscom0, + m._impl.light_targetbodyid, + m._impl.mapM2M, + m.mat_rgba, + m.mesh_face, + m.mesh_faceadr, + m.mesh_graph, + m.mesh_graphadr, + m._impl.mesh_polyadr, + m._impl.mesh_polymap, + m._impl.mesh_polymapadr, + m._impl.mesh_polymapnum, + m._impl.mesh_polynormal, + m._impl.mesh_polynum, + m._impl.mesh_polyvert, + m._impl.mesh_polyvertadr, + m._impl.mesh_polyvertnum, + m.mesh_vert, + m.mesh_vertadr, + m.mesh_vertnum, + m._impl.mocap_bodyid, + m.nC, + m.na, + m.nbody, + m.ncam, + m.neq, + m._impl.nflexedge, + m._impl.nflexelem, + m._impl.nflexvert, + m.ngeom, + m.ngravcomp, + m.nhfield, + m.njnt, + m.nlight, + m._impl.nlsp, + m.nmeshface, + m.nmocap, + m.nsite, + m.ntendon, + m.nu, + m.nv, + m._impl.nxn_geom_pair_filtered, + m._impl.nxn_pairid, + m._impl.nxn_pairid_filtered, + m.pair_dim, + m.pair_friction, + m.pair_gap, + m.pair_margin, + m.pair_solimp, + m.pair_solref, + m.pair_solreffriction, + m._impl.plugin, + m._impl.plugin_attr, + m._impl.qLD_updates, + m._impl.qM_fullm_i, + m._impl.qM_fullm_j, + m._impl.qM_madr_ij, + m._impl.qM_mulm_i, + m._impl.qM_mulm_j, + m._impl.qM_tiles, + m.qpos0, + m.qpos_spring, + m._impl.rangefinder_sensor_adr, + m._impl.sensor_acc_adr, + m.sensor_adr, + m.sensor_cutoff, + m.sensor_datatype, + m._impl.sensor_e_kinetic, + m._impl.sensor_e_potential, + m._impl.sensor_limitfrc_adr, + m._impl.sensor_limitpos_adr, + m._impl.sensor_limitvel_adr, + m.sensor_objid, + m.sensor_objtype, + m._impl.sensor_pos_adr, + m._impl.sensor_rangefinder_adr, + m._impl.sensor_rangefinder_bodyid, + m.sensor_refid, + m.sensor_reftype, + m._impl.sensor_rne_postconstraint, + m._impl.sensor_subtree_vel, + m._impl.sensor_tendonactfrc_adr, + m._impl.sensor_touch_adr, + m.sensor_type, + m._impl.sensor_vel_adr, + m.site_bodyid, + m.site_pos, + m.site_quat, + m.site_size, + m.site_type, + m._impl.subtree_mass, + m.tendon_actfrclimited, + m.tendon_actfrcrange, + m.tendon_adr, + m.tendon_armature, + m.tendon_damping, + m.tendon_frictionloss, + m._impl.tendon_geom_adr, + m.tendon_invweight0, + m._impl.tendon_jnt_adr, + m.tendon_length0, + m.tendon_lengthspring, + m._impl.tendon_limited_adr, + m.tendon_margin, + m.tendon_num, + m.tendon_range, + m._impl.tendon_site_pair_adr, + m.tendon_solimp_fri, + m.tendon_solimp_lim, + m.tendon_solref_fri, + m.tendon_solref_lim, + m.tendon_stiffness, + m._impl.wrap_geom_adr, + m._impl.wrap_jnt_adr, + m.wrap_objid, + m.wrap_prm, + m._impl.wrap_pulley_scale, + m._impl.wrap_site_pair_adr, + m.wrap_type, + m.opt._impl.broadphase, + m.opt._impl.broadphase_filter, + m.opt.cone, + m.opt.density, + m.opt.disableflags, + m.opt.enableflags, + m.opt._impl.epa_iterations, + m.opt._impl.gjk_iterations, + m.opt._impl.graph_conditional, + m.opt.gravity, + m.opt._impl.has_fluid, + m.opt.impratio, + m.opt.integrator, + m.opt._impl.is_sparse, + m.opt.iterations, + m.opt.ls_iterations, + m.opt._impl.ls_parallel, + m.opt.ls_tolerance, + m.opt.magnetic, + m.opt._impl.run_collision_detection, + m.opt._impl.sdf_initpoints, + m.opt._impl.sdf_iterations, + m.opt.solver, + m.opt.timestep, + m.opt.tolerance, + m.opt.viscosity, + m.opt.wind, + m.stat.meaninertia, + d._impl.nconmax, + d._impl.njmax, + d.act, + d.act_dot, + d._impl.act_dot_rk, + d._impl.act_t0, + d.actuator_force, + d._impl.actuator_length, + d._impl.actuator_moment, + d._impl.actuator_trntype_body_ncon, + d._impl.actuator_velocity, + d._impl.cacc, + d.cam_xmat, + d.cam_xpos, + d._impl.cdof, + d._impl.cdof_dot, + d._impl.cfrc_ext, + d._impl.cfrc_int, + d._impl.cinert, + d._impl.collision_hftri_index, + d._impl.collision_pair, + d._impl.collision_pairid, + d._impl.collision_worldid, + d._impl.crb, + d.ctrl, + d.cvel, + d._impl.energy, + d._impl.epa_face, + d._impl.epa_horizon, + d._impl.epa_index, + d._impl.epa_map, + d._impl.epa_norm2, + d._impl.epa_pr, + d._impl.epa_vert, + d._impl.epa_vert1, + d._impl.epa_vert2, + d._impl.epa_vert_index1, + d._impl.epa_vert_index2, + d.eq_active, + d._impl.flexedge_length, + d._impl.flexedge_velocity, + d._impl.flexvert_xpos, + d._impl.fluid_applied, + d._impl.geom_skip, + d.geom_xmat, + d.geom_xpos, + d._impl.inverse_mul_m_skip, + d._impl.light_xdir, + d._impl.light_xpos, + d.mocap_pos, + d.mocap_quat, + d._impl.ncollision, + d._impl.ncon, + d._impl.ncon_hfield, + d._impl.ne, + d._impl.ne_connect, + d._impl.ne_jnt, + d._impl.ne_ten, + d._impl.ne_weld, + d._impl.nefc, + d._impl.nf, + d._impl.nl, + d._impl.nsolving, + d._impl.qLD, + d._impl.qLD_integration, + d._impl.qLDiagInv, + d._impl.qLDiagInv_integration, + d._impl.qM, + d._impl.qM_integration, + d.qacc, + d._impl.qacc_integration, + d._impl.qacc_rk, + d.qacc_smooth, + d.qacc_warmstart, + d.qfrc_actuator, + d.qfrc_applied, + d.qfrc_bias, + d.qfrc_constraint, + d._impl.qfrc_damper, + d.qfrc_fluid, + d.qfrc_gravcomp, + d._impl.qfrc_integration, + d.qfrc_passive, + d.qfrc_smooth, + d._impl.qfrc_spring, + d.qpos, + d._impl.qpos_t0, + d.qvel, + d._impl.qvel_rk, + d._impl.qvel_t0, + d._impl.sap_cumulative_sum, + d._impl.sap_projection_lower, + d._impl.sap_projection_upper, + d._impl.sap_range, + d._impl.sap_segment_index, + d._impl.sap_sort_index, + d._impl.sensor_rangefinder_dist, + d._impl.sensor_rangefinder_geomid, + d._impl.sensor_rangefinder_pnt, + d._impl.sensor_rangefinder_vec, + d.sensordata, + d.site_xmat, + d.site_xpos, + d._impl.solver_niter, + d._impl.subtree_angmom, + d._impl.subtree_bodyvel, + d.subtree_com, + d._impl.subtree_linvel, + d._impl.ten_J, + d._impl.ten_Jdot, + d._impl.ten_actfrc, + d._impl.ten_bias_coef, + d._impl.ten_length, + d._impl.ten_velocity, + d._impl.ten_wrapadr, + d._impl.ten_wrapnum, + d.time, + d._impl.wrap_geom_xpos, + d._impl.wrap_obj, + d._impl.wrap_xpos, + d.xanchor, + d.xaxis, + d.xfrc_applied, + d.ximat, + d.xipos, + d.xmat, + d.xpos, + d.xquat, + d._impl.contact__dim, + d._impl.contact__dist, + d._impl.contact__efc_address, + d._impl.contact__frame, + d._impl.contact__friction, + d._impl.contact__geom, + d._impl.contact__includemargin, + d._impl.contact__pos, + d._impl.contact__solimp, + d._impl.contact__solref, + d._impl.contact__solreffriction, + d._impl.contact__worldid, + d._impl.efc__D, + d._impl.efc__J, + d._impl.efc__Jaref, + d._impl.efc__Ma, + d._impl.efc__Mgrad, + d._impl.efc__active, + d._impl.efc__alpha, + d._impl.efc__aref, + d._impl.efc__beta, + d._impl.efc__beta_den, + d._impl.efc__beta_num, + d._impl.efc__cholesky_L_tmp, + d._impl.efc__cholesky_y_tmp, + d._impl.efc__condim, + d._impl.efc__cost, + d._impl.efc__cost_candidate, + d._impl.efc__done, + d._impl.efc__force, + d._impl.efc__frictionloss, + d._impl.efc__gauss, + d._impl.efc__grad, + d._impl.efc__grad_dot, + d._impl.efc__gtol, + d._impl.efc__h, + d._impl.efc__hi, + d._impl.efc__hi_alpha, + d._impl.efc__hi_next, + d._impl.efc__hi_next_alpha, + d._impl.efc__id, + d._impl.efc__jv, + d._impl.efc__lo, + d._impl.efc__lo_alpha, + d._impl.efc__lo_next, + d._impl.efc__lo_next_alpha, + d._impl.efc__ls_done, + d._impl.efc__margin, + d._impl.efc__mid, + d._impl.efc__mid_alpha, + d._impl.efc__mv, + d._impl.efc__p0, + d._impl.efc__pos, + d._impl.efc__prev_Mgrad, + d._impl.efc__prev_cost, + d._impl.efc__prev_grad, + d._impl.efc__quad, + d._impl.efc__quad_gauss, + d._impl.efc__search, + d._impl.efc__search_dot, + d._impl.efc__type, + d._impl.efc__u, + d._impl.efc__uu, + d._impl.efc__uv, + d._impl.efc__vel, + d._impl.efc__vv, + ) + d = d.tree_replace({ + 'act': out[0], + 'act_dot': out[1], + '_impl.act_dot_rk': out[2], + '_impl.act_t0': out[3], + 'actuator_force': out[4], + '_impl.actuator_length': out[5], + '_impl.actuator_moment': out[6], + '_impl.actuator_trntype_body_ncon': out[7], + '_impl.actuator_velocity': out[8], + '_impl.cacc': out[9], + 'cam_xmat': out[10], + 'cam_xpos': out[11], + '_impl.cdof': out[12], + '_impl.cdof_dot': out[13], + '_impl.cfrc_ext': out[14], + '_impl.cfrc_int': out[15], + '_impl.cinert': out[16], + '_impl.collision_hftri_index': out[17], + '_impl.collision_pair': out[18], + '_impl.collision_pairid': out[19], + '_impl.collision_worldid': out[20], + '_impl.crb': out[21], + 'ctrl': out[22], + 'cvel': out[23], + '_impl.energy': out[24], + '_impl.epa_face': out[25], + '_impl.epa_horizon': out[26], + '_impl.epa_index': out[27], + '_impl.epa_map': out[28], + '_impl.epa_norm2': out[29], + '_impl.epa_pr': out[30], + '_impl.epa_vert': out[31], + '_impl.epa_vert1': out[32], + '_impl.epa_vert2': out[33], + '_impl.epa_vert_index1': out[34], + '_impl.epa_vert_index2': out[35], + 'eq_active': out[36], + '_impl.flexedge_length': out[37], + '_impl.flexedge_velocity': out[38], + '_impl.flexvert_xpos': out[39], + '_impl.fluid_applied': out[40], + '_impl.geom_skip': out[41], + 'geom_xmat': out[42], + 'geom_xpos': out[43], + '_impl.inverse_mul_m_skip': out[44], + '_impl.light_xdir': out[45], + '_impl.light_xpos': out[46], + 'mocap_pos': out[47], + 'mocap_quat': out[48], + '_impl.ncollision': out[49], + '_impl.ncon': out[50], + '_impl.ncon_hfield': out[51], + '_impl.ne': out[52], + '_impl.ne_connect': out[53], + '_impl.ne_jnt': out[54], + '_impl.ne_ten': out[55], + '_impl.ne_weld': out[56], + '_impl.nefc': out[57], + '_impl.nf': out[58], + '_impl.nl': out[59], + '_impl.nsolving': out[60], + '_impl.qLD': out[61], + '_impl.qLD_integration': out[62], + '_impl.qLDiagInv': out[63], + '_impl.qLDiagInv_integration': out[64], + '_impl.qM': out[65], + '_impl.qM_integration': out[66], + 'qacc': out[67], + '_impl.qacc_integration': out[68], + '_impl.qacc_rk': out[69], + 'qacc_smooth': out[70], + 'qacc_warmstart': out[71], + 'qfrc_actuator': out[72], + 'qfrc_applied': out[73], + 'qfrc_bias': out[74], + 'qfrc_constraint': out[75], + '_impl.qfrc_damper': out[76], + 'qfrc_fluid': out[77], + 'qfrc_gravcomp': out[78], + '_impl.qfrc_integration': out[79], + 'qfrc_passive': out[80], + 'qfrc_smooth': out[81], + '_impl.qfrc_spring': out[82], + 'qpos': out[83], + '_impl.qpos_t0': out[84], + 'qvel': out[85], + '_impl.qvel_rk': out[86], + '_impl.qvel_t0': out[87], + '_impl.sap_cumulative_sum': out[88], + '_impl.sap_projection_lower': out[89], + '_impl.sap_projection_upper': out[90], + '_impl.sap_range': out[91], + '_impl.sap_segment_index': out[92], + '_impl.sap_sort_index': out[93], + '_impl.sensor_rangefinder_dist': out[94], + '_impl.sensor_rangefinder_geomid': out[95], + '_impl.sensor_rangefinder_pnt': out[96], + '_impl.sensor_rangefinder_vec': out[97], + 'sensordata': out[98], + 'site_xmat': out[99], + 'site_xpos': out[100], + '_impl.solver_niter': out[101], + '_impl.subtree_angmom': out[102], + '_impl.subtree_bodyvel': out[103], + 'subtree_com': out[104], + '_impl.subtree_linvel': out[105], + '_impl.ten_J': out[106], + '_impl.ten_Jdot': out[107], + '_impl.ten_actfrc': out[108], + '_impl.ten_bias_coef': out[109], + '_impl.ten_length': out[110], + '_impl.ten_velocity': out[111], + '_impl.ten_wrapadr': out[112], + '_impl.ten_wrapnum': out[113], + 'time': out[114], + '_impl.wrap_geom_xpos': out[115], + '_impl.wrap_obj': out[116], + '_impl.wrap_xpos': out[117], + 'xanchor': out[118], + 'xaxis': out[119], + 'xfrc_applied': out[120], + 'ximat': out[121], + 'xipos': out[122], + 'xmat': out[123], + 'xpos': out[124], + 'xquat': out[125], + '_impl.contact__dim': out[126], + '_impl.contact__dist': out[127], + '_impl.contact__efc_address': out[128], + '_impl.contact__frame': out[129], + '_impl.contact__friction': out[130], + '_impl.contact__geom': out[131], + '_impl.contact__includemargin': out[132], + '_impl.contact__pos': out[133], + '_impl.contact__solimp': out[134], + '_impl.contact__solref': out[135], + '_impl.contact__solreffriction': out[136], + '_impl.contact__worldid': out[137], + '_impl.efc__D': out[138], + '_impl.efc__J': out[139], + '_impl.efc__Jaref': out[140], + '_impl.efc__Ma': out[141], + '_impl.efc__Mgrad': out[142], + '_impl.efc__active': out[143], + '_impl.efc__alpha': out[144], + '_impl.efc__aref': out[145], + '_impl.efc__beta': out[146], + '_impl.efc__beta_den': out[147], + '_impl.efc__beta_num': out[148], + '_impl.efc__cholesky_L_tmp': out[149], + '_impl.efc__cholesky_y_tmp': out[150], + '_impl.efc__condim': out[151], + '_impl.efc__cost': out[152], + '_impl.efc__cost_candidate': out[153], + '_impl.efc__done': out[154], + '_impl.efc__force': out[155], + '_impl.efc__frictionloss': out[156], + '_impl.efc__gauss': out[157], + '_impl.efc__grad': out[158], + '_impl.efc__grad_dot': out[159], + '_impl.efc__gtol': out[160], + '_impl.efc__h': out[161], + '_impl.efc__hi': out[162], + '_impl.efc__hi_alpha': out[163], + '_impl.efc__hi_next': out[164], + '_impl.efc__hi_next_alpha': out[165], + '_impl.efc__id': out[166], + '_impl.efc__jv': out[167], + '_impl.efc__lo': out[168], + '_impl.efc__lo_alpha': out[169], + '_impl.efc__lo_next': out[170], + '_impl.efc__lo_next_alpha': out[171], + '_impl.efc__ls_done': out[172], + '_impl.efc__margin': out[173], + '_impl.efc__mid': out[174], + '_impl.efc__mid_alpha': out[175], + '_impl.efc__mv': out[176], + '_impl.efc__p0': out[177], + '_impl.efc__pos': out[178], + '_impl.efc__prev_Mgrad': out[179], + '_impl.efc__prev_cost': out[180], + '_impl.efc__prev_grad': out[181], + '_impl.efc__quad': out[182], + '_impl.efc__quad_gauss': out[183], + '_impl.efc__search': out[184], + '_impl.efc__search_dot': out[185], + '_impl.efc__type': out[186], + '_impl.efc__u': out[187], + '_impl.efc__uu': out[188], + '_impl.efc__uv': out[189], + '_impl.efc__vel': out[190], + '_impl.efc__vv': out[191], + }) + return d + + +@jax.custom_batching.custom_vmap +@ffi.marshal_jax_warp_callable +def step(m: types.Model, d: types.Data): + return _step_jax_impl(m, d) + + +@step.def_vmap +@ffi.marshal_custom_vmap +def step_vmap(unused_axis_size, is_batched, m, d): + d = step(m, d) + return d, is_batched[1] diff --git a/mjx/mujoco/mjx/warp/forward_test.py b/mjx/mujoco/mjx/warp/forward_test.py new file mode 100644 index 00000000..c8e0ddef --- /dev/null +++ b/mjx/mujoco/mjx/warp/forward_test.py @@ -0,0 +1,274 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Tests for forward functions.""" + +import functools +import logging + +from absl.testing import absltest +from absl.testing import parameterized +import jax +import jax.numpy as jp +import mujoco +from mujoco import mjx +from mujoco.mjx._src import io +from mujoco.mjx._src import test_util +import mujoco.mjx.warp as mjxw +from mujoco.mjx.warp import test_util as tu +from mujoco.mjx.warp import warp as wp # pylint: disable=g-importing-member +import numpy as np + +try: + from mujoco.mjx.warp import forward # pylint: disable=g-import-not-at-top +except ImportError: + forward = None + + +class ForwardTest(parameterized.TestCase): + + def setUp(self): + super().setUp() + if mjxw.WARP_INSTALLED: + wp.clear_kernel_cache() + np.random.seed(0) + + @parameterized.parameters( + 'pendula.xml', + 'humanoid/humanoid.xml', + ) + def test_jit_caching(self, xml): + """Tests jit caching on the full step function.""" + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp not installed.') + if not io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device available.') + + batch_size = 7 + m = test_util.load_test_file(xml) + mx = mjx.put_model(m, impl='warp') + + keys = jp.arange(batch_size) + dx_batch = jax.vmap(functools.partial(tu.make_data, m))(keys) + + step_fn = jax.jit(jax.vmap(forward.step, in_axes=(None, 0))) + + was_logging_compiles = jax.config.jax_log_compiles + jax_logger = logging.getLogger('jax') + was_propagating = jax_logger.propagate + jax.config.update('jax_log_compiles', True) + jax_logger.propagate = False # do not print to stdout for this test + with self.assertLogs('jax', level='INFO') as log: + dx_batch1 = step_fn(mx, dx_batch) + jax.tree_util.tree_map(lambda x: x.block_until_ready(), dx_batch1) + + # Re-generate data and run step_fn again to test jit caching. + keys = jp.arange(batch_size, batch_size * 2) + dx_batch = jax.vmap(functools.partial(tu.make_data, m))(keys) + dx_batch2 = step_fn(mx, dx_batch) + jax.tree_util.tree_map(lambda x: x.block_until_ready(), dx_batch2) + jax.config.update('jax_log_compiles', was_logging_compiles) + jax_logger.propagate = was_propagating + + compilation_logs = [ + r for r in log.records if 'Compiling jit(step)' in r.getMessage() + ] + self.assertLen( + compilation_logs, + 1, + msg=( + f'Expected 1 compilation, got {len(compilation_logs)} compilations.' + ), + ) + + @parameterized.product( + xml=( + 'humanoid/humanoid.xml', + 'pendula.xml', + ), + batch_size=(1, 7), + ) + def test_forward(self, xml: str, batch_size: int): + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp not installed.') + if not io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device available.') + + m = test_util.load_test_file(xml) + m.opt.iterations = 10 + m.opt.ls_iterations = 10 + mx = mjx.put_model(m, impl='warp') + + d = mujoco.MjData(m) + worldids = jp.arange(batch_size) + dx_batch = jax.vmap(functools.partial(tu.make_data, m))(worldids) + + dx_batch = jax.jit(jax.vmap(forward.forward, in_axes=(None, 0)))( + mx, dx_batch + ) + + for i in range(batch_size): + dx = dx_batch[i] + + d.qpos[:] = dx.qpos + d.qvel[:] = dx.qvel + d.ctrl[:] = dx.ctrl + d.mocap_pos[:] = dx.mocap_pos + d.mocap_quat[:] = dx.mocap_quat + mujoco.mj_forward(m, d) + + # fwd_position + tu.assert_attr_eq(dx, d, 'xpos') + tu.assert_attr_eq(dx, d, 'xquat') + tu.assert_attr_eq(dx, d, 'xipos') + tu.assert_eq(d.ximat.reshape((-1, 3, 3)), dx.ximat, 'ximat') + tu.assert_attr_eq(dx, d, 'xanchor') + tu.assert_attr_eq(dx, d, 'xaxis') + tu.assert_attr_eq(dx, d, 'geom_xpos') + tu.assert_eq(dx.geom_xmat, d.geom_xmat.reshape((-1, 3, 3)), 'geom_xmat') + if m.nsite: + tu.assert_attr_eq(dx, d, 'site_xpos') + tu.assert_eq(dx.site_xmat, d.site_xmat.reshape((-1, 3, 3)), 'site_xmat') + tu.assert_attr_eq(dx._impl, d, 'cdof') + tu.assert_attr_eq(dx._impl, d, 'cinert') + tu.assert_attr_eq(dx, d, 'subtree_com') + if m.nlight: + tu.assert_attr_eq(dx._impl, d, 'light_xpos') + tu.assert_attr_eq(dx._impl, d, 'light_xdir') + if m.ncam: + tu.assert_attr_eq(dx, d, 'cam_xpos') + tu.assert_eq(dx.cam_xmat, d.cam_xmat.reshape((-1, 3, 3)), 'cam_xmat') + tu.assert_attr_eq(dx._impl, d, 'ten_length') + tu.assert_attr_eq(dx._impl, d, 'ten_J') + tu.assert_attr_eq(dx._impl, d, 'ten_wrapadr') + tu.assert_attr_eq(dx._impl, d, 'ten_wrapnum') + tu.assert_attr_eq(dx._impl, d, 'wrap_xpos') + tu.assert_attr_eq(dx._impl, d, 'wrap_obj') + tu.assert_attr_eq(dx._impl, d, 'crb') + + if not mx.opt._impl.is_sparse: + qm = np.zeros((m.nv, m.nv)) + mujoco.mj_fullM(m, qm, d.qM) + else: + qm = d.qM + tu.assert_eq(qm, dx._impl.qM, 'qM') + # qLD is fused in a cholesky factorize and solve, and not written to. + + tu.assert_contact_eq(d, dx, worldid=i) + + tu.assert_attr_eq(dx._impl, d, 'actuator_length') + actuator_moment = np.zeros((m.nu, m.nv)) + mujoco.mju_sparse2dense( + actuator_moment, + d.actuator_moment, + d.moment_rownnz, + d.moment_rowadr, + d.moment_colind, + ) + tu.assert_eq(dx._impl.actuator_moment, actuator_moment, 'actuator_moment') + + # fwd_velocity + tu.assert_attr_eq(dx._impl, d, 'actuator_velocity') + tu.assert_attr_eq(dx, d, 'cvel') + tu.assert_attr_eq(dx._impl, d, 'cdof_dot') + tu.assert_attr_eq(dx._impl, d, 'qfrc_spring') + tu.assert_attr_eq(dx._impl, d, 'qfrc_damper') + tu.assert_attr_eq(dx, d, 'qfrc_gravcomp') + tu.assert_attr_eq(dx, d, 'qfrc_fluid') + tu.assert_attr_eq(dx, d, 'qfrc_passive') + tu.assert_attr_eq(dx, d, 'qfrc_bias') + tu.assert_efc_eq(d, dx, worldid=i) + + # fwd_actuation + tu.assert_attr_eq(dx, d, 'act_dot') + tu.assert_attr_eq(dx, d, 'actuator_force') + tu.assert_attr_eq(dx, d, 'qfrc_actuator') + + # fwd_acceleration + tu.assert_attr_eq(dx, d, 'qfrc_smooth') + tu.assert_attr_eq(dx, d, 'qacc_smooth') + + # solve + np.testing.assert_allclose( + dx.qacc_warmstart, + d.qacc_warmstart, + err_msg='qacc_warmstart', + rtol=1e-5, + atol=1.0, + ) + np.testing.assert_allclose( + dx.qacc, d.qacc, err_msg='qacc', rtol=1e-5, atol=1.0 + ) + + +class StepTest(parameterized.TestCase): + + def setUp(self): + super().setUp() + if mjxw.WARP_INSTALLED: + wp.clear_kernel_cache() + np.random.seed(0) + + @parameterized.product( + xml=( + 'humanoid/humanoid.xml', + 'pendula.xml', + ), + batch_size=(1, 7), + ) + def test_step(self, xml: str, batch_size: int): + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp not installed.') + if not io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device available.') + + m = test_util.load_test_file(xml) + m.opt.iterations = 10 + m.opt.ls_iterations = 10 + mx = mjx.put_model(m, impl='warp') + + d = mujoco.MjData(m) + worldids = jp.arange(batch_size) + dx_batch = jax.vmap(functools.partial(tu.make_data, m))(worldids) + dx_batch_orig = dx_batch + + for _ in range(10): + dx_batch = jax.jit(jax.vmap(forward.step, in_axes=(None, 0)))( + mx, dx_batch + ) + + for i in range(batch_size): + dx = dx_batch[i] + dx_orig = dx_batch_orig[i] + + d.qpos[:] = dx_orig.qpos + d.qvel[:] = dx_orig.qvel + d.ctrl[:] = dx_orig.ctrl + d.mocap_pos[:] = dx_orig.mocap_pos + d.mocap_quat[:] = dx_orig.mocap_quat + d.time = dx_orig.time + mujoco.mj_step(m, d, 10) + + tu.assert_attr_eq(dx, d, 'qpos') + tu.assert_attr_eq(dx, d, 'qvel') + tu.assert_attr_eq(dx, d, 'time') + tu.assert_attr_eq(dx, d, 'ctrl') + tu.assert_attr_eq(dx, d, 'act') + tu.assert_attr_eq(dx, d, 'mocap_pos') + tu.assert_attr_eq(dx, d, 'mocap_quat') + tu.assert_attr_eq(dx, d, 'sensordata') + + +if __name__ == '__main__': + absltest.main() diff --git a/mjx/mujoco/mjx/warp/smooth.py b/mjx/mujoco/mjx/warp/smooth.py new file mode 100644 index 00000000..087d5aab --- /dev/null +++ b/mjx/mujoco/mjx/warp/smooth.py @@ -0,0 +1,291 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""DO NOT EDIT. This file is auto-generated.""" +import dataclasses +import jax +from mujoco.mjx._src import types +from mujoco.mjx.warp import ffi +import mujoco.mjx.third_party.mujoco_warp as mjwarp +from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types +import warp as wp + + +_m = mjwarp.Model( + **{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init} +) +_d = mjwarp.Data( + **{f.name: None for f in dataclasses.fields(mjwarp.Data) if f.init} +) +_o = mjwarp.Option( + **{f.name: None for f in dataclasses.fields(mjwarp.Option) if f.init} +) +_s = mjwarp.Statistic( + **{f.name: None for f in dataclasses.fields(mjwarp.Statistic) if f.init} +) +_c = mjwarp.Contact( + **{f.name: None for f in dataclasses.fields(mjwarp.Contact) if f.init} +) +_e = mjwarp.Constraint( + **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} +) + + +@ffi.format_args_for_warp +def _kinematics_shim( + # Model + nworld: int, + body_dofadr: wp.array(dtype=int), + body_ipos: wp.array2d(dtype=wp.vec3), + body_iquat: wp.array2d(dtype=wp.quat), + body_jntadr: wp.array(dtype=int), + body_jntnum: wp.array(dtype=int), + body_parentid: wp.array(dtype=int), + body_pos: wp.array2d(dtype=wp.vec3), + body_quat: wp.array2d(dtype=wp.quat), + body_tree: tuple[wp.array(dtype=int), ...], + flex_edge: wp.array(dtype=wp.vec2i), + flex_vertadr: wp.array(dtype=int), + flex_vertbodyid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + geom_pos: wp.array2d(dtype=wp.vec3), + geom_quat: wp.array2d(dtype=wp.quat), + jnt_axis: wp.array2d(dtype=wp.vec3), + jnt_pos: wp.array2d(dtype=wp.vec3), + jnt_qposadr: wp.array(dtype=int), + jnt_type: wp.array(dtype=int), + mocap_bodyid: wp.array(dtype=int), + nflexedge: int, + nflexvert: int, + ngeom: int, + nmocap: int, + nsite: int, + qpos0: wp.array2d(dtype=float), + site_bodyid: wp.array(dtype=int), + site_pos: wp.array2d(dtype=wp.vec3), + site_quat: wp.array2d(dtype=wp.quat), + # Data + flexedge_length: wp.array2d(dtype=float), + flexedge_velocity: wp.array2d(dtype=float), + flexvert_xpos: wp.array2d(dtype=wp.vec3), + geom_skip: wp.array(dtype=bool), + geom_xmat: wp.array2d(dtype=wp.mat33), + geom_xpos: wp.array2d(dtype=wp.vec3), + mocap_pos: wp.array2d(dtype=wp.vec3), + mocap_quat: wp.array2d(dtype=wp.quat), + qpos: wp.array2d(dtype=float), + qvel: wp.array2d(dtype=float), + site_xmat: wp.array2d(dtype=wp.mat33), + site_xpos: wp.array2d(dtype=wp.vec3), + xanchor: wp.array2d(dtype=wp.vec3), + xaxis: wp.array2d(dtype=wp.vec3), + ximat: wp.array2d(dtype=wp.mat33), + xipos: wp.array2d(dtype=wp.vec3), + xmat: wp.array2d(dtype=wp.mat33), + xpos: wp.array2d(dtype=wp.vec3), + xquat: wp.array2d(dtype=wp.quat), +): + _m.stat = _s + _m.opt = _o + _d.efc = _e + _d.contact = _c + _m.body_dofadr = body_dofadr + _m.body_ipos = body_ipos + _m.body_iquat = body_iquat + _m.body_jntadr = body_jntadr + _m.body_jntnum = body_jntnum + _m.body_parentid = body_parentid + _m.body_pos = body_pos + _m.body_quat = body_quat + _m.body_tree = body_tree + _m.flex_edge = flex_edge + _m.flex_vertadr = flex_vertadr + _m.flex_vertbodyid = flex_vertbodyid + _m.geom_bodyid = geom_bodyid + _m.geom_pos = geom_pos + _m.geom_quat = geom_quat + _m.jnt_axis = jnt_axis + _m.jnt_pos = jnt_pos + _m.jnt_qposadr = jnt_qposadr + _m.jnt_type = jnt_type + _m.mocap_bodyid = mocap_bodyid + _m.nflexedge = nflexedge + _m.nflexvert = nflexvert + _m.ngeom = ngeom + _m.nmocap = nmocap + _m.nsite = nsite + _m.qpos0 = qpos0 + _m.site_bodyid = site_bodyid + _m.site_pos = site_pos + _m.site_quat = site_quat + _d.flexedge_length = flexedge_length + _d.flexedge_velocity = flexedge_velocity + _d.flexvert_xpos = flexvert_xpos + _d.geom_skip = geom_skip + _d.geom_xmat = geom_xmat + _d.geom_xpos = geom_xpos + _d.mocap_pos = mocap_pos + _d.mocap_quat = mocap_quat + _d.qpos = qpos + _d.qvel = qvel + _d.site_xmat = site_xmat + _d.site_xpos = site_xpos + _d.xanchor = xanchor + _d.xaxis = xaxis + _d.ximat = ximat + _d.xipos = xipos + _d.xmat = xmat + _d.xpos = xpos + _d.xquat = xquat + _d.nworld = nworld + mjwarp.kinematics(_m, _d) + + +def _kinematics_jax_impl(m: types.Model, d: types.Data): + output_dims = { + 'flexedge_length': d._impl.flexedge_length.shape, + 'flexedge_velocity': d._impl.flexedge_velocity.shape, + 'flexvert_xpos': d._impl.flexvert_xpos.shape, + 'geom_skip': d._impl.geom_skip.shape, + 'geom_xmat': d.geom_xmat.shape, + 'geom_xpos': d.geom_xpos.shape, + 'mocap_pos': d.mocap_pos.shape, + 'mocap_quat': d.mocap_quat.shape, + 'qpos': d.qpos.shape, + 'qvel': d.qvel.shape, + 'site_xmat': d.site_xmat.shape, + 'site_xpos': d.site_xpos.shape, + 'xanchor': d.xanchor.shape, + 'xaxis': d.xaxis.shape, + 'ximat': d.ximat.shape, + 'xipos': d.xipos.shape, + 'xmat': d.xmat.shape, + 'xpos': d.xpos.shape, + 'xquat': d.xquat.shape, + } + jf = ffi.jax_callable_variadic_tuple( + _kinematics_shim, + num_outputs=19, + output_dims=output_dims, + vmap_method=None, + graph_compatible=True, + in_out_argnames={ + 'flexedge_length', + 'flexedge_velocity', + 'flexvert_xpos', + 'geom_skip', + 'geom_xmat', + 'geom_xpos', + 'mocap_pos', + 'mocap_quat', + 'qpos', + 'qvel', + 'site_xmat', + 'site_xpos', + 'xanchor', + 'xaxis', + 'ximat', + 'xipos', + 'xmat', + 'xpos', + 'xquat', + }, + ) + out = jf( + d.qpos.shape[0], + m.body_dofadr, + m.body_ipos, + m.body_iquat, + m.body_jntadr, + m.body_jntnum, + m.body_parentid, + m.body_pos, + m.body_quat, + m._impl.body_tree, + m._impl.flex_edge, + m._impl.flex_vertadr, + m._impl.flex_vertbodyid, + m.geom_bodyid, + m.geom_pos, + m.geom_quat, + m.jnt_axis, + m.jnt_pos, + m.jnt_qposadr, + m.jnt_type, + m._impl.mocap_bodyid, + m._impl.nflexedge, + m._impl.nflexvert, + m.ngeom, + m.nmocap, + m.nsite, + m.qpos0, + m.site_bodyid, + m.site_pos, + m.site_quat, + d._impl.flexedge_length, + d._impl.flexedge_velocity, + d._impl.flexvert_xpos, + d._impl.geom_skip, + d.geom_xmat, + d.geom_xpos, + d.mocap_pos, + d.mocap_quat, + d.qpos, + d.qvel, + d.site_xmat, + d.site_xpos, + d.xanchor, + d.xaxis, + d.ximat, + d.xipos, + d.xmat, + d.xpos, + d.xquat, + ) + d = d.tree_replace({ + '_impl.flexedge_length': out[0], + '_impl.flexedge_velocity': out[1], + '_impl.flexvert_xpos': out[2], + '_impl.geom_skip': out[3], + 'geom_xmat': out[4], + 'geom_xpos': out[5], + 'mocap_pos': out[6], + 'mocap_quat': out[7], + 'qpos': out[8], + 'qvel': out[9], + 'site_xmat': out[10], + 'site_xpos': out[11], + 'xanchor': out[12], + 'xaxis': out[13], + 'ximat': out[14], + 'xipos': out[15], + 'xmat': out[16], + 'xpos': out[17], + 'xquat': out[18], + }) + return d + + +@jax.custom_batching.custom_vmap +@ffi.marshal_jax_warp_callable +def kinematics(m: types.Model, d: types.Data): + return _kinematics_jax_impl(m, d) + + +@kinematics.def_vmap +@ffi.marshal_custom_vmap +def kinematics_vmap(unused_axis_size, is_batched, m, d): + d = kinematics(m, d) + return d, is_batched[1] diff --git a/mjx/mujoco/mjx/warp/smooth_test.py b/mjx/mujoco/mjx/warp/smooth_test.py new file mode 100644 index 00000000..7a87c24f --- /dev/null +++ b/mjx/mujoco/mjx/warp/smooth_test.py @@ -0,0 +1,229 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Tests for codegen'd smooth functions.""" + +import functools + +from absl.testing import absltest +import jax +from jax import numpy as jp +import mujoco +from mujoco import mjx +from mujoco.mjx._src import io +from mujoco.mjx._src import math +import mujoco.mjx.warp as mjxw +from mujoco.mjx.warp import test_util as tu +from mujoco.mjx.warp import warp as wp # pylint: disable=g-importing-member +import numpy as np + +try: + from mujoco.mjx.warp import smooth # pylint: disable=g-import-not-at-top +except ImportError: + smooth = None + + +class SmoothTest(absltest.TestCase): + + def setUp(self): + super().setUp() + if mjxw.WARP_INSTALLED: + wp.clear_kernel_cache() + np.random.seed(0) + + def test_kinematics(self): + """Tests kinematics with unbatched data.""" + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp not installed.') + if not io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device available.') + + m = tu.load_test_file('pendula.xml') + + d = mujoco.MjData(m) + mx = mjx.put_model(m, impl='warp') + + rng = jax.random.PRNGKey(0) + dx = mjx.make_data(m, impl='warp') + rng, key = jax.random.split(rng) + qpos = jax.random.uniform(key, (m.nq,)) + _, key1, key2 = jax.random.split(rng, 3) + mocap_pos = jax.random.normal(key1, (m.nmocap, 3)) + mocap_quat = jax.random.normal(key2, (m.nmocap, 4)) + mocap_quat = math.normalize(mocap_quat) + dx = dx.replace(qpos=qpos, mocap_pos=mocap_pos, mocap_quat=mocap_quat) + + dx = jax.jit(smooth.kinematics)(mx, dx) + + d.qpos[:] = qpos + d.mocap_pos[:] = mocap_pos + d.mocap_quat[:] = mocap_quat + mujoco.mj_forward(m, d) + + tu.assert_attr_eq(d, dx, 'xanchor') + tu.assert_attr_eq(d, dx, 'xaxis') + tu.assert_attr_eq(d, dx, 'xpos') + tu.assert_attr_eq(d, dx, 'xquat') + tu.assert_eq(d.xmat.reshape((-1, 3, 3)), dx.xmat, 'xmat') + tu.assert_attr_eq(d, dx, 'xipos') + tu.assert_eq(d.ximat.reshape((-1, 3, 3)), dx.ximat, 'ximat') + tu.assert_attr_eq(d, dx, 'geom_xpos') + tu.assert_eq(d.geom_xmat.reshape((-1, 3, 3)), dx.geom_xmat, 'geom_xmat') + tu.assert_attr_eq(d, dx, 'site_xpos') + tu.assert_eq(d.site_xmat.reshape((-1, 3, 3)), dx.site_xmat, 'site_xmat') + + def test_kinematics_vmap(self): + """Tests kinematics with batched data.""" + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp not installed.') + if not io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device available.') + + m = tu.load_test_file('pendula.xml') + + batch_size = 7 + d = mujoco.MjData(m) + mx = mjx.put_model(m, impl='warp') + + worldids = jp.arange(batch_size) + dx_batch = jax.vmap(functools.partial(tu.make_data, m))(worldids) + fields = ('xanchor', 'xaxis', 'xpos', 'xquat', 'xmat', 'xipos', 'ximat', + 'geom_xpos', 'geom_xmat', 'site_xpos', 'site_xmat') # fmt: skip + for f in fields: + dx_batch = dx_batch.replace(**{f: jp.zeros_like(getattr(dx_batch, f))}) + + dx_batch = jax.jit(jax.vmap(smooth.kinematics, in_axes=(None, 0)))( + mx, dx_batch + ) + + for i in range(batch_size): + dx = dx_batch[i] + + d.qpos[:] = dx.qpos + d.mocap_pos[:] = dx.mocap_pos + d.mocap_quat[:] = dx.mocap_quat + mujoco.mj_forward(m, d) + + tu.assert_attr_eq(d, dx, 'xanchor') + tu.assert_attr_eq(d, dx, 'xaxis') + tu.assert_attr_eq(d, dx, 'xpos') + tu.assert_attr_eq(d, dx, 'xquat') + tu.assert_eq(d.xmat.reshape((-1, 3, 3)), dx.xmat, 'xmat') + tu.assert_attr_eq(d, dx, 'xipos') + tu.assert_eq(d.ximat.reshape((-1, 3, 3)), dx.ximat, 'ximat') + tu.assert_attr_eq(d, dx, 'geom_xpos') + tu.assert_eq(d.geom_xmat.reshape((-1, 3, 3)), dx.geom_xmat, 'geom_xmat') + tu.assert_attr_eq(d, dx, 'site_xpos') + tu.assert_eq(d.site_xmat.reshape((-1, 3, 3)), dx.site_xmat, 'site_xmat') + + def test_kinematics_nested_vmap(self): + """Tests kinematics with nested batch data.""" + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp not installed.') + if not io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device available.') + + m = tu.load_test_file('pendula.xml') + + d = mujoco.MjData(m) + mx = mjx.put_model(m, impl='warp') + + worldids = jp.arange(16).reshape((4, 4)) + dx_batch = jax.vmap(jax.vmap(functools.partial(tu.make_data, m)))(worldids) + fields = ('xanchor', 'xaxis', 'xpos', 'xquat', 'xmat', 'xipos', 'ximat', + 'geom_xpos', 'geom_xmat', 'site_xpos', 'site_xmat') # fmt: skip + for f in fields: + dx_batch = dx_batch.replace(**{f: jp.zeros_like(getattr(dx_batch, f))}) + + dx_batch = jax.jit( + jax.vmap( + jax.vmap(smooth.kinematics, in_axes=(None, 0)), in_axes=(None, 0) + ) + )(mx, dx_batch) + + for i in range(4): + for j in range(4): + dx = dx_batch[i, j] + + d.qpos[:] = dx.qpos + d.mocap_pos[:] = dx.mocap_pos + d.mocap_quat[:] = dx.mocap_quat + mujoco.mj_forward(m, d) + + tu.assert_attr_eq(d, dx, 'xanchor') + tu.assert_attr_eq(d, dx, 'xaxis') + tu.assert_attr_eq(d, dx, 'xpos') + tu.assert_attr_eq(d, dx, 'xquat') + tu.assert_eq(d.xmat.reshape((-1, 3, 3)), dx.xmat, 'xmat') + tu.assert_attr_eq(d, dx, 'xipos') + tu.assert_eq(d.ximat.reshape((-1, 3, 3)), dx.ximat, 'ximat') + tu.assert_attr_eq(d, dx, 'geom_xpos') + tu.assert_eq(d.geom_xmat.reshape((-1, 3, 3)), dx.geom_xmat, 'geom_xmat') + tu.assert_attr_eq(d, dx, 'site_xpos') + tu.assert_eq(d.site_xmat.reshape((-1, 3, 3)), dx.site_xmat, 'site_xmat') + + def test_kinematics_model_vmap(self): + """Tests kinematics with vmap on model and data fields.""" + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp not installed.') + if not io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device available.') + + m = tu.load_test_file('pendula.xml') + + batch_size = 7 + d = mujoco.MjData(m) + mx = mjx.put_model(m, impl='warp') + # Add batch dimension to one model field. + mx = mx.replace( + geom_pos=jax.random.normal( + jax.random.PRNGKey(0), (batch_size, m.ngeom, 3) + ) + ) + + worldids = jp.arange(batch_size) + dx_batch = jax.vmap(functools.partial(tu.make_data, m))(worldids) + fields = ('xanchor', 'xaxis', 'xpos', 'xquat', 'xmat', 'xipos', 'ximat', + 'geom_xpos', 'geom_xmat', 'site_xpos', 'site_xmat') # fmt: skip + for f in fields: + dx_batch = dx_batch.replace(**{f: jp.zeros_like(getattr(dx_batch, f))}) + + dx_batch = jax.jit(jax.vmap(smooth.kinematics, in_axes=(None, 0)))( + mx, dx_batch + ) + + for i in range(batch_size): + dx = dx_batch[i] + + d.qpos[:] = dx.qpos + d.mocap_pos[:] = dx.mocap_pos + d.mocap_quat[:] = dx.mocap_quat + m.geom_pos[:] = mx.geom_pos[i] + mujoco.mj_forward(m, d) + + tu.assert_attr_eq(d, dx, 'xanchor') + tu.assert_attr_eq(d, dx, 'xaxis') + tu.assert_attr_eq(d, dx, 'xpos') + tu.assert_attr_eq(d, dx, 'xquat') + tu.assert_eq(d.xmat.reshape((-1, 3, 3)), dx.xmat, 'xmat') + tu.assert_attr_eq(d, dx, 'xipos') + tu.assert_eq(d.ximat.reshape((-1, 3, 3)), dx.ximat, 'ximat') + tu.assert_attr_eq(d, dx, 'geom_xpos') + tu.assert_eq(d.geom_xmat.reshape((-1, 3, 3)), dx.geom_xmat, 'geom_xmat') + tu.assert_attr_eq(d, dx, 'site_xpos') + tu.assert_eq(d.site_xmat.reshape((-1, 3, 3)), dx.site_xmat, 'site_xmat') + + +if __name__ == '__main__': + absltest.main() diff --git a/mjx/mujoco/mjx/warp/test_util.py b/mjx/mujoco/mjx/warp/test_util.py new file mode 100644 index 00000000..b5c25f3c --- /dev/null +++ b/mjx/mujoco/mjx/warp/test_util.py @@ -0,0 +1,204 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Utilities for testing MJX MjWarp integration.""" +import jax +import jax.numpy as jp +import mujoco +from mujoco import mjx +from mujoco.mjx._src import test_util as mjx_test_util +import numpy as np + +try: + from mujoco.mjx.warp import forward as mjxw_forward # pylint: disable=g-import-not-at-top +except ImportError: + mjxw_forward = None + +# tolerance for difference between MuJoCo and MJX smooth calculations - mostly +# due to float precision +_TOLERANCE = 5e-5 + + +def assert_eq(a, b, name): + tol = _TOLERANCE * 10 # avoid test noise + err_msg = f'mismatch: {name}' + np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol) + + +def assert_attr_eq(a, b, attr): + assert_eq(getattr(a, attr), getattr(b, attr), attr) + + +def make_data( + m: mujoco.MjModel, worldid: int, nconmax: int = 1_000, njmax: int = 100 +): + """Make data for a given worldid using keyframes when available.""" + dx = mjx.make_data(m, impl='warp', nconmax=nconmax, njmax=njmax) + + rng = jax.random.PRNGKey(worldid) + rng, key = jax.random.split(rng) + qpos = m.qpos0 + jax.random.uniform(key, (m.nq,), minval=-0.05, maxval=0.05) + key_qpos = ( + jp.array(m.key_qpos)[worldid] if m.nkey > 0 else jp.zeros_like(qpos) + ) + qpos = jp.where(worldid < m.nkey, key_qpos, qpos) + + rng, key = jax.random.split(rng) + qvel = jax.random.uniform(key, (m.nv,), minval=-0.05, maxval=0.05) + key_qvel = ( + jp.array(m.key_qvel)[worldid] if m.nkey > 0 else jp.zeros_like(qvel) + ) + qvel = jp.where(worldid < m.nkey, key_qvel, qvel) + + rng, key1, key2 = jax.random.split(rng, 3) + mocap_pos = jax.random.normal(key1, (m.nmocap, 3)) + key_mpos = ( + jp.array(m.key_mpos)[worldid] + if m.nkey > 0 and m.nmocap > 0 + else jp.zeros_like(mocap_pos) + ) + mocap_pos = jp.where(worldid < m.nkey, key_mpos, mocap_pos) + mocap_quat = jax.random.normal(key2, (m.nmocap, 4)) + key_mquat = ( + jp.array(m.key_mquat)[worldid] + if m.nkey > 0 and m.nmocap > 0 + else jp.zeros_like(mocap_quat) + ) + mocap_quat = jp.where(worldid < m.nkey, key_mquat, mocap_quat) + + _, key = jax.random.split(rng) + ctrl = jax.random.uniform(key, (m.nu,), minval=-0.1, maxval=0.1) + key_ctrl = ( + jp.array(m.key_ctrl)[worldid] if m.nkey > 0 else jp.zeros_like(ctrl) + ) + ctrl = jp.where(worldid < m.nkey, key_ctrl, ctrl) + + dx = dx.replace( + qpos=qpos, + qvel=qvel, + mocap_pos=mocap_pos, + mocap_quat=mocap_quat, + ctrl=ctrl, + ) + + # Mimic forward call within a vmap trace (since make_data gets called in a + # vmap in tests). Some fields get vmapped (qpos, qvel) while others need to be + # broadcasted by the custom vmap rule (e.g. geom_xpos). + mx = mjx.put_model(m, impl='warp') + dx = mjxw_forward.forward(mx, dx) + + return dx + + +def _mjx_contact(dx, worldid: int): + keys = [] + for i in range(dx._impl.ncon[0]): + if dx._impl.contact__worldid[i] != worldid: + continue + + g1, g2 = tuple(map(int, dx._impl.contact__geom[i])) + dist = float(dx._impl.contact__dist[i]) + keys.append((g1, g2, -dist, i)) + + keys = sorted(keys) + geom1 = np.array([k[0] for k in keys]) + geom2 = np.array([k[1] for k in keys]) + dist = np.array([dx._impl.contact__dist[k[-1]] for k in keys]) + normal = np.array([dx._impl.contact__frame[k[-1]][0] for k in keys]) + return geom1, geom2, dist, normal + + +def _mj_contact(d): + keys = [] + for i in range(d.ncon): + g1, g2 = tuple(map(int, d.contact.geom[i])) + dist = float(d.contact.dist[i]) + keys.append((g1, g2, -dist, i)) + + keys = sorted(keys) + geom1 = np.array([k[0] for k in keys]) + geom2 = np.array([k[1] for k in keys]) + dist = np.array([d.contact.dist[k[-1]] for k in keys]) + normal = np.array([d.contact.frame[k[-1]][:3] for k in keys]) + return geom1, geom2, dist, normal + + +def assert_contact_eq(d, dx, worldid: int): + *geom, dist, normal = _mj_contact(d) + *geomp, distp, normalp = _mjx_contact(dx, worldid) + assert_eq(geomp, geom, 'geom') + assert_eq(distp, dist, 'dist') + assert_eq(normalp, normal, 'normal') + + +def _mjx_efc(dx, worldid: int): + """Gets unpacked efc data for a given worldid.""" + select = lambda x: x if dx._impl.nefc.ndim == 0 else x[worldid] + nefc = select(dx._impl.nefc) + keys = np.arange(nefc) + if not keys.size: + empty = np.array([]) + return 0, empty, empty, np.zeros((0, dx.qvel.shape[0])), empty, empty + efc_pos = select(dx._impl.efc__pos[:nefc]) + efc_type = select(dx._impl.efc__type[:nefc]) + keys_sorted = np.lexsort((-efc_pos, efc_type)) + keys = keys[keys_sorted] + + nefc = len(keys) + type_ = efc_type[keys] + pos = efc_pos[keys] + j = select(dx._impl.efc__J[:nefc])[keys] + aref = select(dx._impl.efc__aref[:nefc])[keys] + d_ = select(dx._impl.efc__D[:nefc])[keys] + return nefc, type_, pos, j, aref, d_ + + +def _mj_efc(d): + """Gets unpacked efc data.""" + efc_j = np.zeros((d.efc_J_rownnz.shape[0], d.qvel.shape[0])) + if d.efc_J.shape[0] < efc_j.shape[0] * efc_j.shape[1]: + mujoco.mju_sparse2dense( + efc_j, + d.efc_J, + d.efc_J_rownnz, + d.efc_J_rowadr, + d.efc_J_colind, + ) + else: + efc_j = d.efc_J.reshape((-1, d.qvel.shape[0])) + + keys = np.lexsort((-d.efc_pos, d.efc_type)) + type_ = d.efc_type[keys] + pos = d.efc_pos[keys] + efc_j = efc_j[keys] + aref = d.efc_aref[keys] + d_ = d.efc_D[keys] + return d.nefc, type_, pos, efc_j, aref, d_ + + +def assert_efc_eq(d, dx, worldid: int): + nefc, type_, pos, j, aref, d_ = _mj_efc(d) + nefcp, typep, posp, jp_, arefp, dp = _mjx_efc(dx, worldid) + + assert_eq(nefcp, nefc, 'nefc') + assert_eq(typep, type_, 'type') + assert_eq(posp, pos, 'pos') + assert_eq(jp_, j, 'J') + assert_eq(arefp, aref, 'aref') + assert_eq(dp, d_, 'D') + + +def load_test_file(name: str) -> mujoco.MjModel: + """Loads a mujoco.MjModel based on the file name.""" + return mjx_test_util.load_test_file(name) diff --git a/mjx/mujoco/mjx/warp/testspeed.py b/mjx/mujoco/mjx/warp/testspeed.py new file mode 100644 index 00000000..24802a07 --- /dev/null +++ b/mjx/mujoco/mjx/warp/testspeed.py @@ -0,0 +1,427 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Run benchmarks.""" + +import functools +import os +import time +from typing import Any, Callable, Sequence, Tuple + +from absl import app +from absl import flags +import jax +import mujoco +from mujoco import mjx +from mujoco.mjx._src import test_util +from mujoco.mjx.warp import collision_driver as wp_collision +from mujoco.mjx.warp import forward as wp_forward +from mujoco.mjx.warp import smooth as wp_smooth +import mujoco.mjx.third_party.mujoco_warp as mjwarp +import warp as wp +from mujoco.mjx.third_party.warp.jax_experimental import ffi as warp_ffi + +_MODELFILE = flags.DEFINE_string( + 'modelfile', + 'humanoid/humanoid.xml', + 'path to model', +) +_FUNCTION = flags.DEFINE_string( + 'function', 'kinematics', 'function to benchmark' +) +_NSTEP = flags.DEFINE_integer('nstep', 1000, 'number of steps per rollout') +_NENV = flags.DEFINE_integer('nenv', 8192, 'number of environments to simulate') +_UNROLL = flags.DEFINE_integer('unroll', 4, 'number of steps to unroll') +_NCONMAX = flags.DEFINE_integer('nconmax', 30_000, 'max contacts') +_NJMAX = flags.DEFINE_integer('njmax', 80_000, 'max constraints') +_WP_KERNEL_CACHE_DIR = flags.DEFINE_string( + 'wp_kernel_cache_dir', + None, + 'Path to the Warp kernel cache directory.', +) +_COMPILER_OPTIONS = {'xla_gpu_graph_min_graph_size': 1} +jax_jit = functools.partial(jax.jit, compiler_options=_COMPILER_OPTIONS) + + +def _measure(fn, *args) -> Tuple[float, float]: + """Reports jit time and op time for a function.""" + + beg = time.perf_counter() + compiled_fn = fn.lower(*args).compile() + end = time.perf_counter() + jit_time = end - beg + + # warmup + result = compiled_fn(*args) + jax.block_until_ready(result) + + times = [] + + for i in range(1): + beg = time.perf_counter() + result = compiled_fn(*args) + jax.block_until_ready(result) + end = time.perf_counter() + run_time = end - beg + times.append(run_time) + print('Measure run: ', i, f', run time: {run_time:.3f}') + + return jit_time, sum(times) / len(times) + + +def benchmark( + m: mujoco.MjModel, + mx: mjx.Model, + step_fn: Callable[..., Any], + nstep: int = 1000, + nenv: int = 8192, + unroll_steps: int = 4, +) -> Tuple[float, float, int]: + """Benchmark a model.""" + + @jax.vmap + def init(key): + d = mjx.make_data( + m, impl=mx.impl, nconmax=_NCONMAX.value, njmax=_NJMAX.value + ) + return d + + key = jax.random.split(jax.random.key(0), nenv) + d = jax_jit(init)(key) + jax.block_until_ready(d) + + @jax_jit + def unroll(d): + def fn(d, _): + d = d.replace(qpos=d.qpos + 0 * d.qpos) + return step_fn(mx, d), None + + return jax.lax.scan(fn, d, None, length=nstep, unroll=unroll_steps) + + jit_time, run_time = _measure(unroll, d) + steps = nstep * nenv + + return jit_time, run_time, steps + + +def benchmark_raw_jax_warp( + m: mujoco.MjModel, + nstep: int = 1000, + nenv: int = 8192, + unroll_steps: int = 4, + function: str = 'kinematics', +): + if function not in ('kinematics', 'forward', 'step', 'collision'): + raise NotImplementedError( + f'{function} is not implemented for raw warp speed test.' + ) + + def warp_fn( + qpos_in: wp.array2d(dtype=wp.float32), + qpos_out: wp.array2d(dtype=wp.float32), + time_: wp.array(dtype=wp.float32), + xpos: wp.array2d(dtype=wp.vec3), + xquat: wp.array2d(dtype=wp.quat), + xmat: wp.array2d(dtype=wp.mat33), + xipos: wp.array2d(dtype=wp.vec3), + ximat: wp.array2d(dtype=wp.mat33), + xanchor: wp.array2d(dtype=wp.vec3), + xaxis: wp.array2d(dtype=wp.vec3), + geom_xpos: wp.array2d(dtype=wp.vec3), + geom_xmat: wp.array2d(dtype=wp.mat33), + ): + wp.copy(d.qpos, qpos_in) + if function == 'kinematics': + mjwarp.kinematics(mw, d) + elif function == 'forward': + mjwarp.forward(mw, d) + elif function == 'step': + mjwarp.step(mw, d) + elif function == 'collision': + mjwarp.collision(mw, d) + else: + raise NotImplementedError(f'{function} not implemented in speed test.') + wp.copy(qpos_out, d.qpos) + wp.copy(time_, d.time) + wp.copy(xpos, d.xpos) + wp.copy(xquat, d.xquat) + wp.copy(xmat, d.xmat) + wp.copy(xipos, d.xipos) + wp.copy(ximat, d.ximat) + wp.copy(xanchor, d.xanchor) + wp.copy(xaxis, d.xaxis) + wp.copy(geom_xpos, d.geom_xpos) + wp.copy(geom_xmat, d.geom_xmat) + + def unroll( + qpos, + qpos_out, + time_, + xpos, + xquat, + xmat, + xipos, + ximat, + xanchor, + xaxis, + geom_xpos, + geom_xmat, + ): + def step(carry, _): + qpos_, *_ = carry + out = warp_fn_jax(qpos_ + 0.0 * qpos_) + out = tuple(out) + return (out[0],) + out, None + + ( + qpos, + qpos_out, + time_, + xpos, + xquat, + xmat, + xipos, + ximat, + xanchor, + xaxis, + geom_xpos, + geom_xmat, + ), _ = jax.lax.scan( + step, + ( + qpos, + qpos_out, + time_, + xpos, + xquat, + xmat, + xipos, + ximat, + xanchor, + xaxis, + geom_xpos, + geom_xmat, + ), + length=nstep, + unroll=unroll_steps, + ) + + return ( + qpos, + qpos_out, + time_, + xpos, + xquat, + xmat, + xipos, + ximat, + xanchor, + xaxis, + geom_xpos, + geom_xmat, + ) + + output_dims = { + 'time_': (nenv,), + 'qpos_out': (nenv, m.nq), + 'xpos': (nenv, m.nbody, 3), + 'xquat': (nenv, m.nbody, 4), + 'xmat': (nenv, m.nbody, 3, 3), + 'xipos': (nenv, m.nbody, 3), + 'ximat': (nenv, m.nbody, 3, 3), + 'xanchor': (nenv, m.njnt, 3), + 'xaxis': (nenv, m.njnt, 3), + 'geom_xpos': (nenv, m.ngeom, 3), + 'geom_xmat': (nenv, m.ngeom, 3, 3), + } + warp_fn_jax = warp_ffi.jax_callable( + warp_fn, + num_outputs=11, + output_dims=output_dims, + ) + + @jax.vmap + def init(key): + d = mjx.make_data(m, impl='jax') + return d + + key = jax.random.split(jax.random.key(0), nenv) + dx = jax_jit(init)(key) + d_ = mujoco.MjData(m) + mw = mjwarp.put_model(m) + mw.opt.graph_conditional = False + d = mjwarp.put_data( + m, d_, nworld=nenv, nconmax=_NCONMAX.value, njmax=_NJMAX.value + ) + + jax_unroll_fn = jax_jit(unroll) + jit_time, run_time = _measure( + jax_unroll_fn, + dx.qpos, + dx.qpos, + dx.time, + dx.xpos, + dx.xquat, + dx.xmat, + dx.xipos, + dx.ximat, + dx.xanchor, + dx.xaxis, + dx.geom_xpos, + dx.geom_xmat, + ) + steps = nstep * nenv + + return jit_time, run_time, steps + + +def _compile_fn(fn, m, d): + fn(m, d) + fn(m, d) + with wp.ScopedCapture() as capture: + fn(m, d) + return capture.graph + + +def benchmark_raw_warp( + m: mujoco.MjModel, + nstep: int = 1000, + nenv: int = 8192, + unroll_steps: int = 4, + function: str = 'kinematics', +): + """Benchmarks raw warp.""" + del unroll_steps + if function not in ('kinematics', 'forward', 'step', 'collision'): + raise NotImplementedError( + f'{function} is not implemented for raw warp speed test.' + ) + + mw = mjwarp.put_model(m) + # TODO(btaba): re-enable graph conditional once JAX supports it, for fair + # comparison. + mw.opt.graph_conditional = False + dw = mjwarp.make_data( + m, nworld=nenv, nconmax=_NCONMAX.value, njmax=_NJMAX.value + ) + + if function == 'kinematics': + fn = mjwarp.kinematics + elif function == 'forward': + fn = mjwarp.forward + elif function == 'step': + fn = mjwarp.step + elif function == 'collision': + fn = mjwarp.collision + else: + raise NotImplementedError(f'{function} not implemented in speed test.') + + start = time.time() + graph = _compile_fn(fn, mw, dw) + jit_time = time.time() - start + + start = time.time() + for _ in range(nstep): + wp.capture_launch(graph) + wp.synchronize() + + run_time = time.time() - start + return jit_time, run_time, nstep * nenv + + +def _main(_: Sequence[str]): + """Runs testpeed function.""" + os.environ['MJX_WARP_ENABLED'] = 'true' + + if _WP_KERNEL_CACHE_DIR.value: + wp.config.kernel_cache_dir = _WP_KERNEL_CACHE_DIR.value + + modelfile = _MODELFILE.value + function_ = _FUNCTION.value + nstep, nenv, unroll = _NSTEP.value, _NENV.value, _UNROLL.value + + try: + m = test_util.load_test_file(modelfile) + except Exception as _: + m = mujoco.MjModel.from_xml_path(modelfile) + + mx = mjx.put_model(m, impl='jax') + mw = mjx.put_model(m, impl='warp') + + if function_ == 'kinematics': + func_warp = jax.vmap(wp_smooth.kinematics, in_axes=(None, 0)) + func_jax = jax.vmap(mjx.kinematics, in_axes=(None, 0)) + elif function_ == 'forward': + func_warp = jax.vmap(wp_forward.forward, in_axes=(None, 0)) + func_jax = jax.vmap(mjx.forward, in_axes=(None, 0)) + elif function_ == 'step': + func_warp = jax.vmap(wp_forward.step, in_axes=(None, 0)) + func_jax = jax.vmap(mjx.step, in_axes=(None, 0)) + elif function_ == 'collision': + func_warp = jax.vmap(wp_collision.collision, in_axes=(None, 0)) + func_jax = jax.vmap(mjx.collision, in_axes=(None, 0)) + else: + raise ValueError(f'Unknown function: {function_}') + + print('testspeed.py:\n') + print(f' modelfile : {modelfile}') + print(f' function : {function_}') + print(f' nenv : {nenv}') + print(f' nstep : {nstep}') + print(f' timestep : {m.opt.timestep}') + print(f' unroll : {unroll}\n') + + for name, mx_, op in ( + ('JAX WARP FFI', mw, func_warp), + ('Pure JAX', mx, func_jax), + ): + if op is not None: + jit_time, run_time, steps = benchmark(m, mx_, op, nstep, nenv, unroll) + + print(f' {name}:') + print(f' JIT time : {jit_time:.2f} s') + print(f' simulation time : {run_time:.2f} s') + print(f' steps per second : {steps / run_time:,.0f}') + print( + f' realtime factor : {steps * m.opt.timestep / run_time:.2f} x' + ) + print(f' time per step : {1e6 * run_time / steps:.2f} µs\n') + + jit_time, run_time, steps = benchmark_raw_jax_warp( + m, nstep, nenv, unroll, function=function_ + ) + print(' Pure JAX-WARP:') + print(f' JIT time : {jit_time:.2f} s') + print(f' simulation time : {run_time:.2f} s') + print(f' steps per second : {steps / run_time:,.0f}') + print(f' realtime factor : {steps * m.opt.timestep / run_time:.2f} x') + print(f' time per step : {1e6 * run_time / steps:.2f} µs\n') + + jit_time, run_time, steps = benchmark_raw_warp( + m, nstep, nenv, unroll, function=function_ + ) + print(' Pure WARP:') + print(f' JIT time : {jit_time:.2f} s') + print(f' simulation time : {run_time:.2f} s') + print(f' steps per second : {steps / run_time:,.0f}') + print(f' realtime factor : {steps * m.opt.timestep / run_time:.2f} x') + print(f' time per step : {1e6 * run_time / steps:.2f} µs\n') + + +def main(): + app.run(_main) + + +if __name__ == '__main__': + main() diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py new file mode 100644 index 00000000..17cab394 --- /dev/null +++ b/mjx/mujoco/mjx/warp/types.py @@ -0,0 +1,1602 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""MJX Warp types. + +DO NOT EDIT. This file is auto-generated. +""" +import dataclasses +from typing import Tuple +import jax +from jax import tree_util +from jax.interpreters import batching +from mujoco.mjx._src import dataclasses as mjx_dataclasses +import numpy as np + +PyTreeNode = mjx_dataclasses.PyTreeNode + + +@dataclasses.dataclass(frozen=True) +@tree_util.register_pytree_node_class +class TileSet: + """Tiling configuration for decomposable block diagonal matrix. + + For non-square, non-block-diagonal tiles, use two tilesets. + + Attributes: + adr: address of each tile in the set + size: size of all the tiles in this set + """ + + adr: np.ndarray + size: int + + def tree_flatten(self): + children = list((getattr(self, k) for k in self.__dataclass_fields__)) + return (children, None) + + @classmethod + def tree_unflatten(cls, aux_data, children): + del aux_data + return cls(*children) + + +@dataclasses.dataclass(frozen=True) +@tree_util.register_pytree_node_class +class BlockDim: + """Block dimension 'block_dim' settings for wp.launch_tiled. + + TODO(team): experimental and may be removed + """ + + actuator_velocity: int + cholesky_factorize: int + cholesky_factorize_solve: int + cholesky_solve: int + energy_vel_kinetic: int + euler_dense: int + mul_m_dense: int + qderiv_actuator_passive_actuation: int + qderiv_actuator_passive_no_actuation: int + qfrc_actuator: int + ray: int + segmented_sort: int + tendon_velocity: int + update_gradient_cholesky: int + + def tree_flatten(self): + children = list((getattr(self, k) for k in self.__dataclass_fields__)) + return (children, None) + + @classmethod + def tree_unflatten(cls, aux_data, children): + del aux_data + return cls(*children) + + +class StatisticWarp(PyTreeNode): + """Derived fields from Statistic.""" + + meaninertia: float + + +class OptionWarp(PyTreeNode): + """Derived fields from Option.""" + + broadphase: int + broadphase_filter: int + epa_iterations: int + gjk_iterations: int + graph_conditional: bool + has_fluid: bool + is_sparse: bool + ls_parallel: bool + run_collision_detection: bool + sdf_initpoints: int + sdf_iterations: int + + +class ModelWarp(PyTreeNode): + """Derived fields from Model.""" + + M_colind: np.ndarray + M_rowadr: np.ndarray + M_rownnz: np.ndarray + actuator_affine_bias_gain: bool + actuator_moment_tiles_nu: Tuple[TileSet, ...] + actuator_moment_tiles_nv: Tuple[TileSet, ...] + actuator_trntype_body_adr: np.ndarray + block_dim: BlockDim + body_tree: Tuple[np.ndarray, ...] + condim_max: int + dof_tri_col: np.ndarray + dof_tri_row: np.ndarray + eq_connect_adr: np.ndarray + eq_jnt_adr: np.ndarray + eq_ten_adr: np.ndarray + eq_wld_adr: np.ndarray + flex_bending: np.ndarray + flex_damping: np.ndarray + flex_dim: np.ndarray + flex_edge: np.ndarray + flex_edgeadr: np.ndarray + flex_edgeflap: np.ndarray + flex_elem: np.ndarray + flex_elemedge: np.ndarray + flex_elemedgeadr: np.ndarray + flex_stiffness: np.ndarray + flex_vertadr: np.ndarray + flex_vertbodyid: np.ndarray + flex_vertnum: np.ndarray + flexedge_length0: np.ndarray + geom_pair_type_count: Tuple[int, ...] + geom_plugin_index: np.ndarray + geompair2hfgeompair: np.ndarray + has_sdf_geom: bool + jnt_limited_ball_adr: np.ndarray + jnt_limited_slide_hinge_adr: np.ndarray + light_bodyid: np.ndarray + light_targetbodyid: np.ndarray + mapM2M: np.ndarray + mesh_polyadr: np.ndarray + mesh_polymap: np.ndarray + mesh_polymapadr: np.ndarray + mesh_polymapnum: np.ndarray + mesh_polynormal: np.ndarray + mesh_polynum: np.ndarray + mesh_polyvert: np.ndarray + mesh_polyvertadr: np.ndarray + mesh_polyvertnum: np.ndarray + mocap_bodyid: np.ndarray + nflex: int + nflexedge: int + nflexelem: int + nflexelemdata: int + nflexvert: int + nlsp: int + nmeshpoly: int + nmeshpolymap: int + nmeshpolyvert: int + nxn_geom_pair: np.ndarray + nxn_geom_pair_filtered: np.ndarray + nxn_pairid: np.ndarray + nxn_pairid_filtered: np.ndarray + plugin: np.ndarray + plugin_attr: np.ndarray + qLD_updates: Tuple[np.ndarray, ...] + qM_fullm_i: np.ndarray + qM_fullm_j: np.ndarray + qM_madr_ij: np.ndarray + qM_mulm_i: np.ndarray + qM_mulm_j: np.ndarray + qM_tiles: Tuple[TileSet, ...] + rangefinder_sensor_adr: np.ndarray + sensor_acc_adr: np.ndarray + sensor_e_kinetic: bool + sensor_e_potential: bool + sensor_limitfrc_adr: np.ndarray + sensor_limitpos_adr: np.ndarray + sensor_limitvel_adr: np.ndarray + sensor_pos_adr: np.ndarray + sensor_rangefinder_adr: np.ndarray + sensor_rangefinder_bodyid: np.ndarray + sensor_rne_postconstraint: bool + sensor_subtree_vel: bool + sensor_tendonactfrc_adr: np.ndarray + sensor_touch_adr: np.ndarray + sensor_vel_adr: np.ndarray + subtree_mass: jax.Array + ten_wrapadr_site: np.ndarray + ten_wrapnum_site: np.ndarray + tendon_geom_adr: np.ndarray + tendon_jnt_adr: np.ndarray + tendon_limited_adr: np.ndarray + tendon_site_pair_adr: np.ndarray + wrap_geom_adr: np.ndarray + wrap_jnt_adr: np.ndarray + wrap_pulley_scale: np.ndarray + wrap_site_adr: np.ndarray + wrap_site_pair_adr: np.ndarray + + +class DataWarp(PyTreeNode): + """Derived fields from Data.""" + + act_dot_rk: jax.Array + act_t0: jax.Array + act_vel_integration: jax.Array + actuator_length: jax.Array + actuator_moment: jax.Array + actuator_trntype_body_ncon: jax.Array + actuator_velocity: jax.Array + cacc: jax.Array + cdof: jax.Array + cdof_dot: jax.Array + cfrc_ext: jax.Array + cfrc_int: jax.Array + cinert: jax.Array + collision_hftri_index: jax.Array + collision_pair: jax.Array + collision_pairid: jax.Array + collision_worldid: jax.Array + contact__dim: jax.Array + contact__dist: jax.Array + contact__efc_address: jax.Array + contact__frame: jax.Array + contact__friction: jax.Array + contact__geom: jax.Array + contact__includemargin: jax.Array + contact__pos: jax.Array + contact__solimp: jax.Array + contact__solref: jax.Array + contact__solreffriction: jax.Array + contact__worldid: jax.Array + crb: jax.Array + efc__D: jax.Array + efc__J: jax.Array + efc__Jaref: jax.Array + efc__Ma: jax.Array + efc__Mgrad: jax.Array + efc__active: jax.Array + efc__alpha: jax.Array + efc__aref: jax.Array + efc__beta: jax.Array + efc__beta_den: jax.Array + efc__beta_num: jax.Array + efc__cholesky_L_tmp: jax.Array + efc__cholesky_y_tmp: jax.Array + efc__condim: jax.Array + efc__cost: jax.Array + efc__cost_candidate: jax.Array + efc__done: jax.Array + efc__force: jax.Array + efc__frictionloss: jax.Array + efc__gauss: jax.Array + efc__grad: jax.Array + efc__grad_dot: jax.Array + efc__gtol: jax.Array + efc__h: jax.Array + efc__hi: jax.Array + efc__hi_alpha: jax.Array + efc__hi_next: jax.Array + efc__hi_next_alpha: jax.Array + efc__id: jax.Array + efc__jv: jax.Array + efc__lo: jax.Array + efc__lo_alpha: jax.Array + efc__lo_next: jax.Array + efc__lo_next_alpha: jax.Array + efc__ls_done: jax.Array + efc__margin: jax.Array + efc__mid: jax.Array + efc__mid_alpha: jax.Array + efc__mv: jax.Array + efc__p0: jax.Array + efc__pos: jax.Array + efc__prev_Mgrad: jax.Array + efc__prev_cost: jax.Array + efc__prev_grad: jax.Array + efc__quad: jax.Array + efc__quad_gauss: jax.Array + efc__search: jax.Array + efc__search_dot: jax.Array + efc__type: jax.Array + efc__u: jax.Array + efc__uu: jax.Array + efc__uv: jax.Array + efc__vel: jax.Array + efc__vv: jax.Array + energy: jax.Array + energy_vel_mul_m_skip: jax.Array + epa_face: jax.Array + epa_horizon: jax.Array + epa_index: jax.Array + epa_map: jax.Array + epa_norm2: jax.Array + epa_pr: jax.Array + epa_vert: jax.Array + epa_vert1: jax.Array + epa_vert2: jax.Array + epa_vert_index1: jax.Array + epa_vert_index2: jax.Array + flexedge_length: jax.Array + flexedge_velocity: jax.Array + flexvert_xpos: jax.Array + fluid_applied: jax.Array + geom_skip: jax.Array + inverse_mul_m_skip: jax.Array + light_xdir: jax.Array + light_xpos: jax.Array + ncollision: jax.Array + ncon: jax.Array + ncon_hfield: jax.Array + nconmax: int + ne: jax.Array + ne_connect: jax.Array + ne_jnt: jax.Array + ne_ten: jax.Array + ne_weld: jax.Array + nefc: jax.Array + nf: jax.Array + njmax: int + nl: jax.Array + nsolving: jax.Array + nworld: int + qLD: jax.Array + qLD_integration: jax.Array + qLDiagInv: jax.Array + qLDiagInv_integration: jax.Array + qM: jax.Array + qM_integration: jax.Array + qacc_discrete: jax.Array + qacc_integration: jax.Array + qacc_rk: jax.Array + qfrc_damper: jax.Array + qfrc_integration: jax.Array + qfrc_spring: jax.Array + qpos_t0: jax.Array + qvel_rk: jax.Array + qvel_t0: jax.Array + ray_bodyexclude: jax.Array + ray_dist: jax.Array + ray_geomid: jax.Array + sap_cumulative_sum: jax.Array + sap_projection_lower: jax.Array + sap_projection_upper: jax.Array + sap_range: jax.Array + sap_segment_index: jax.Array + sap_sort_index: jax.Array + sensor_rangefinder_dist: jax.Array + sensor_rangefinder_geomid: jax.Array + sensor_rangefinder_pnt: jax.Array + sensor_rangefinder_vec: jax.Array + solver_niter: jax.Array + subtree_angmom: jax.Array + subtree_bodyvel: jax.Array + subtree_linvel: jax.Array + ten_J: jax.Array + ten_Jdot: jax.Array + ten_actfrc: jax.Array + ten_bias_coef: jax.Array + ten_length: jax.Array + ten_velocity: jax.Array + ten_wrapadr: jax.Array + ten_wrapnum: jax.Array + wrap_geom_xpos: jax.Array + wrap_obj: jax.Array + wrap_xpos: jax.Array + shape = property(lambda self: self.cacc.shape) + + +DATA_NON_VMAP = { + 'collision_hftri_index', + 'collision_pair', + 'collision_pairid', + 'collision_worldid', + 'contact__dim', + 'contact__dist', + 'contact__efc_address', + 'contact__frame', + 'contact__friction', + 'contact__geom', + 'contact__includemargin', + 'contact__pos', + 'contact__solimp', + 'contact__solref', + 'contact__solreffriction', + 'contact__worldid', + 'efc__u', + 'efc__uu', + 'efc__uv', + 'efc__vv', + 'epa_face', + 'epa_horizon', + 'epa_index', + 'epa_map', + 'epa_norm2', + 'epa_pr', + 'epa_vert', + 'epa_vert1', + 'epa_vert2', + 'epa_vert_index1', + 'epa_vert_index2', + 'geom_skip', + 'ncollision', + 'ncon', + 'nconmax', + 'njmax', + 'nsolving', + 'nworld', + 'ray_bodyexclude', +} + + +def _to_elt(cont, _, d, axis): + return DataWarp(**{ + f.name: ( + cont(getattr(d, f.name), axis) + if f.name not in DATA_NON_VMAP + else getattr(d, f.name) + ) + for f in DataWarp.fields() + }) + + +def _from_elt(cont, axis_size, d, axis_dest): + return DataWarp(**{ + f.name: ( + cont(axis_size, getattr(d, f.name), axis_dest) + if f.name not in DATA_NON_VMAP + else getattr(d, f.name) + ) + for f in DataWarp.fields() + }) + + +batching.register_vmappable(DataWarp, int, int, _to_elt, _from_elt, None) + +NDIM = { + 'Data': { + 'act': 2, + 'act_dot': 2, + 'act_dot_rk': 2, + 'act_t0': 2, + 'act_vel_integration': 2, + 'actuator_force': 2, + 'actuator_length': 2, + 'actuator_moment': 3, + 'actuator_trntype_body_ncon': 2, + 'actuator_velocity': 2, + 'cacc': 3, + 'cam_xmat': 4, + 'cam_xpos': 3, + 'cdof': 3, + 'cdof_dot': 3, + 'cfrc_ext': 3, + 'cfrc_int': 3, + 'cinert': 3, + 'collision_hftri_index': 1, + 'collision_pair': 2, + 'collision_pairid': 1, + 'collision_worldid': 1, + 'contact__dim': 1, + 'contact__dist': 1, + 'contact__efc_address': 2, + 'contact__frame': 3, + 'contact__friction': 2, + 'contact__geom': 2, + 'contact__includemargin': 1, + 'contact__pos': 2, + 'contact__solimp': 2, + 'contact__solref': 2, + 'contact__solreffriction': 2, + 'contact__worldid': 1, + 'crb': 3, + 'ctrl': 2, + 'cvel': 3, + 'efc__D': 2, + 'efc__J': 3, + 'efc__Jaref': 2, + 'efc__Ma': 2, + 'efc__Mgrad': 2, + 'efc__active': 2, + 'efc__alpha': 1, + 'efc__aref': 2, + 'efc__beta': 1, + 'efc__beta_den': 1, + 'efc__beta_num': 1, + 'efc__cholesky_L_tmp': 3, + 'efc__cholesky_y_tmp': 2, + 'efc__condim': 2, + 'efc__cost': 1, + 'efc__cost_candidate': 2, + 'efc__done': 1, + 'efc__force': 2, + 'efc__frictionloss': 2, + 'efc__gauss': 1, + 'efc__grad': 2, + 'efc__grad_dot': 1, + 'efc__gtol': 1, + 'efc__h': 3, + 'efc__hi': 2, + 'efc__hi_alpha': 1, + 'efc__hi_next': 2, + 'efc__hi_next_alpha': 1, + 'efc__id': 2, + 'efc__jv': 2, + 'efc__lo': 2, + 'efc__lo_alpha': 1, + 'efc__lo_next': 2, + 'efc__lo_next_alpha': 1, + 'efc__ls_done': 1, + 'efc__margin': 2, + 'efc__mid': 2, + 'efc__mid_alpha': 1, + 'efc__mv': 2, + 'efc__p0': 2, + 'efc__pos': 2, + 'efc__prev_Mgrad': 2, + 'efc__prev_cost': 1, + 'efc__prev_grad': 2, + 'efc__quad': 3, + 'efc__quad_gauss': 2, + 'efc__search': 2, + 'efc__search_dot': 1, + 'efc__type': 2, + 'efc__u': 2, + 'efc__uu': 1, + 'efc__uv': 1, + 'efc__vel': 2, + 'efc__vv': 1, + 'energy': 2, + 'energy_vel_mul_m_skip': 1, + 'epa_face': 3, + 'epa_horizon': 2, + 'epa_index': 2, + 'epa_map': 2, + 'epa_norm2': 2, + 'epa_pr': 3, + 'epa_vert': 3, + 'epa_vert1': 3, + 'epa_vert2': 3, + 'epa_vert_index1': 2, + 'epa_vert_index2': 2, + 'eq_active': 2, + 'flexedge_length': 2, + 'flexedge_velocity': 2, + 'flexvert_xpos': 3, + 'fluid_applied': 3, + 'geom_skip': 1, + 'geom_xmat': 4, + 'geom_xpos': 3, + 'inverse_mul_m_skip': 1, + 'light_xdir': 3, + 'light_xpos': 3, + 'mocap_pos': 3, + 'mocap_quat': 3, + 'ncollision': 1, + 'ncon': 1, + 'ncon_hfield': 2, + 'nconmax': 0, + 'ne': 1, + 'ne_connect': 1, + 'ne_jnt': 1, + 'ne_ten': 1, + 'ne_weld': 1, + 'nefc': 1, + 'nf': 1, + 'njmax': 0, + 'nl': 1, + 'nsolving': 1, + 'nworld': 0, + 'qLD': 3, + 'qLD_integration': 3, + 'qLDiagInv': 2, + 'qLDiagInv_integration': 2, + 'qM': 3, + 'qM_integration': 3, + 'qacc': 2, + 'qacc_discrete': 2, + 'qacc_integration': 2, + 'qacc_rk': 2, + 'qacc_smooth': 2, + 'qacc_warmstart': 2, + 'qfrc_actuator': 2, + 'qfrc_applied': 2, + 'qfrc_bias': 2, + 'qfrc_constraint': 2, + 'qfrc_damper': 2, + 'qfrc_fluid': 2, + 'qfrc_gravcomp': 2, + 'qfrc_integration': 2, + 'qfrc_inverse': 2, + 'qfrc_passive': 2, + 'qfrc_smooth': 2, + 'qfrc_spring': 2, + 'qpos': 2, + 'qpos_t0': 2, + 'qvel': 2, + 'qvel_rk': 2, + 'qvel_t0': 2, + 'ray_bodyexclude': 1, + 'ray_dist': 2, + 'ray_geomid': 2, + 'sap_cumulative_sum': 2, + 'sap_projection_lower': 3, + 'sap_projection_upper': 2, + 'sap_range': 2, + 'sap_segment_index': 2, + 'sap_sort_index': 3, + 'sensor_rangefinder_dist': 2, + 'sensor_rangefinder_geomid': 2, + 'sensor_rangefinder_pnt': 3, + 'sensor_rangefinder_vec': 3, + 'sensordata': 2, + 'site_xmat': 4, + 'site_xpos': 3, + 'solver_niter': 1, + 'subtree_angmom': 3, + 'subtree_bodyvel': 3, + 'subtree_com': 3, + 'subtree_linvel': 3, + 'ten_J': 3, + 'ten_Jdot': 3, + 'ten_actfrc': 2, + 'ten_bias_coef': 2, + 'ten_length': 2, + 'ten_velocity': 2, + 'ten_wrapadr': 2, + 'ten_wrapnum': 2, + 'time': 1, + 'wrap_geom_xpos': 3, + 'wrap_obj': 3, + 'wrap_xpos': 3, + 'xanchor': 3, + 'xaxis': 3, + 'xfrc_applied': 3, + 'ximat': 4, + 'xipos': 3, + 'xmat': 4, + 'xpos': 3, + 'xquat': 3, + }, + 'Model': { + 'M_colind': 1, + 'M_rowadr': 1, + 'M_rownnz': 1, + 'actuator_acc0': 1, + 'actuator_actadr': 1, + 'actuator_actearly': 1, + 'actuator_actlimited': 1, + 'actuator_actnum': 1, + 'actuator_actrange': 3, + 'actuator_affine_bias_gain': 0, + 'actuator_biasprm': 3, + 'actuator_biastype': 1, + 'actuator_cranklength': 1, + 'actuator_ctrllimited': 1, + 'actuator_ctrlrange': 3, + 'actuator_dynprm': 3, + 'actuator_dyntype': 1, + 'actuator_forcelimited': 1, + 'actuator_forcerange': 3, + 'actuator_gainprm': 3, + 'actuator_gaintype': 1, + 'actuator_gear': 3, + 'actuator_lengthrange': 2, + 'actuator_moment_tiles_nu': -1, + 'actuator_moment_tiles_nv': -1, + 'actuator_trnid': 2, + 'actuator_trntype': 1, + 'actuator_trntype_body_adr': 1, + 'block_dim__actuator_velocity': 0, + 'block_dim__cholesky_factorize': 0, + 'block_dim__cholesky_factorize_solve': 0, + 'block_dim__cholesky_solve': 0, + 'block_dim__energy_vel_kinetic': 0, + 'block_dim__euler_dense': 0, + 'block_dim__mul_m_dense': 0, + 'block_dim__qderiv_actuator_passive_actuation': 0, + 'block_dim__qderiv_actuator_passive_no_actuation': 0, + 'block_dim__qfrc_actuator': 0, + 'block_dim__ray': 0, + 'block_dim__segmented_sort': 0, + 'block_dim__tendon_velocity': 0, + 'block_dim__update_gradient_cholesky': 0, + 'body_conaffinity': 1, + 'body_contype': 1, + 'body_dofadr': 1, + 'body_dofnum': 1, + 'body_geomadr': 1, + 'body_geomnum': 1, + 'body_gravcomp': 2, + 'body_inertia': 3, + 'body_invweight0': 3, + 'body_ipos': 3, + 'body_iquat': 3, + 'body_jntadr': 1, + 'body_jntnum': 1, + 'body_mass': 2, + 'body_mocapid': 1, + 'body_parentid': 1, + 'body_pos': 3, + 'body_quat': 3, + 'body_rootid': 1, + 'body_subtreemass': 2, + 'body_tree': -1, + 'body_weldid': 1, + 'cam_bodyid': 1, + 'cam_fovy': 1, + 'cam_intrinsic': 2, + 'cam_mat0': 4, + 'cam_mode': 1, + 'cam_pos': 3, + 'cam_pos0': 3, + 'cam_poscom0': 3, + 'cam_quat': 3, + 'cam_resolution': 2, + 'cam_sensorsize': 2, + 'cam_targetbodyid': 1, + 'condim_max': 0, + 'dof_Madr': 1, + 'dof_armature': 2, + 'dof_bodyid': 1, + 'dof_damping': 2, + 'dof_frictionloss': 2, + 'dof_invweight0': 2, + 'dof_jntid': 1, + 'dof_parentid': 1, + 'dof_solimp': 3, + 'dof_solref': 3, + 'dof_tri_col': 1, + 'dof_tri_row': 1, + 'eq_active0': 1, + 'eq_connect_adr': 1, + 'eq_data': 3, + 'eq_jnt_adr': 1, + 'eq_obj1id': 1, + 'eq_obj2id': 1, + 'eq_objtype': 1, + 'eq_solimp': 3, + 'eq_solref': 3, + 'eq_ten_adr': 1, + 'eq_type': 1, + 'eq_wld_adr': 1, + 'exclude_signature': 1, + 'flex_bending': 3, + 'flex_damping': 1, + 'flex_dim': 1, + 'flex_edge': 2, + 'flex_edgeadr': 1, + 'flex_edgeflap': 2, + 'flex_elem': 1, + 'flex_elemedge': 1, + 'flex_elemedgeadr': 1, + 'flex_stiffness': 1, + 'flex_vertadr': 1, + 'flex_vertbodyid': 1, + 'flex_vertnum': 1, + 'flexedge_length0': 1, + 'geom_aabb': 3, + 'geom_bodyid': 1, + 'geom_conaffinity': 1, + 'geom_condim': 1, + 'geom_contype': 1, + 'geom_dataid': 1, + 'geom_friction': 3, + 'geom_gap': 2, + 'geom_group': 1, + 'geom_margin': 2, + 'geom_matid': 2, + 'geom_pair_type_count': -1, + 'geom_plugin_index': 1, + 'geom_pos': 3, + 'geom_priority': 1, + 'geom_quat': 3, + 'geom_rbound': 2, + 'geom_rgba': 3, + 'geom_size': 3, + 'geom_solimp': 3, + 'geom_solmix': 2, + 'geom_solref': 3, + 'geom_type': 1, + 'geompair2hfgeompair': 1, + 'has_sdf_geom': 0, + 'hfield_adr': 1, + 'hfield_data': 1, + 'hfield_ncol': 1, + 'hfield_nrow': 1, + 'hfield_size': 2, + 'jnt_actfrclimited': 1, + 'jnt_actfrcrange': 3, + 'jnt_actgravcomp': 1, + 'jnt_axis': 3, + 'jnt_bodyid': 1, + 'jnt_dofadr': 1, + 'jnt_limited': 1, + 'jnt_limited_ball_adr': 1, + 'jnt_limited_slide_hinge_adr': 1, + 'jnt_margin': 2, + 'jnt_pos': 3, + 'jnt_qposadr': 1, + 'jnt_range': 3, + 'jnt_solimp': 3, + 'jnt_solref': 3, + 'jnt_stiffness': 2, + 'jnt_type': 1, + 'light_bodyid': 1, + 'light_dir': 3, + 'light_dir0': 3, + 'light_mode': 1, + 'light_pos': 3, + 'light_pos0': 3, + 'light_poscom0': 3, + 'light_targetbodyid': 1, + 'mapM2M': 1, + 'mat_rgba': 3, + 'mesh_face': 2, + 'mesh_faceadr': 1, + 'mesh_graph': 1, + 'mesh_graphadr': 1, + 'mesh_polyadr': 1, + 'mesh_polymap': 1, + 'mesh_polymapadr': 1, + 'mesh_polymapnum': 1, + 'mesh_polynormal': 2, + 'mesh_polynum': 1, + 'mesh_polyvert': 1, + 'mesh_polyvertadr': 1, + 'mesh_polyvertnum': 1, + 'mesh_vert': 2, + 'mesh_vertadr': 1, + 'mesh_vertnum': 1, + 'mocap_bodyid': 1, + 'nC': 0, + 'nM': 0, + 'na': 0, + 'nbody': 0, + 'ncam': 0, + 'neq': 0, + 'nexclude': 0, + 'nflex': 0, + 'nflexedge': 0, + 'nflexelem': 0, + 'nflexelemdata': 0, + 'nflexvert': 0, + 'ngeom': 0, + 'ngravcomp': 0, + 'nhfield': 0, + 'nhfielddata': 0, + 'njnt': 0, + 'nlight': 0, + 'nlsp': 0, + 'nmeshface': 0, + 'nmeshgraph': 0, + 'nmeshpoly': 0, + 'nmeshpolymap': 0, + 'nmeshpolyvert': 0, + 'nmeshvert': 0, + 'nmocap': 0, + 'npair': 0, + 'nq': 0, + 'nsensor': 0, + 'nsensordata': 0, + 'nsite': 0, + 'ntendon': 0, + 'nu': 0, + 'nv': 0, + 'nwrap': 0, + 'nxn_geom_pair': 2, + 'nxn_geom_pair_filtered': 2, + 'nxn_pairid': 1, + 'nxn_pairid_filtered': 1, + 'opt__broadphase': 0, + 'opt__broadphase_filter': 0, + 'opt__cone': 0, + 'opt__density': 1, + 'opt__disableflags': 0, + 'opt__enableflags': 0, + 'opt__epa_iterations': 0, + 'opt__gjk_iterations': 0, + 'opt__graph_conditional': 0, + 'opt__gravity': 2, + 'opt__has_fluid': 0, + 'opt__impratio': 1, + 'opt__integrator': 0, + 'opt__is_sparse': 0, + 'opt__iterations': 0, + 'opt__ls_iterations': 0, + 'opt__ls_parallel': 0, + 'opt__ls_tolerance': 1, + 'opt__magnetic': 2, + 'opt__run_collision_detection': 0, + 'opt__sdf_initpoints': 0, + 'opt__sdf_iterations': 0, + 'opt__solver': 0, + 'opt__timestep': 1, + 'opt__tolerance': 1, + 'opt__viscosity': 1, + 'opt__wind': 2, + 'pair_dim': 1, + 'pair_friction': 3, + 'pair_gap': 2, + 'pair_geom1': 1, + 'pair_geom2': 1, + 'pair_margin': 2, + 'pair_solimp': 3, + 'pair_solref': 3, + 'pair_solreffriction': 3, + 'plugin': 1, + 'plugin_attr': 2, + 'qLD_updates': -1, + 'qM_fullm_i': 1, + 'qM_fullm_j': 1, + 'qM_madr_ij': 1, + 'qM_mulm_i': 1, + 'qM_mulm_j': 1, + 'qM_tiles': -1, + 'qpos0': 2, + 'qpos_spring': 2, + 'rangefinder_sensor_adr': 1, + 'sensor_acc_adr': 1, + 'sensor_adr': 1, + 'sensor_cutoff': 1, + 'sensor_datatype': 1, + 'sensor_dim': 1, + 'sensor_e_kinetic': 0, + 'sensor_e_potential': 0, + 'sensor_limitfrc_adr': 1, + 'sensor_limitpos_adr': 1, + 'sensor_limitvel_adr': 1, + 'sensor_objid': 1, + 'sensor_objtype': 1, + 'sensor_pos_adr': 1, + 'sensor_rangefinder_adr': 1, + 'sensor_rangefinder_bodyid': 1, + 'sensor_refid': 1, + 'sensor_reftype': 1, + 'sensor_rne_postconstraint': 0, + 'sensor_subtree_vel': 0, + 'sensor_tendonactfrc_adr': 1, + 'sensor_touch_adr': 1, + 'sensor_type': 1, + 'sensor_vel_adr': 1, + 'site_bodyid': 1, + 'site_pos': 3, + 'site_quat': 3, + 'site_size': 2, + 'site_type': 1, + 'stat__meaninertia': 0, + 'subtree_mass': 2, + 'ten_wrapadr_site': 1, + 'ten_wrapnum_site': 1, + 'tendon_actfrclimited': 1, + 'tendon_actfrcrange': 3, + 'tendon_adr': 1, + 'tendon_armature': 2, + 'tendon_damping': 2, + 'tendon_frictionloss': 2, + 'tendon_geom_adr': 1, + 'tendon_invweight0': 2, + 'tendon_jnt_adr': 1, + 'tendon_length0': 2, + 'tendon_lengthspring': 3, + 'tendon_limited': 1, + 'tendon_limited_adr': 1, + 'tendon_margin': 2, + 'tendon_num': 1, + 'tendon_range': 3, + 'tendon_site_pair_adr': 1, + 'tendon_solimp_fri': 3, + 'tendon_solimp_lim': 3, + 'tendon_solref_fri': 3, + 'tendon_solref_lim': 3, + 'tendon_stiffness': 2, + 'wrap_geom_adr': 1, + 'wrap_jnt_adr': 1, + 'wrap_objid': 1, + 'wrap_prm': 1, + 'wrap_pulley_scale': 1, + 'wrap_site_adr': 1, + 'wrap_site_pair_adr': 1, + 'wrap_type': 1, + }, + 'Option': { + 'broadphase': 0, + 'broadphase_filter': 0, + 'cone': 0, + 'density': 1, + 'disableflags': 0, + 'enableflags': 0, + 'epa_iterations': 0, + 'gjk_iterations': 0, + 'graph_conditional': 0, + 'gravity': 2, + 'has_fluid': 0, + 'impratio': 1, + 'integrator': 0, + 'is_sparse': 0, + 'iterations': 0, + 'ls_iterations': 0, + 'ls_parallel': 0, + 'ls_tolerance': 1, + 'magnetic': 2, + 'run_collision_detection': 0, + 'sdf_initpoints': 0, + 'sdf_iterations': 0, + 'solver': 0, + 'timestep': 1, + 'tolerance': 1, + 'viscosity': 1, + 'wind': 2, + }, + 'Statistic': {'meaninertia': 0}, +} +BATCH_DIM = { + 'Data': { + 'act': True, + 'act_dot': True, + 'act_dot_rk': True, + 'act_t0': True, + 'act_vel_integration': True, + 'actuator_force': True, + 'actuator_length': True, + 'actuator_moment': True, + 'actuator_trntype_body_ncon': True, + 'actuator_velocity': True, + 'cacc': True, + 'cam_xmat': True, + 'cam_xpos': True, + 'cdof': True, + 'cdof_dot': True, + 'cfrc_ext': True, + 'cfrc_int': True, + 'cinert': True, + 'collision_hftri_index': False, + 'collision_pair': False, + 'collision_pairid': False, + 'collision_worldid': False, + 'contact__dim': False, + 'contact__dist': False, + 'contact__efc_address': False, + 'contact__frame': False, + 'contact__friction': False, + 'contact__geom': False, + 'contact__includemargin': False, + 'contact__pos': False, + 'contact__solimp': False, + 'contact__solref': False, + 'contact__solreffriction': False, + 'contact__worldid': False, + 'crb': True, + 'ctrl': True, + 'cvel': True, + 'efc__D': True, + 'efc__J': True, + 'efc__Jaref': True, + 'efc__Ma': True, + 'efc__Mgrad': True, + 'efc__active': True, + 'efc__alpha': True, + 'efc__aref': True, + 'efc__beta': True, + 'efc__beta_den': True, + 'efc__beta_num': True, + 'efc__cholesky_L_tmp': True, + 'efc__cholesky_y_tmp': True, + 'efc__condim': True, + 'efc__cost': True, + 'efc__cost_candidate': True, + 'efc__done': True, + 'efc__force': True, + 'efc__frictionloss': True, + 'efc__gauss': True, + 'efc__grad': True, + 'efc__grad_dot': True, + 'efc__gtol': True, + 'efc__h': True, + 'efc__hi': True, + 'efc__hi_alpha': True, + 'efc__hi_next': True, + 'efc__hi_next_alpha': True, + 'efc__id': True, + 'efc__jv': True, + 'efc__lo': True, + 'efc__lo_alpha': True, + 'efc__lo_next': True, + 'efc__lo_next_alpha': True, + 'efc__ls_done': True, + 'efc__margin': True, + 'efc__mid': True, + 'efc__mid_alpha': True, + 'efc__mv': True, + 'efc__p0': True, + 'efc__pos': True, + 'efc__prev_Mgrad': True, + 'efc__prev_cost': True, + 'efc__prev_grad': True, + 'efc__quad': True, + 'efc__quad_gauss': True, + 'efc__search': True, + 'efc__search_dot': True, + 'efc__type': True, + 'efc__u': False, + 'efc__uu': False, + 'efc__uv': False, + 'efc__vel': True, + 'efc__vv': False, + 'energy': True, + 'energy_vel_mul_m_skip': True, + 'epa_face': False, + 'epa_horizon': False, + 'epa_index': False, + 'epa_map': False, + 'epa_norm2': False, + 'epa_pr': False, + 'epa_vert': False, + 'epa_vert1': False, + 'epa_vert2': False, + 'epa_vert_index1': False, + 'epa_vert_index2': False, + 'eq_active': True, + 'flexedge_length': True, + 'flexedge_velocity': True, + 'flexvert_xpos': True, + 'fluid_applied': True, + 'geom_skip': False, + 'geom_xmat': True, + 'geom_xpos': True, + 'inverse_mul_m_skip': True, + 'light_xdir': True, + 'light_xpos': True, + 'mocap_pos': True, + 'mocap_quat': True, + 'ncollision': False, + 'ncon': False, + 'ncon_hfield': True, + 'nconmax': False, + 'ne': True, + 'ne_connect': True, + 'ne_jnt': True, + 'ne_ten': True, + 'ne_weld': True, + 'nefc': True, + 'nf': True, + 'njmax': False, + 'nl': True, + 'nsolving': False, + 'nworld': False, + 'qLD': True, + 'qLD_integration': True, + 'qLDiagInv': True, + 'qLDiagInv_integration': True, + 'qM': True, + 'qM_integration': True, + 'qacc': True, + 'qacc_discrete': True, + 'qacc_integration': True, + 'qacc_rk': True, + 'qacc_smooth': True, + 'qacc_warmstart': True, + 'qfrc_actuator': True, + 'qfrc_applied': True, + 'qfrc_bias': True, + 'qfrc_constraint': True, + 'qfrc_damper': True, + 'qfrc_fluid': True, + 'qfrc_gravcomp': True, + 'qfrc_integration': True, + 'qfrc_inverse': True, + 'qfrc_passive': True, + 'qfrc_smooth': True, + 'qfrc_spring': True, + 'qpos': True, + 'qpos_t0': True, + 'qvel': True, + 'qvel_rk': True, + 'qvel_t0': True, + 'ray_bodyexclude': False, + 'ray_dist': True, + 'ray_geomid': True, + 'sap_cumulative_sum': True, + 'sap_projection_lower': True, + 'sap_projection_upper': True, + 'sap_range': True, + 'sap_segment_index': True, + 'sap_sort_index': True, + 'sensor_rangefinder_dist': True, + 'sensor_rangefinder_geomid': True, + 'sensor_rangefinder_pnt': True, + 'sensor_rangefinder_vec': True, + 'sensordata': True, + 'site_xmat': True, + 'site_xpos': True, + 'solver_niter': True, + 'subtree_angmom': True, + 'subtree_bodyvel': True, + 'subtree_com': True, + 'subtree_linvel': True, + 'ten_J': True, + 'ten_Jdot': True, + 'ten_actfrc': True, + 'ten_bias_coef': True, + 'ten_length': True, + 'ten_velocity': True, + 'ten_wrapadr': True, + 'ten_wrapnum': True, + 'time': True, + 'wrap_geom_xpos': True, + 'wrap_obj': True, + 'wrap_xpos': True, + 'xanchor': True, + 'xaxis': True, + 'xfrc_applied': True, + 'ximat': True, + 'xipos': True, + 'xmat': True, + 'xpos': True, + 'xquat': True, + }, + 'Model': { + 'M_colind': False, + 'M_rowadr': False, + 'M_rownnz': False, + 'actuator_acc0': False, + 'actuator_actadr': False, + 'actuator_actearly': False, + 'actuator_actlimited': False, + 'actuator_actnum': False, + 'actuator_actrange': True, + 'actuator_affine_bias_gain': False, + 'actuator_biasprm': True, + 'actuator_biastype': False, + 'actuator_cranklength': False, + 'actuator_ctrllimited': False, + 'actuator_ctrlrange': True, + 'actuator_dynprm': True, + 'actuator_dyntype': False, + 'actuator_forcelimited': False, + 'actuator_forcerange': True, + 'actuator_gainprm': True, + 'actuator_gaintype': False, + 'actuator_gear': True, + 'actuator_lengthrange': False, + 'actuator_moment_tiles_nu': False, + 'actuator_moment_tiles_nv': False, + 'actuator_trnid': False, + 'actuator_trntype': False, + 'actuator_trntype_body_adr': False, + 'block_dim__actuator_velocity': False, + 'block_dim__cholesky_factorize': False, + 'block_dim__cholesky_factorize_solve': False, + 'block_dim__cholesky_solve': False, + 'block_dim__energy_vel_kinetic': False, + 'block_dim__euler_dense': False, + 'block_dim__mul_m_dense': False, + 'block_dim__qderiv_actuator_passive_actuation': False, + 'block_dim__qderiv_actuator_passive_no_actuation': False, + 'block_dim__qfrc_actuator': False, + 'block_dim__ray': False, + 'block_dim__segmented_sort': False, + 'block_dim__tendon_velocity': False, + 'block_dim__update_gradient_cholesky': False, + 'body_conaffinity': False, + 'body_contype': False, + 'body_dofadr': False, + 'body_dofnum': False, + 'body_geomadr': False, + 'body_geomnum': False, + 'body_gravcomp': True, + 'body_inertia': True, + 'body_invweight0': True, + 'body_ipos': True, + 'body_iquat': True, + 'body_jntadr': False, + 'body_jntnum': False, + 'body_mass': True, + 'body_mocapid': False, + 'body_parentid': False, + 'body_pos': True, + 'body_quat': True, + 'body_rootid': False, + 'body_subtreemass': True, + 'body_tree': False, + 'body_weldid': False, + 'cam_bodyid': False, + 'cam_fovy': False, + 'cam_intrinsic': False, + 'cam_mat0': True, + 'cam_mode': False, + 'cam_pos': True, + 'cam_pos0': True, + 'cam_poscom0': True, + 'cam_quat': True, + 'cam_resolution': False, + 'cam_sensorsize': False, + 'cam_targetbodyid': False, + 'condim_max': False, + 'dof_Madr': False, + 'dof_armature': True, + 'dof_bodyid': False, + 'dof_damping': True, + 'dof_frictionloss': True, + 'dof_invweight0': True, + 'dof_jntid': False, + 'dof_parentid': False, + 'dof_solimp': True, + 'dof_solref': True, + 'dof_tri_col': False, + 'dof_tri_row': False, + 'eq_active0': False, + 'eq_connect_adr': False, + 'eq_data': True, + 'eq_jnt_adr': False, + 'eq_obj1id': False, + 'eq_obj2id': False, + 'eq_objtype': False, + 'eq_solimp': True, + 'eq_solref': True, + 'eq_ten_adr': False, + 'eq_type': False, + 'eq_wld_adr': False, + 'exclude_signature': False, + 'flex_bending': False, + 'flex_damping': False, + 'flex_dim': False, + 'flex_edge': False, + 'flex_edgeadr': False, + 'flex_edgeflap': False, + 'flex_elem': False, + 'flex_elemedge': False, + 'flex_elemedgeadr': False, + 'flex_stiffness': False, + 'flex_vertadr': False, + 'flex_vertbodyid': False, + 'flex_vertnum': False, + 'flexedge_length0': False, + 'geom_aabb': False, + 'geom_bodyid': False, + 'geom_conaffinity': False, + 'geom_condim': False, + 'geom_contype': False, + 'geom_dataid': False, + 'geom_friction': True, + 'geom_gap': True, + 'geom_group': False, + 'geom_margin': True, + 'geom_matid': True, + 'geom_pair_type_count': False, + 'geom_plugin_index': False, + 'geom_pos': True, + 'geom_priority': False, + 'geom_quat': True, + 'geom_rbound': True, + 'geom_rgba': True, + 'geom_size': True, + 'geom_solimp': True, + 'geom_solmix': True, + 'geom_solref': True, + 'geom_type': False, + 'geompair2hfgeompair': False, + 'has_sdf_geom': False, + 'hfield_adr': False, + 'hfield_data': False, + 'hfield_ncol': False, + 'hfield_nrow': False, + 'hfield_size': False, + 'jnt_actfrclimited': False, + 'jnt_actfrcrange': True, + 'jnt_actgravcomp': False, + 'jnt_axis': True, + 'jnt_bodyid': False, + 'jnt_dofadr': False, + 'jnt_limited': False, + 'jnt_limited_ball_adr': False, + 'jnt_limited_slide_hinge_adr': False, + 'jnt_margin': True, + 'jnt_pos': True, + 'jnt_qposadr': False, + 'jnt_range': True, + 'jnt_solimp': True, + 'jnt_solref': True, + 'jnt_stiffness': True, + 'jnt_type': False, + 'light_bodyid': False, + 'light_dir': True, + 'light_dir0': True, + 'light_mode': False, + 'light_pos': True, + 'light_pos0': True, + 'light_poscom0': True, + 'light_targetbodyid': False, + 'mapM2M': False, + 'mat_rgba': True, + 'mesh_face': False, + 'mesh_faceadr': False, + 'mesh_graph': False, + 'mesh_graphadr': False, + 'mesh_polyadr': False, + 'mesh_polymap': False, + 'mesh_polymapadr': False, + 'mesh_polymapnum': False, + 'mesh_polynormal': False, + 'mesh_polynum': False, + 'mesh_polyvert': False, + 'mesh_polyvertadr': False, + 'mesh_polyvertnum': False, + 'mesh_vert': False, + 'mesh_vertadr': False, + 'mesh_vertnum': False, + 'mocap_bodyid': False, + 'nC': False, + 'nM': False, + 'na': False, + 'nbody': False, + 'ncam': False, + 'neq': False, + 'nexclude': False, + 'nflex': False, + 'nflexedge': False, + 'nflexelem': False, + 'nflexelemdata': False, + 'nflexvert': False, + 'ngeom': False, + 'ngravcomp': False, + 'nhfield': False, + 'nhfielddata': False, + 'njnt': False, + 'nlight': False, + 'nlsp': False, + 'nmeshface': False, + 'nmeshgraph': False, + 'nmeshpoly': False, + 'nmeshpolymap': False, + 'nmeshpolyvert': False, + 'nmeshvert': False, + 'nmocap': False, + 'npair': False, + 'nq': False, + 'nsensor': False, + 'nsensordata': False, + 'nsite': False, + 'ntendon': False, + 'nu': False, + 'nv': False, + 'nwrap': False, + 'nxn_geom_pair': False, + 'nxn_geom_pair_filtered': False, + 'nxn_pairid': False, + 'nxn_pairid_filtered': False, + 'opt__broadphase': False, + 'opt__broadphase_filter': False, + 'opt__cone': False, + 'opt__density': True, + 'opt__disableflags': False, + 'opt__enableflags': False, + 'opt__epa_iterations': False, + 'opt__gjk_iterations': False, + 'opt__graph_conditional': False, + 'opt__gravity': True, + 'opt__has_fluid': False, + 'opt__impratio': True, + 'opt__integrator': False, + 'opt__is_sparse': False, + 'opt__iterations': False, + 'opt__ls_iterations': False, + 'opt__ls_parallel': False, + 'opt__ls_tolerance': True, + 'opt__magnetic': True, + 'opt__run_collision_detection': False, + 'opt__sdf_initpoints': False, + 'opt__sdf_iterations': False, + 'opt__solver': False, + 'opt__timestep': True, + 'opt__tolerance': True, + 'opt__viscosity': True, + 'opt__wind': True, + 'pair_dim': False, + 'pair_friction': True, + 'pair_gap': True, + 'pair_geom1': False, + 'pair_geom2': False, + 'pair_margin': True, + 'pair_solimp': True, + 'pair_solref': True, + 'pair_solreffriction': True, + 'plugin': False, + 'plugin_attr': False, + 'qLD_updates': False, + 'qM_fullm_i': False, + 'qM_fullm_j': False, + 'qM_madr_ij': False, + 'qM_mulm_i': False, + 'qM_mulm_j': False, + 'qM_tiles': False, + 'qpos0': True, + 'qpos_spring': True, + 'rangefinder_sensor_adr': False, + 'sensor_acc_adr': False, + 'sensor_adr': False, + 'sensor_cutoff': False, + 'sensor_datatype': False, + 'sensor_dim': False, + 'sensor_e_kinetic': False, + 'sensor_e_potential': False, + 'sensor_limitfrc_adr': False, + 'sensor_limitpos_adr': False, + 'sensor_limitvel_adr': False, + 'sensor_objid': False, + 'sensor_objtype': False, + 'sensor_pos_adr': False, + 'sensor_rangefinder_adr': False, + 'sensor_rangefinder_bodyid': False, + 'sensor_refid': False, + 'sensor_reftype': False, + 'sensor_rne_postconstraint': False, + 'sensor_subtree_vel': False, + 'sensor_tendonactfrc_adr': False, + 'sensor_touch_adr': False, + 'sensor_type': False, + 'sensor_vel_adr': False, + 'site_bodyid': False, + 'site_pos': True, + 'site_quat': True, + 'site_size': False, + 'site_type': False, + 'stat__meaninertia': False, + 'subtree_mass': True, + 'ten_wrapadr_site': False, + 'ten_wrapnum_site': False, + 'tendon_actfrclimited': False, + 'tendon_actfrcrange': True, + 'tendon_adr': False, + 'tendon_armature': True, + 'tendon_damping': True, + 'tendon_frictionloss': True, + 'tendon_geom_adr': False, + 'tendon_invweight0': True, + 'tendon_jnt_adr': False, + 'tendon_length0': True, + 'tendon_lengthspring': True, + 'tendon_limited': False, + 'tendon_limited_adr': False, + 'tendon_margin': True, + 'tendon_num': False, + 'tendon_range': True, + 'tendon_site_pair_adr': False, + 'tendon_solimp_fri': True, + 'tendon_solimp_lim': True, + 'tendon_solref_fri': True, + 'tendon_solref_lim': True, + 'tendon_stiffness': True, + 'wrap_geom_adr': False, + 'wrap_jnt_adr': False, + 'wrap_objid': False, + 'wrap_prm': False, + 'wrap_pulley_scale': False, + 'wrap_site_adr': False, + 'wrap_site_pair_adr': False, + 'wrap_type': False, + }, + 'Option': { + 'broadphase': False, + 'broadphase_filter': False, + 'cone': False, + 'density': True, + 'disableflags': False, + 'enableflags': False, + 'epa_iterations': False, + 'gjk_iterations': False, + 'graph_conditional': False, + 'gravity': True, + 'has_fluid': False, + 'impratio': True, + 'integrator': False, + 'is_sparse': False, + 'iterations': False, + 'ls_iterations': False, + 'ls_parallel': False, + 'ls_tolerance': True, + 'magnetic': True, + 'run_collision_detection': False, + 'sdf_initpoints': False, + 'sdf_iterations': False, + 'solver': False, + 'timestep': True, + 'tolerance': True, + 'viscosity': True, + 'wind': True, + }, + 'Statistic': {'meaninertia': False}, +} diff --git a/mjx/pyproject.toml b/mjx/pyproject.toml index c80bcf1d..38a21389 100644 --- a/mjx/pyproject.toml +++ b/mjx/pyproject.toml @@ -33,6 +33,8 @@ dependencies = [ "mujoco>=3.3.5.dev0", "scipy", "trimesh", + "warp-lang==1.8.0; python_version >= '3.12' and sys_platform == 'linux'", + "warp-lang==1.8.0; python_version >= '3.12' and sys_platform == 'win32'", ] [project.scripts] @@ -65,3 +67,8 @@ pyink-use-majority-quotes = true extend-exclude = '''( .ipynb$ )''' + +[tool.pytest.ini_options] +norecursedirs = [ + "**/third_party", +] diff --git a/mjx/requirements.txt b/mjx/requirements.txt index 9b8cf8c1..9340b6e2 100644 --- a/mjx/requirements.txt +++ b/mjx/requirements.txt @@ -4,31 +4,51 @@ etils[epath]==1.10.0; python_version >= '3.10' \ --hash=sha256:0777fe60a234b4c65ca53470fc64f2dd2d0c6bca7fcc623fdaa8d7fa5a317098 etils[epath]==1.5.2; python_version == '3.9' \ --hash=sha256:6dc882d355e1e98a5d1a148d6323679dc47c9a5792939b9de72615aa4737eb0b -jax==0.4.34; python_version >= '3.10' \ - --hash=sha256:b957ca1fc91f7343f91a186af9f19c7f342c946f95a8c11c7f1e5cdfe2e58d9e +jax==0.5.3; python_version >= '3.10' and (sys_platform != 'darwin' or platform_machine != 'x86_64') \ + --hash=sha256:1483dc237b4f47e41755d69429e8c3c138736716147cd43bb2b99b259d4e3c41 \ + --hash=sha256:f17fcb0fd61dc289394af6ce4de2dada2312f2689bb0d73642c6f026a95fbb2c +jax==0.4.38; python_version >= '3.10' and sys_platform == 'darwin' and platform_machine == 'x86_64' \ + --hash=sha256:78987306f7041ea8500d99df1a17c33ed92620c2268c4c3677fb24e06712be64 jax==0.4.30; python_version == '3.9' \ --hash=sha256:289b30ae03b52f7f4baf6ef082a9f4e3e29c1080e22d13512c5ecf02d5f1a55b -jaxlib==0.4.34; python_version >= '3.10' \ - --hash=sha256:6b43a974c5d91a19912d138f2658dd8dbb7d30dcdff5c961d896c673e872b611 \ - --hash=sha256:87f25a477cd279840e53718403f97092eba0e8a945fcab47bcf435b6f9119dda \ - --hash=sha256:7be673a876ebd1aef440fb7e3ebaf99a91abeb550c9728c644b7d7c7b5d7c108 \ - --hash=sha256:c303f5acaf6c56ce5ff133a923c9b6247bdebedde15bd2c893c24be4d8f71306 \ - --hash=sha256:72e22e99a5dc890a64443c3fc12f13f20091f578c405a76de077ba42b4c62cd7 \ - --hash=sha256:901cb4040ed24eae40071d8114ea8d10dff436277fa74a1a5b9e7206f641151c \ - --hash=sha256:48272e9034ff868d4328cf0055a07882fd2be93f59dfb6283af7de491f9d1290 \ - --hash=sha256:1a30771d85fa77f9ab8f18e63240f455ab3a3f87660ed7b8d5eea6ceecbe5c1e \ - --hash=sha256:096f0ca309d41fa692a9d1f2f9baab1c5c8ca0749876ebb3f748e738a27c7ff4 \ - --hash=sha256:c7b3e724a30426a856070aba0192b5d199e95b4411070e7ad96ad8b196877b10 \ - --hash=sha256:133070d4fec5525ffea4dc72956398c1cf647a04dcb37f8a935ee82af78d9965 \ - --hash=sha256:3bcfa639ca3cfaf86c8ceebd5fc0d47300fd98a078014a1d0cc03133e1523d5f \ - --hash=sha256:571ef03259835458111596a71a2f4a6fabf4ec34595df4cea555035362ac5bf0 \ - --hash=sha256:c9d3adcae43a33aad4332be9c2aedc5ef751d1e755f917a5afb30c7872eacaa8 \ - --hash=sha256:8ee3f93836e53c86556ccd9449a4ea43516ee05184d031a71dd692e81259f7d9 \ - --hash=sha256:b0001c8f0e2b1c7bc99e4f314b524a340d25653505c1a1484d4041a9d3617f6f \ - --hash=sha256:d840e64b85f8865404d6d225b9bb340e158df1457152a361b05680e24792b232 \ - --hash=sha256:3e60bc826933082e99b19b87c21818a8d26fcdb01f418d47cedff554746fd6cc \ - --hash=sha256:45d719a2ce0ebf21255a277b71d756f3609b7b5be70cddc5d88fd58c35219de0 \ - --hash=sha256:b7a212a3cb5c6acc201c32ae4f4b5f5a9ac09457fbb77ba8db5ce7e7d4adc214 +jaxlib==0.5.3; python_version >= '3.10' and (sys_platform != 'darwin' or platform_machine != 'x86_64') \ + --hash=sha256:48ff5c89fb8a0fe04d475e9ddc074b4879a91d7ab68a51cec5cd1e87f81e6c47 \ + --hash=sha256:972400db4af6e85270d81db5e6e620d31395f0472e510c50dfcd4cb3f72b7220 \ + --hash=sha256:52be6c9775aff738a61170d8c047505c75bb799a45518e66a7a0908127b11785 \ + --hash=sha256:b41a6fcaeb374fabc4ee7e74cfed60843bdab607cd54f60a68b7f7655cde2b66 \ + --hash=sha256:b62bd8b29e5a4f9bfaa57c8daf6e04820b2c994f448f3dec602d64255545e9f2 \ + --hash=sha256:a4666f81d72c060ed3e581ded116a9caa9b0a70a148a54cb12a1d3afca3624b5 \ + --hash=sha256:29e1530fc81833216f1e28b578d0c59697654f72ee31c7a44ed7753baf5ac466 \ + --hash=sha256:8eb54e38d789557579f900ea3d70f104a440f8555a9681ed45f4a122dcbfd92e \ + --hash=sha256:d394dbde4a1c6bd67501cfb29d3819a10b900cb534cc0fc603319f7092f24cfa \ + --hash=sha256:bddf6360377aa1c792e47fd87f307c342e331e5ff3582f940b1bca00f6b4bc73 \ + --hash=sha256:5a5e88ab1cd6fdf78d69abe3544e8f09cce200dd339bb85fbe3c2ea67f2a5e68 \ + --hash=sha256:520665929649f29f7d948d4070dbaf3e032a4c1f7c11f2863eac73320fcee784 \ + --hash=sha256:31321c25282a06a6dfc940507bc14d0a0ac838d8ced6c07aa00a7fae34ce7b3f \ + --hash=sha256:e904b92dedfbc7e545725a8d7676987030ae9c069001d94701bc109c6dab4100 \ + --hash=sha256:bb7593cb7fffcb13963f22fa5229ed960b8fb4ae5ec3b0820048cbd67f1e8e31 \ + --hash=sha256:8019f73a10b1290f988dd3768c684f3a8a147239091c3b790ce7e47e3bbc00bd +jaxlib==0.4.38; python_version >= '3.10' and sys_platform == 'darwin' and platform_machine == 'x86_64' \ + --hash=sha256:55c19b9d3f33a6fc59f644aa5a21fba02639ccdd776cb4a9b5526625f57839ff \ + --hash=sha256:30b2f52cb50d74734af2f477c2533a7a583e3bb7b2c8acdeb361ee77d940577a \ + --hash=sha256:ee19c163a8fdf0839d4c18b88a5fbfb4e731ba7c437416d3e5483e570bb764e4 \ + --hash=sha256:61aeccb9a27c67fdb8450f6357240019cd4511cb9d62a44e4764756d384853ad \ + --hash=sha256:d6ab745a89d0fb737a36fe1d8b86659e3fffe6ee8303b20651b26193d5edc0ef \ + --hash=sha256:b67fdeabd6dfed08b7768f3bdffb521160085f8305669bd197beef61d08de08b \ + --hash=sha256:3fb0eaae7369157afecbead50aaf29e73ffddfa77a2335d721bd9794f3c510e4 \ + --hash=sha256:43db58c4c427627296366a56c10318e1f00f503690e17f94bb4344293e1995e0 \ + --hash=sha256:2751ff7037d6a997d0be0e77cc4be381c5a9f9bb8b314edb755c13a6fd969f45 \ + --hash=sha256:35226968fc9de6873d1571670eac4117f5ed80e955f7a1775204d1044abe16c6 \ + --hash=sha256:3fefea985f0415816f3bbafd3f03a437050275ef9bac9a72c1314e1644ac57c1 \ + --hash=sha256:f33bcafe32c97a562ecf6894d7c41674c80c0acdedfa5423d49af51147149874 \ + --hash=sha256:496f45b0e001a2341309cd0c74af0b670537dced79c168cb230cfcc773f0aa86 \ + --hash=sha256:dad6c0a96567c06d083c0469fec40f201210b099365bd698be31a6d2ec88fd59 \ + --hash=sha256:966cdec36cfa978f5b4582bcb4147fe511725b94c1a752dac3a5f52ce46b6fa3 \ + --hash=sha256:41e55ae5818a882e5789e848f6f16687ac132bcfbb5a5fa114a5d18b78d05f2d \ + --hash=sha256:6fe326b8af366387dd47ccf312583b2b17fed12712c9b74a648b18a13cbdbabf \ + --hash=sha256:248cca3771ebf24b070f49701364ceada33e6139445b06c782cca5ac5ad92bf4 \ + --hash=sha256:2ce77ba8cda9259a4bca97afc1c722e4291a6c463a63f8d372c6edc85117d625 \ + --hash=sha256:4103db0b3a38a5dc132741237453c24d8547290a22079ba1b577d6c88c95300a jaxlib==0.4.30; python_version == '3.9' \ --hash=sha256:54987e97a22db70f3829b437b9329e4799d653634bacc8b398554d3b90c76b2a \ --hash=sha256:f74a6b0e09df4b5e2ee399ebb9f0e01190e26e84ccb0a758fadb516415c07f18 \ @@ -83,6 +103,12 @@ trimesh==4.5.2 \ --hash=sha256:2e50f3a7fd135c3045da887a1b9f91230528f3ce11d2ec1ba44750d82d6b4f73 wheel==0.45.0 \ --hash=sha256:52f0baa5e6522155090a09c6bd95718cc46956d1b51d537ea5454249edb671c7 +warp-lang==1.8.0; python_version >= '3.12' \ + --hash=sha256:75a88d2795596f06fcf79eead94e2f194a6195dadbafbd4c9b7c8b4b05456cc4 \ + --hash=sha256:1be62e7b3e8019ccccc00916c2798afc6c4bfce17d1d6025dc8ef5c790e42c1d \ + --hash=sha256:373464bee59be37018d134b5924bf8fdb33d1313f484fd1b7c64296971c26ae2 \ + --hash=sha256:0ecf3b07c1d6d16592ab0318c86f452e16fb3337cfe327bcb76d793a98226dd8 \ + --hash=sha256:2c7627bee127b522551e02f44c35d7d6a4e5632ff4b58199208012e7f2f6c608 # Transitive dependencies of etils[epath] fsspec==2024.10.0 \