diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py index 40653b62..44224545 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py @@ -28,6 +28,7 @@ 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.bvh import refit_bvh as refit_bvh 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 @@ -104,6 +105,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import GainType as GainType from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType as GeomType from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType as IntegratorType from mujoco.mjx.third_party.mujoco_warp._src.types import JointType as JointType +from mujoco.mjx.third_party.mujoco_warp._src.types import ObjType as ObjType from mujoco.mjx.third_party.mujoco_warp._src.types import Option as Option from mujoco.mjx.third_party.mujoco_warp._src.types import RenderContext as RenderContext from mujoco.mjx.third_party.mujoco_warp._src.types import SolverType as SolverType diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/benchmark.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/benchmark.py deleted file mode 100644 index 293aca63..00000000 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/benchmark.py +++ /dev/null @@ -1,171 +0,0 @@ -# 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 benchmarking MuJoCo Warp.""" - -import time -from typing import Callable, Tuple - -import numpy as np -import warp as wp - -from mujoco.mjx.third_party.mujoco_warp._src import warp_util -from mujoco.mjx.third_party.mujoco_warp._src.types import Data -from mujoco.mjx.third_party.mujoco_warp._src.types import Model -from mujoco.mjx.third_party.mujoco_warp._src.types import RenderContext -from mujoco.mjx.third_party.mujoco_warp._src.util_misc import halton - - -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: - opt_timestep: wp.array[float], - actuator_ctrllimited: wp.array[bool], - actuator_ctrlrange: wp.array2d[wp.vec2], - # Data in: - ctrl_in: wp.array2d[float], - # In: - ctrl_center: wp.array[float], - step: int, - ctrlnoisestd: float, - ctrlnoiserate: float, - # Data out: - ctrl_out: wp.array2d[float], -): - worldid, actid = wp.tid() - - # convert rate and scale to discrete time (Ornstein-Uhlenbeck) - rate = wp.exp(-opt_timestep[worldid % opt_timestep.shape[0]] / ctrlnoiserate) - scale = ctrlnoisestd * wp.sqrt(1.0 - rate * rate) - - midpoint = 0.0 - halfrange = 1.0 - ctrlrange = actuator_ctrlrange[worldid % actuator_ctrlrange.shape[0], actid] - is_limited = actuator_ctrllimited[actid] - if is_limited: - midpoint = 0.5 * (ctrlrange[1] + ctrlrange[0]) - halfrange = 0.5 * (ctrlrange[1] - ctrlrange[0]) - if ctrl_center.shape[0] > 0: - midpoint = ctrl_center[actid] - - # exponential convergence to midpoint at ctrlnoiserate - ctrl = rate * ctrl_in[worldid, actid] + (1.0 - rate) * midpoint - - # add noise - ctrl += scale * halfrange * (2.0 * halton((step + 1) * (worldid + 1), actid + 2) - 1.0) - - # clip to range if limited - if is_limited: - ctrl = wp.clamp(ctrl, ctrlrange[0], ctrlrange[1]) - - ctrl_out[worldid, actid] = ctrl - - -def benchmark( - fn: Callable[[Model, Data], None], - m: Model, - d: Data, - nstep: int, - ctrls: np.ndarray | None = None, - event_trace: bool = False, - measure_alloc: bool = False, - measure_solver_niter: bool = False, - render_context: RenderContext | None = None, -) -> Tuple[float, float, dict, list, list, list, int]: - """Benchmark a function of Model and Data. - - Args: - fn: Function to benchmark. - m: The model containing kinematic and dynamic information (device). - d: The data object containing the current state and output information (device). - nstep: Number of timesteps. - ctrls: Control sequence to apply during benchmarking. - event_trace: If True, time routines decorated with @event_scope. - measure_alloc: If True, record number of contacts and constraints. - measure_solver_niter: If True, record the number of solver iterations. - render_context: The render context to use for rendering. - - Returns: - - Time to JIT fn. - - Total time to run the benchmark. - - Trace. - - Number of contacts. - - Number of constraints. - - Number of solver iterations. - - Number of converged worlds. - """ - trace = {} - nacon, nefc, solver_niter = [], [], [] - center = wp.array([], dtype=wp.float32) - - with warp_util.EventTracer(enabled=event_trace) as tracer: - # capture the whole function as a CUDA graph - jit_beg = time.perf_counter() - - if render_context is not None: - with wp.ScopedCapture() as capture: - fn(m, d, render_context) - else: - with wp.ScopedCapture() as capture: - fn(m, d) - - jit_end = time.perf_counter() - jit_duration = jit_end - jit_beg - - graph = capture.graph - - time_vec = np.zeros(nstep) - for i in range(nstep): - with wp.ScopedStream(wp.get_stream()): - if ctrls is not None: - center = wp.array(ctrls[i], dtype=wp.float32) - wp.launch( - ctrl_noise, - dim=(d.nworld, m.nu), - inputs=[m.opt.timestep, m.actuator_ctrllimited, m.actuator_ctrlrange, d.ctrl, center, i, 0.01, 0.1], - outputs=[d.ctrl], - ) - wp.synchronize() - - 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: - nacon.append(np.max([d.nacon.numpy()[0], d.ncollision.numpy()[0]])) - nefc.append(np.max(d.nefc.numpy())) - if measure_solver_niter: - solver_niter.append(d.solver_niter.numpy()) - - nsuccess = np.sum(~np.any(np.isnan(d.qpos.numpy()), axis=1)) - run_duration = np.sum(time_vec) - - return jit_duration, run_duration, trace, nacon, nefc, solver_niter, nsuccess 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 index 628ad5a0..855f9686 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py @@ -45,8 +45,8 @@ def create_blocked_cholesky_func(block_size: int): wp.tile_matmul(L_block, wp.tile_transpose(L_block), A_kk_tile, alpha=-1.0) # 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)) + wp.tile_cholesky_inplace(A_kk_tile) + wp.tile_store(L, A_kk_tile, offset=(k, k)) # Process the blocks below the current block for i in range(end, matrix_size, block_size): @@ -57,7 +57,7 @@ def create_blocked_cholesky_func(block_size: int): L_2_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(k, j), storage="shared") wp.tile_matmul(L_tile, wp.tile_transpose(L_2_tile), A_ik_tile, alpha=-1.0) - wp.tile_lower_solve_inplace(L_kk_tile, wp.tile_transpose(A_ik_tile)) + wp.tile_lower_solve_inplace(A_kk_tile, wp.tile_transpose(A_ik_tile)) wp.tile_store(L, A_ik_tile, offset=(i, k)) return blocked_cholesky_func @@ -98,11 +98,12 @@ def create_blocked_cholesky_solve_func(block_size: int, matrix_size_static: int) tmp_tile = wp.tile_view(rhs_tile, shape=(block_size, 1), offset=(i, 0)) for j in range(i_end, matrix_size, block_size): L_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(j, i), storage="shared") - x_tile = wp.tile_load(x, shape=(block_size, 1), offset=(j, 0), storage="shared", bounds_check=False) + x_tile = wp.tile_view(rhs_tile, shape=(block_size, 1), offset=(j, 0)) wp.tile_matmul(wp.tile_transpose(L_tile), x_tile, tmp_tile, alpha=-1.0) L_tile = wp.tile_load(L, shape=(block_size, block_size), offset=(i, i), storage="shared") wp.tile_upper_solve_inplace(wp.tile_transpose(L_tile), tmp_tile) - wp.tile_store(x, tmp_tile, offset=(i, 0), bounds_check=False) + + wp.tile_store(x, rhs_tile, offset=(0, 0), bounds_check=False) return blocked_cholesky_solve_func diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/bvh.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/bvh.py index 32899e07..37823711 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/bvh.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/bvh.py @@ -245,6 +245,19 @@ def compute_bvh_group_roots( group_root_out[tid] = root +# Warp exposes mesh group-root lookup as a kernel builtin in this version. +@wp.kernel +def compute_mesh_group_roots( + # In: + mesh_id: wp.uint64, + # Out: + group_root_out: wp.array[int], +): + tid = wp.tid() + root = wp.mesh_get_group_root(mesh_id, tid) + group_root_out[tid] = root + + @wp.kernel def _compute_flex_bvh_bounds( # Model: @@ -1083,7 +1096,7 @@ def build_flex_bvh( group_root = wp.empty(nworld, dtype=int) wp.launch( - kernel=compute_bvh_group_roots, + kernel=compute_mesh_group_roots, dim=nworld, inputs=[flex_mesh.id], outputs=[group_root], diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/cli.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/cli.py new file mode 100644 index 00000000..912df58f --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/cli.py @@ -0,0 +1,226 @@ +# Copyright 2026 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. +# ============================================================================== + +"""Shared utilities and flags for MJWarp CLI tools.""" + +import time +from typing import Callable, Tuple, get_type_hints + +import mujoco +import numpy as np +import warp as wp +from absl import app +from absl import flags +from etils import epath + +import mujoco.mjx.third_party.mujoco_warp as mjw +from mujoco.mjx.third_party.mujoco_warp._src import warp_util +from mujoco.mjx.third_party.mujoco_warp._src.io import find_keys +from mujoco.mjx.third_party.mujoco_warp._src.io import make_trajectory +from mujoco.mjx.third_party.mujoco_warp._src.io import override_model +from mujoco.mjx.third_party.mujoco_warp._src.util_misc import halton + +# shared flags for cli tool +NWORLD = flags.DEFINE_integer("nworld", 8192, "number of parallel rollouts") +NSTEP = flags.DEFINE_integer("nstep", 1000, "number of steps per rollout") +NCONMAX = flags.DEFINE_integer("nconmax", None, "override maximum number of contacts per world") +NJMAX = flags.DEFINE_integer("njmax", None, "override maximum number of constraints per world") +NJMAX_NNZ = flags.DEFINE_integer("njmax_nnz", None, "override maximum number of non-zeros in constraint Jacobian") +NCCDMAX = flags.DEFINE_integer("nccdmax", None, "override maximum number of CCD contacts per world") +OVERRIDE = flags.DEFINE_multi_string("override", [], "Model overrides (notation: foo.bar = baz)", short_name="o") +KEYFRAME = flags.DEFINE_integer("keyframe", 0, "keyframe to initialize simulation.") +EVENT_TRACE = flags.DEFINE_bool("event_trace", False, "print an event trace report") +NOISE_STD = flags.DEFINE_float("noise_std", 0.01, "add noise to ctrl signal (standard deviation)") +NOISE_RATE = flags.DEFINE_float("noise_rate", 0.1, "add noise to ctrl signal (noise rate)") + +DEVICE = flags.DEFINE_string("device", None, "override the default Warp device") +REPLAY = flags.DEFINE_string("replay", None, "keyframe sequence to replay, keyframe name must prefix match") + +RENDER_WIDTH = flags.DEFINE_integer("render_width", 64, "render width (pixels)") +RENDER_HEIGHT = flags.DEFINE_integer("render_height", 64, "render height (pixels)") +RENDER_RGB = flags.DEFINE_bool("render_rgb", True, "render RGB image") +RENDER_DEPTH = flags.DEFINE_bool("render_depth", True, "render depth image") +RENDER_TEXTURES = flags.DEFINE_bool("render_textures", True, "use textures") +RENDER_SHADOWS = flags.DEFINE_bool("render_shadows", False, "use shadows") + + +def load_model(path: epath.Path) -> mujoco.MjModel: + """Load a MuJoCo model from a path, handling resources and plugins.""" + if not path.exists(): + resource_path = epath.resource_path("mjx") / "third_party/mujoco_warp" / path + if not resource_path.exists(): + raise FileNotFoundError(f"file not found: {path}\nalso tried: {resource_path}") + path = resource_path + + if path.suffix == ".mjb": + return mujoco.MjModel.from_binary_path(path.as_posix()) + + spec = mujoco.MjSpec.from_file(path.as_posix()) + if any(p.plugin_name.startswith("mujoco.sdf") for p in spec.plugins): + from mujoco.mjx.third_party.mujoco_warp.test_data.collision_sdf.utils import register_sdf_plugins as register_sdf_plugins + + register_sdf_plugins(mjw) + + mjm = spec.compile() + + if OVERRIDE.value: + override_model(mjm, OVERRIDE.value) + + return mjm + + +@wp.kernel +def _ctrl_noise( + # Model: + opt_timestep: wp.array[float], + actuator_ctrllimited: wp.array[bool], + actuator_ctrlrange: wp.array2d[wp.vec2], + # Data in: + ctrl_in: wp.array2d[float], + # In: + ctrl_center: wp.array[float], + step: int, + ctrlnoisestd: float, + ctrlnoiserate: float, + # Data out: + ctrl_out: wp.array2d[float], +): + worldid, actid = wp.tid() + + # convert rate and scale to discrete time (Ornstein-Uhlenbeck) + rate = wp.exp(-opt_timestep[worldid % opt_timestep.shape[0]] / ctrlnoiserate) + scale = ctrlnoisestd * wp.sqrt(1.0 - rate * rate) + + midpoint = 0.0 + halfrange = 1.0 + ctrlrange = actuator_ctrlrange[worldid % actuator_ctrlrange.shape[0], actid] + is_limited = actuator_ctrllimited[actid] + if is_limited: + midpoint = 0.5 * (ctrlrange[1] + ctrlrange[0]) + halfrange = 0.5 * (ctrlrange[1] - ctrlrange[0]) + if ctrl_center.shape[0] > 0: + midpoint = ctrl_center[actid] + + # exponential convergence to midpoint at ctrlnoiserate + ctrl = rate * ctrl_in[worldid, actid] + (1.0 - rate) * midpoint + + # add noise + ctrl += scale * halfrange * (2.0 * halton((step + 1) * (worldid + 1), actid + 2) - 1.0) + + # clip to range if limited + if is_limited: + ctrl = wp.clamp(ctrl, ctrlrange[0], ctrlrange[1]) + + ctrl_out[worldid, actid] = ctrl + + +def init_structs( + fn: Callable[..., None], mjm: mujoco.MjModel +) -> Tuple[mjw.Model, mjw.Data, mjw.RenderContext | None, list[np.ndarray] | None]: + """Initialize device structs.""" + mjd = mujoco.MjData(mjm) + ctrls = None + if REPLAY.value: + keys = find_keys(mjm, REPLAY.value) + if not keys: + raise app.UsageError(f"Key prefix not found: {REPLAY.value}") + ctrls = make_trajectory(mjm, keys) + mujoco.mj_resetDataKeyframe(mjm, mjd, keys[0]) + elif mjm.nkey > 0 and KEYFRAME.value > -1: + mujoco.mj_resetDataKeyframe(mjm, mjd, KEYFRAME.value) + ctrls = [mjd.ctrl.copy() for _ in range(NSTEP.value)] + + with wp.ScopedDevice(wp.get_device(DEVICE.value)): + m = mjw.put_model(mjm) + if OVERRIDE.value: + override_model(m, OVERRIDE.value) + d = mjw.put_data( + mjm, mjd, nworld=NWORLD.value, nconmax=NCONMAX.value, njmax=NJMAX.value, njmax_nnz=NJMAX_NNZ.value, nccdmax=NCCDMAX.value + ) + + if mjw.RenderContext not in get_type_hints(fn).values(): + return m, d, None, ctrls + + rc = mjw.create_render_context( + mjm, + NWORLD.value, + (RENDER_WIDTH.value, RENDER_HEIGHT.value), + RENDER_RGB.value, + RENDER_DEPTH.value, + RENDER_TEXTURES.value, + RENDER_SHADOWS.value, + ) + + return m, d, rc, ctrls + + +def unroll( + fn: Callable[..., None], + m: mjw.Model, + d: mjw.Data, + rc: mjw.RenderContext | None, + callback: Callable[[int, dict, float], None] | None = None, + ctrls: list[np.ndarray] | None = None, +) -> dict: + """Unroll a function on batched Data and return some statistics. + + Args: + fn: Function to unroll (e.g. mjw.step). + m: Model. + d: Data. + rc: Render context (optional). + callback: Optional callback called after each step with (step count, trace, latency). + ctrls: Optional control trajectory. + + Returns: + jit_duration: Time to JIT capture the function. + """ + with wp.ScopedDevice(wp.get_device(DEVICE.value)): + with warp_util.EventTracer(enabled=EVENT_TRACE.value) as tracer: + jit_beg = time.perf_counter() + with wp.ScopedCapture() as capture: + fn(*(m, d) if rc is None else (m, d, rc)) + jit_end = time.perf_counter() + + for i in range(NSTEP.value): + with wp.ScopedStream(wp.get_stream()): + if ctrls is not None: + center = wp.array(ctrls[i], dtype=wp.float32) + wp.launch( + _ctrl_noise, + dim=(d.nworld, m.nu), + inputs=[ + m.opt.timestep, + m.actuator_ctrllimited, + m.actuator_ctrlrange, + d.ctrl, + center, + i, + NOISE_STD.value, + NOISE_RATE.value, + ], + outputs=[d.ctrl], + ) + wp.synchronize() + + run_beg = time.perf_counter() + wp.capture_launch(capture.graph) + wp.synchronize() + run_end = time.perf_counter() + + if callback: + callback(i, tracer.trace(), run_end - run_beg) + + return jit_end - jit_beg diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_flex.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_flex.py index cd5c9087..af166dd8 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_flex.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_flex.py @@ -19,6 +19,7 @@ from mujoco.mjx.third_party.mujoco_warp._src import collision_primitive_core from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL 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 @@ -28,6 +29,93 @@ from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope wp.set_module_options({"enable_backward": False}) +# TODO(team): generalize into a shared contact parameter mixing function +# (mj_contactParam) that works for both geom-geom and geom-flex contacts. +@wp.func +def _mix_flex_contact_params( + # In: + a_condim: int, + a_priority: int, + a_solmix: float, + a_solref: wp.vec2, + a_solimp: vec5, + a_friction: wp.vec3, + a_gap: float, + b_condim: int, + b_priority: int, + b_solmix: float, + b_solref: wp.vec2, + b_solimp: vec5, + b_friction: wp.vec3, + b_gap: float, +): + """Mix contact parameters between geom and flex, matching mj_contactParam.""" + gap = a_gap + b_gap + + if a_priority > b_priority: + condim = a_condim + solref = a_solref + solimp = a_solimp + fri = a_friction + elif a_priority < b_priority: + condim = b_condim + solref = b_solref + solimp = b_solimp + fri = b_friction + else: + # same priority + condim = wp.max(a_condim, b_condim) + + # compute solver mix factor + if a_solmix >= MJ_MINVAL and b_solmix >= MJ_MINVAL: + mix = a_solmix / (a_solmix + b_solmix) + elif a_solmix < MJ_MINVAL and b_solmix < MJ_MINVAL: + mix = 0.5 + elif a_solmix < MJ_MINVAL: + mix = 0.0 + else: + mix = 1.0 + + # solref: mix if both standard, min if either direct + if a_solref[0] > 0.0 and b_solref[0] > 0.0: + solref = wp.vec2( + mix * a_solref[0] + (1.0 - mix) * b_solref[0], + mix * a_solref[1] + (1.0 - mix) * b_solref[1], + ) + else: + solref = wp.vec2( + wp.min(a_solref[0], b_solref[0]), + wp.min(a_solref[1], b_solref[1]), + ) + + # solimp: mix + solimp = vec5( + mix * a_solimp[0] + (1.0 - mix) * b_solimp[0], + mix * a_solimp[1] + (1.0 - mix) * b_solimp[1], + mix * a_solimp[2] + (1.0 - mix) * b_solimp[2], + mix * a_solimp[3] + (1.0 - mix) * b_solimp[3], + mix * a_solimp[4] + (1.0 - mix) * b_solimp[4], + ) + + # friction: max + fri = wp.vec3( + wp.max(a_friction[0], b_friction[0]), + wp.max(a_friction[1], b_friction[1]), + wp.max(a_friction[2], b_friction[2]), + ) + + # unpack 5D friction with MJ_MINMU floor + friction = vec5( + wp.max(MJ_MINMU, fri[0]), + wp.max(MJ_MINMU, fri[0]), + wp.max(MJ_MINMU, fri[1]), + wp.max(MJ_MINMU, fri[2]), + wp.max(MJ_MINMU, fri[2]), + ) + + return condim, gap, solref, solimp, friction + + @wp.func def _write_flex_contact( # Data in: @@ -264,13 +352,21 @@ def _flex_plane_narrowphase( nflexvert: int, geom_type: wp.array[int], geom_condim: wp.array[int], + geom_priority: wp.array[int], + geom_solmix: wp.array2d[float], geom_solref: wp.array2d[wp.vec2], geom_solimp: wp.array2d[vec5], geom_friction: wp.array2d[wp.vec3], geom_margin: wp.array2d[float], + geom_gap: wp.array2d[float], flex_condim: wp.array[int], + flex_priority: wp.array[int], + flex_solmix: wp.array[float], + flex_solref: wp.array[wp.vec2], + flex_solimp: wp.array[vec5], flex_friction: wp.array[wp.vec3], flex_margin: wp.array[float], + flex_gap: wp.array[float], flex_vertadr: wp.array[int], flex_radius: wp.array[float], flex_vertflexid: wp.array[int], @@ -303,8 +399,6 @@ def _flex_plane_narrowphase( flexid = flex_vertflexid[vertid] radius = flex_radius[flexid] flex_margin_val = flex_margin[flexid] - flex_condim_val = flex_condim[flexid] - flex_fric = flex_friction[flexid] # Convert global vertid to local vertex index within this flex local_vertid = vertid - flex_vertadr[flexid] @@ -327,20 +421,21 @@ def _flex_plane_narrowphase( dist = signed_dist - radius if dist < margin: - geom_condim_val = geom_condim[geomid] - condim = wp.max(geom_condim_val, flex_condim_val) - solref = geom_solref[worldid % geom_solref.shape[0], geomid] - solimp = geom_solimp[worldid % geom_solimp.shape[0], geomid] - geom_fric = geom_friction[worldid % geom_friction.shape[0], geomid] - fric0 = wp.max(geom_fric[0], flex_fric[0]) - fric1 = wp.max(geom_fric[1], flex_fric[1]) - fric2 = wp.max(geom_fric[2], flex_fric[2]) - friction = vec5( - wp.max(MJ_MINMU, fric0), - wp.max(MJ_MINMU, fric0), - wp.max(MJ_MINMU, fric1), - wp.max(MJ_MINMU, fric2), - wp.max(MJ_MINMU, fric2), + condim, gap, solref, solimp, friction = _mix_flex_contact_params( + geom_condim[geomid], + geom_priority[geomid], + geom_solmix[worldid % geom_solmix.shape[0], geomid], + geom_solref[worldid % geom_solref.shape[0], geomid], + geom_solimp[worldid % geom_solimp.shape[0], geomid], + geom_friction[worldid % geom_friction.shape[0], geomid], + geom_gap[worldid % geom_gap.shape[0], geomid], + flex_condim[flexid], + flex_priority[flexid], + flex_solmix[flexid], + flex_solref[flexid], + flex_solimp[flexid], + flex_friction[flexid], + flex_gap[flexid], ) contact_pos = vert - plane_normal * (dist * 0.5 + radius) @@ -349,7 +444,7 @@ def _flex_plane_narrowphase( dist, contact_pos, make_frame(plane_normal), - margin, + margin - gap, condim, friction, solref, @@ -386,14 +481,24 @@ def _flex_narrowphase_dim2( geom_contype: wp.array[int], geom_conaffinity: wp.array[int], geom_condim: wp.array[int], + geom_priority: wp.array[int], + geom_solmix: wp.array2d[float], geom_solref: wp.array2d[wp.vec2], geom_solimp: wp.array2d[vec5], geom_size: wp.array2d[wp.vec3], geom_friction: wp.array2d[wp.vec3], geom_margin: wp.array2d[float], + geom_gap: wp.array2d[float], flex_contype: wp.array[int], flex_conaffinity: wp.array[int], + flex_condim: wp.array[int], + flex_priority: wp.array[int], + flex_solmix: wp.array[float], + flex_solref: wp.array[wp.vec2], + flex_solimp: wp.array[vec5], + flex_friction: wp.array[wp.vec3], flex_margin: wp.array[float], + flex_gap: wp.array[float], flex_dim: wp.array[int], flex_vertadr: wp.array[int], flex_elemadr: wp.array[int], @@ -478,17 +583,22 @@ def _flex_narrowphase_dim2( geom_rot = geom_xmat_in[worldid, geomid] geom_size_val = geom_size[worldid % geom_size.shape[0], geomid] - condim = geom_condim[geomid] - gf = geom_friction[worldid % geom_friction.shape[0], geomid] - friction = vec5( - wp.max(MJ_MINMU, gf[0]), - wp.max(MJ_MINMU, gf[0]), - wp.max(MJ_MINMU, gf[1]), - wp.max(MJ_MINMU, gf[2]), - wp.max(MJ_MINMU, gf[2]), + condim, gap, solref, solimp, friction = _mix_flex_contact_params( + geom_condim[geomid], + geom_priority[geomid], + geom_solmix[worldid % geom_solmix.shape[0], geomid], + geom_solref[worldid % geom_solref.shape[0], geomid], + geom_solimp[worldid % geom_solimp.shape[0], geomid], + geom_friction[worldid % geom_friction.shape[0], geomid], + geom_gap[worldid % geom_gap.shape[0], geomid], + flex_condim[flexid], + flex_priority[flexid], + flex_solmix[flexid], + flex_solref[flexid], + flex_solimp[flexid], + flex_friction[flexid], + flex_gap[flexid], ) - solref = geom_solref[worldid % geom_solref.shape[0], geomid] - solimp = geom_solimp[worldid % geom_solimp.shape[0], geomid] _collide_geom_triangle( naconmax_in, @@ -537,14 +647,24 @@ def _flex_narrowphase_dim3( geom_contype: wp.array[int], geom_conaffinity: wp.array[int], geom_condim: wp.array[int], + geom_priority: wp.array[int], + geom_solmix: wp.array2d[float], geom_solref: wp.array2d[wp.vec2], geom_solimp: wp.array2d[vec5], geom_size: wp.array2d[wp.vec3], geom_friction: wp.array2d[wp.vec3], geom_margin: wp.array2d[float], + geom_gap: wp.array2d[float], flex_contype: wp.array[int], flex_conaffinity: wp.array[int], + flex_condim: wp.array[int], + flex_priority: wp.array[int], + flex_solmix: wp.array[float], + flex_solref: wp.array[wp.vec2], + flex_solimp: wp.array[vec5], + flex_friction: wp.array[wp.vec3], flex_margin: wp.array[float], + flex_gap: wp.array[float], flex_dim: wp.array[int], flex_vertadr: wp.array[int], flex_shellnum: wp.array[int], @@ -632,17 +752,22 @@ def _flex_narrowphase_dim3( geom_rot = geom_xmat_in[worldid, geomid] geom_size_val = geom_size[worldid % geom_size.shape[0], geomid] - condim = geom_condim[geomid] - gf = geom_friction[worldid % geom_friction.shape[0], geomid] - friction = vec5( - wp.max(MJ_MINMU, gf[0]), - wp.max(MJ_MINMU, gf[0]), - wp.max(MJ_MINMU, gf[1]), - wp.max(MJ_MINMU, gf[2]), - wp.max(MJ_MINMU, gf[2]), + condim, gap, solref, solimp, friction = _mix_flex_contact_params( + geom_condim[geomid], + geom_priority[geomid], + geom_solmix[worldid % geom_solmix.shape[0], geomid], + geom_solref[worldid % geom_solref.shape[0], geomid], + geom_solimp[worldid % geom_solimp.shape[0], geomid], + geom_friction[worldid % geom_friction.shape[0], geomid], + geom_gap[worldid % geom_gap.shape[0], geomid], + flex_condim[flexid], + flex_priority[flexid], + flex_solmix[flexid], + flex_solref[flexid], + flex_solimp[flexid], + flex_friction[flexid], + flex_gap[flexid], ) - solref = geom_solref[worldid % geom_solref.shape[0], geomid] - solimp = geom_solimp[worldid % geom_solimp.shape[0], geomid] _collide_geom_triangle( naconmax_in, @@ -698,14 +823,24 @@ def flex_narrowphase(m: Model, d: Data): m.geom_contype, m.geom_conaffinity, m.geom_condim, + 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.flex_contype, m.flex_conaffinity, + m.flex_condim, + m.flex_priority, + m.flex_solmix, + m.flex_solref, + m.flex_solimp, + m.flex_friction, m.flex_margin, + m.flex_gap, m.flex_dim, m.flex_vertadr, m.flex_elemadr, @@ -749,14 +884,24 @@ def flex_narrowphase(m: Model, d: Data): m.geom_contype, m.geom_conaffinity, m.geom_condim, + 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.flex_contype, m.flex_conaffinity, + m.flex_condim, + m.flex_priority, + m.flex_solmix, + m.flex_solref, + m.flex_solimp, + m.flex_friction, m.flex_margin, + m.flex_gap, m.flex_dim, m.flex_vertadr, m.flex_shellnum, @@ -797,13 +942,21 @@ def flex_narrowphase(m: Model, d: Data): m.nflexvert, m.geom_type, m.geom_condim, + m.geom_priority, + m.geom_solmix, m.geom_solref, m.geom_solimp, m.geom_friction, m.geom_margin, + m.geom_gap, m.flex_condim, + m.flex_priority, + m.flex_solmix, + m.flex_solref, + m.flex_solimp, m.flex_friction, m.flex_margin, + m.flex_gap, m.flex_vertadr, m.flex_radius, m.flex_vertflexid, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py index 788f1af5..59a7ad72 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py @@ -144,6 +144,7 @@ def _equality_connect( eq_solimp: wp.array2d[vec5], eq_data: wp.array2d[vec11], is_sparse: bool, + body_isdofancestor: wp.array2d[int], eq_connect_adr: wp.array[int], # Data in: qvel_in: wp.array2d[float], @@ -259,6 +260,7 @@ def _equality_connect( body_parentid, body_rootid, dof_bodyid, + body_isdofancestor, subtree_com_in, cdof_in, pos1, @@ -270,6 +272,7 @@ def _equality_connect( body_parentid, body_rootid, dof_bodyid, + body_isdofancestor, subtree_com_in, cdof_in, pos2, @@ -301,6 +304,7 @@ def _equality_connect( body_parentid, body_rootid, dof_bodyid, + body_isdofancestor, subtree_com_in, cdof_in, pos1, @@ -312,6 +316,7 @@ def _equality_connect( body_parentid, body_rootid, dof_bodyid, + body_isdofancestor, subtree_com_in, cdof_in, pos2, @@ -694,6 +699,7 @@ def _equality_flex(is_sparse: bool): eq_flex_adr: wp.array[int], # Data in: qvel_in: wp.array2d[float], + eq_active_in: wp.array2d[bool], flexedge_J_in: wp.array2d[float], flexedge_length_in: wp.array2d[float], njmax_in: int, @@ -718,6 +724,10 @@ def _equality_flex(is_sparse: bool): ): worldid, eqflexid, edgeid = wp.tid() eqid = eq_flex_adr[eqflexid] + + if not eq_active_in[worldid, eqid]: + return + flexid = eq_obj1id[eqid] if edgeid < flex_edgeadr[flexid] or edgeid >= flex_edgeadr[flexid] + flex_edgenum[flexid]: return @@ -813,6 +823,7 @@ def _equality_weld( eq_solimp: wp.array2d[vec5], eq_data: wp.array2d[vec11], is_sparse: bool, + body_isdofancestor: wp.array2d[int], eq_wld_adr: wp.array[int], # Data in: qvel_in: wp.array2d[float], @@ -947,6 +958,7 @@ def _equality_weld( body_parentid, body_rootid, dof_bodyid, + body_isdofancestor, subtree_com_in, cdof_in, pos1, @@ -958,6 +970,7 @@ def _equality_weld( body_parentid, body_rootid, dof_bodyid, + body_isdofancestor, subtree_com_in, cdof_in, pos2, @@ -1003,6 +1016,7 @@ def _equality_weld( body_parentid, body_rootid, dof_bodyid, + body_isdofancestor, subtree_com_in, cdof_in, pos1, @@ -1014,6 +1028,7 @@ def _equality_weld( body_parentid, body_rootid, dof_bodyid, + body_isdofancestor, subtree_com_in, cdof_in, pos2, @@ -1665,97 +1680,181 @@ def _limit_tendon( ) -@wp.kernel -def _contact_pyramidal( - # Model: - nv: int, - opt_timestep: wp.array[float], - opt_disableflags: int, - opt_impratio_invsqrt: wp.array[float], - body_parentid: wp.array[int], - body_rootid: wp.array[int], - body_weldid: wp.array[int], - body_dofnum: wp.array[int], - body_dofadr: wp.array[int], - body_invweight0: wp.array2d[wp.vec2], - dof_bodyid: wp.array[int], - dof_parentid: wp.array[int], - geom_bodyid: wp.array[int], - flex_vertadr: wp.array[int], - flex_vertbodyid: wp.array[int], - is_sparse: bool, - # Data in: - qvel_in: wp.array2d[float], - subtree_com_in: wp.array2d[wp.vec3], - cdof_in: wp.array2d[wp.spatial_vector], - njmax_in: int, - njmax_nnz_in: int, - nacon_in: wp.array[int], - # In: - dist_in: wp.array[float], - condim_in: wp.array[int], - includemargin_in: wp.array[float], - worldid_in: wp.array[int], - geom_in: wp.array[wp.vec2i], - flex_in: wp.array[wp.vec2i], - vert_in: wp.array[wp.vec2i], - pos_in: wp.array[wp.vec3], - frame_in: wp.array[wp.mat33], - friction_in: wp.array[vec5], - solref_in: wp.array[wp.vec2], - solimp_in: wp.array[vec5], - type_in: wp.array[int], - # Data out: - nefc_out: wp.array[int], - contact_efc_address_out: wp.array2d[int], - efc_type_out: wp.array2d[int], - efc_id_out: wp.array2d[int], - efc_J_rownnz_out: wp.array2d[int], - efc_J_rowadr_out: wp.array2d[int], - efc_J_colind_out: wp.array3d[int], - efc_J_out: wp.array3d[float], - efc_pos_out: wp.array2d[float], - efc_margin_out: wp.array2d[float], - efc_D_out: wp.array2d[float], - efc_vel_out: wp.array2d[float], - efc_aref_out: wp.array2d[float], - efc_frictionloss_out: wp.array2d[float], - # Out: - efc_nnz_out: wp.array[int], -): - conid, dimid = wp.tid() +def _efc_contact_init(cone_type: types.ConeType, is_sparse: bool): + IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC + IS_SPARSE = is_sparse - if conid >= nacon_in[0]: - return + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + body_weldid: wp.array[int], + body_dofnum: wp.array[int], + body_dofadr: wp.array[int], + dof_parentid: wp.array[int], + geom_bodyid: wp.array[int], + flex_vertadr: wp.array[int], + flex_vertbodyid: wp.array[int], + # Data in: + njmax_in: int, + njmax_nnz_in: int, + nacon_in: wp.array[int], + # In: + dist_in: wp.array[float], + condim_in: wp.array[int], + includemargin_in: wp.array[float], + worldid_in: wp.array[int], + geom_in: wp.array[wp.vec2i], + flex_in: wp.array[wp.vec2i], + vert_in: wp.array[wp.vec2i], + type_in: wp.array[int], + # Data out: + nefc_out: wp.array[int], + contact_efc_address_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_J_rownnz_out: wp.array2d[int], + efc_J_rowadr_out: wp.array2d[int], + # Out: + efc_nnz_out: wp.array[int], + ): + conid = wp.tid() - if not type_in[conid] & ContactType.CONSTRAINT: - 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: - contact_efc_address_out[conid, dimid] = -1 + if conid >= nacon_in[0]: return - timestep = opt_timestep[worldid % opt_timestep.shape[0]] - impratio_invsqrt = opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] - contact_efc_address_out[conid, dimid] = efcid + if not type_in[conid] & ContactType.CONSTRAINT: + return + + condim = condim_in[conid] + + includemargin = includemargin_in[conid] + pos = dist_in[conid] - includemargin + active = pos < 0 + + if not active: + return + + if wp.static(IS_ELLIPTIC): + ndim = condim + else: + if condim == 1: + ndim = 1 + else: + ndim = 2 * (condim - 1) + + worldid = worldid_in[conid] + + # Allocate contiguous block of efcids for all dimids + base_efcid = wp.atomic_add(nefc_out, worldid, ndim) + for dim in range(ndim): + efcid = base_efcid + dim + if efcid >= njmax_in: + contact_efc_address_out[conid, dim] = -1 + else: + contact_efc_address_out[conid, dim] = efcid + # This is redundant with the _efc_row call later but needed for the jac calculation + efc_id_out[worldid, efcid] = conid + + if wp.static(IS_SPARSE): + geom = geom_in[conid] + + if geom[0] >= 0: + body1 = geom_bodyid[geom[0]] + else: + flex = flex_in[conid] + vert = vert_in[conid] + body1 = flex_vertbodyid[flex_vertadr[flex[0]] + vert[0]] + + if geom[1] >= 0: + body2 = geom_bodyid[geom[1]] + else: + flex = flex_in[conid] + vert = vert_in[conid] + body2 = flex_vertbodyid[flex_vertadr[flex[1]] + vert[1]] + + # skip fixed bodies + body1 = body_weldid[body1] + body2 = body_weldid[body2] + + da1 = int(body_dofadr[body1] + body_dofnum[body1] - 1) + da2 = int(body_dofadr[body2] + body_dofnum[body2] - 1) + + # count non-zeros + rownnz = int(0) + while da1 >= 0 or da2 >= 0: + da = wp.max(da1, da2) + # skip common dofs + if da1 == da and da2 == da: + break + if da1 == da: + da1 = dof_parentid[da1] + if da2 == da: + da2 = dof_parentid[da2] + rownnz += 1 + + rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz * ndim) + if rowadr + rownnz * ndim > njmax_nnz_in: + return + for dim in range(ndim): + efcid = base_efcid + dim + if efcid < njmax_in: + efc_J_rowadr_out[worldid, efcid] = rowadr + dim * rownnz + efc_J_rownnz_out[worldid, efcid] = rownnz + + return kernel + + +def _efc_contact_jac_sparse(cone_type: types.ConeType): + IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC + + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + body_parentid: wp.array[int], + body_rootid: wp.array[int], + body_weldid: wp.array[int], + body_dofnum: wp.array[int], + body_dofadr: wp.array[int], + dof_bodyid: wp.array[int], + dof_parentid: wp.array[int], + geom_bodyid: wp.array[int], + flex_vertadr: wp.array[int], + flex_vertbodyid: wp.array[int], + body_isdofancestor: wp.array2d[int], + # Data in: + qvel_in: wp.array2d[float], + subtree_com_in: wp.array2d[wp.vec3], + cdof_in: wp.array2d[wp.spatial_vector], + contact_efc_address_in: wp.array2d[int], + efc_J_rownnz_in: wp.array2d[int], + efc_J_rowadr_in: wp.array2d[int], + nacon_in: wp.array[int], + # In: + condim_in: wp.array[int], + geom_in: wp.array[wp.vec2i], + flex_in: wp.array[wp.vec2i], + vert_in: wp.array[wp.vec2i], + pos_in: wp.array[wp.vec3], + frame_in: wp.array2d[wp.vec3], + friction_in: wp.array2d[float], + worldid_in: wp.array[int], + # Data out: + efc_J_colind_out: wp.array3d[int], + efc_J_out: wp.array3d[float], + efc_Jqvel_out: wp.array2d[float], + ): + conid, dimid = wp.tid() + + if conid >= nacon_in[0]: + return + + efcid = contact_efc_address_in[conid, dimid] + if efcid < 0: + return + + worldid = worldid_in[conid] + condim = condim_in[conid] geom = geom_in[conid] - if geom[0] >= 0: body1 = geom_bodyid[geom[0]] else: @@ -1770,75 +1869,39 @@ def _contact_pyramidal( vert = vert_in[conid] body2 = flex_vertbodyid[flex_vertadr[flex[1]] + vert[1]] - con_pos = pos_in[conid] - frame = frame_in[conid] - - # pyramidal has common invweight across all edges - body_invweight0_id = worldid % body_invweight0.shape[0] - invweight = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, 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_invsqrt * impratio_invsqrt - - Jqvel = float(0.0) - # skip fixed bodies body1 = body_weldid[body1] body2 = body_weldid[body2] + con_pos = pos_in[conid] + + if not wp.static(IS_ELLIPTIC): + frame_0 = frame_in[conid, 0] + if condim > 1: + dimid2 = dimid / 2 + 1 + frii = friction_in[conid, dimid2 - 1] + da1 = int(body_dofadr[body1] + body_dofnum[body1] - 1) da2 = int(body_dofadr[body2] + body_dofnum[body2] - 1) - - if is_sparse: - pda1 = da1 - pda2 = da2 - rownnz = int(0) - while pda1 >= 0 or pda2 >= 0: - da = wp.max(pda1, pda2) - # skip common dofs - if pda1 == da and pda2 == da: - break - if pda1 == da: - pda1 = dof_parentid[pda1] - if pda2 == da: - pda2 = dof_parentid[pda2] - rownnz += 1 - - # get rowadr - rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz) - if rowadr + rownnz > njmax_nnz_in: - return - efc_J_rowadr_out[worldid, efcid] = rowadr - efc_J_rownnz_out[worldid, efcid] = rownnz - da = wp.max(da1, da2) - if is_sparse: - nnz = int(0) - dofid = int(da) - else: - dofid = int(nv - 1) + rowadr = efc_J_rowadr_in[worldid, efcid] + rownnz = efc_J_rownnz_in[worldid, efcid] + + Jqvel = float(0.0) + nnz = int(0) + dofid = int(da) while True: - if is_sparse: - if nnz >= rownnz: - break - else: - if dofid < 0: - break + if nnz >= rownnz: + break if dofid == da: - # TODO(team): contact_jacobian jac1p, jac1r = support.jac_dof( body_parentid, body_rootid, dof_bodyid, + body_isdofancestor, subtree_com_in, cdof_in, con_pos, @@ -1850,6 +1913,7 @@ def _contact_pyramidal( body_parentid, body_rootid, dof_bodyid, + body_isdofancestor, subtree_com_in, cdof_in, con_pos, @@ -1858,37 +1922,43 @@ def _contact_pyramidal( worldid, ) - J = float(0.0) - Ji = float(0.0) - if condim > 1: - dimid2 = dimid / 2 + 1 + jacp_dif = jac2p - jac1p + jacr_dif = jac2r - jac1r - for xyz in range(3): - jacp_dif = jac2p[xyz] - jac1p[xyz] - J += frame[0, xyz] * jacp_dif + if wp.static(IS_ELLIPTIC): + J = float(0.0) + if dimid < 3: + frame_row = frame_in[conid, dimid] + for xyz in range(3): + J += frame_row[xyz] * jacp_dif[xyz] + else: + frame_row = frame_in[conid, dimid - 3] + for xyz in range(3): + J += frame_row[xyz] * jacr_dif[xyz] + else: + J = float(0.0) + Ji = float(0.0) + + for xyz in range(3): + J += frame_0[xyz] * jacp_dif[xyz] + + if condim > 1: + if dimid2 < 3: + Ji += frame_in[conid, dimid2][xyz] * jacp_dif[xyz] + else: + Ji += frame_in[conid, dimid2 - 3][xyz] * jacr_dif[xyz] if condim > 1: - if dimid2 < 3: - Ji += frame[dimid2, xyz] * jacp_dif + if dimid % 2 == 0: + J += Ji * frii else: - Ji += frame[dimid2 - 3, xyz] * (jac2r[xyz] - jac1r[xyz]) + J -= Ji * frii - if condim > 1: - if dimid % 2 == 0: - J += Ji * frii - else: - J -= Ji * frii - - if is_sparse: - sparseid = rowadr + nnz - efc_J_colind_out[worldid, 0, sparseid] = dofid - efc_J_out[worldid, 0, sparseid] = J - nnz += 1 - else: - efc_J_out[worldid, efcid, dofid] = J + sparseid = rowadr + nnz + efc_J_colind_out[worldid, 0, sparseid] = dofid + efc_J_out[worldid, 0, sparseid] = J + nnz += 1 Jqvel += J * qvel_in[worldid, dofid] - if is_sparse and nnz >= rownnz: - break # Advance tree pointers and recompute da for next iteration if da1 == da: @@ -1896,135 +1966,229 @@ def _contact_pyramidal( if da2 == da: da2 = dof_parentid[da2] da = wp.max(da1, da2) - if is_sparse: - dofid = da - else: - dofid -= 1 - else: - if not is_sparse: - efc_J_out[worldid, efcid, dofid] = 0.0 - dofid -= 1 + dofid = da - if condim == 1: - efc_type = ConstraintType.CONTACT_FRICTIONLESS - else: - efc_type = ConstraintType.CONTACT_PYRAMIDAL + efc_Jqvel_out[worldid, efcid] = Jqvel - _efc_row( - opt_disableflags, - worldid, - timestep, - 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, - ) + return kernel -@wp.kernel -def _contact_elliptic( - # Model: - nv: int, - opt_timestep: wp.array[float], - opt_disableflags: int, - opt_impratio_invsqrt: wp.array[float], - body_parentid: wp.array[int], - body_rootid: wp.array[int], - body_weldid: wp.array[int], - body_dofnum: wp.array[int], - body_dofadr: wp.array[int], - body_invweight0: wp.array2d[wp.vec2], - dof_bodyid: wp.array[int], - dof_parentid: wp.array[int], - geom_bodyid: wp.array[int], - flex_vertadr: wp.array[int], - flex_vertbodyid: wp.array[int], - is_sparse: bool, - # Data in: - qvel_in: wp.array2d[float], - subtree_com_in: wp.array2d[wp.vec3], - cdof_in: wp.array2d[wp.spatial_vector], - njmax_in: int, - njmax_nnz_in: int, - nacon_in: wp.array[int], - # In: - dist_in: wp.array[float], - condim_in: wp.array[int], - includemargin_in: wp.array[float], - worldid_in: wp.array[int], - geom_in: wp.array[wp.vec2i], - flex_in: wp.array[wp.vec2i], - vert_in: wp.array[wp.vec2i], - pos_in: wp.array[wp.vec3], - frame_in: wp.array[wp.mat33], - friction_in: wp.array[vec5], - solref_in: wp.array[wp.vec2], - solreffriction_in: wp.array[wp.vec2], - solimp_in: wp.array[vec5], - type_in: wp.array[int], - # Data out: - nefc_out: wp.array[int], - contact_efc_address_out: wp.array2d[int], - efc_type_out: wp.array2d[int], - efc_id_out: wp.array2d[int], - efc_J_rownnz_out: wp.array2d[int], - efc_J_rowadr_out: wp.array2d[int], - efc_J_colind_out: wp.array3d[int], - efc_J_out: wp.array3d[float], - efc_pos_out: wp.array2d[float], - efc_margin_out: wp.array2d[float], - efc_D_out: wp.array2d[float], - efc_vel_out: wp.array2d[float], - efc_aref_out: wp.array2d[float], - efc_frictionloss_out: wp.array2d[float], - # Out: - efc_nnz_out: wp.array[int], -): - conid, dimid = wp.tid() +def _efc_contact_jac_dense(tile_size: int, cone_type: types.ConeType): + TILE_SIZE = tile_size + IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC - if conid >= nacon_in[0]: - return + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + body_rootid: wp.array[int], + geom_bodyid: wp.array[int], + flex_vertadr: wp.array[int], + flex_vertbodyid: wp.array[int], + body_isdofancestor: wp.array2d[int], + # Data in: + ne_in: wp.array[int], + nf_in: wp.array[int], + nl_in: wp.array[int], + nefc_in: wp.array[int], + qvel_in: wp.array2d[float], + subtree_com_in: wp.array2d[wp.vec3], + cdof_in: wp.array2d[wp.spatial_vector], + contact_efc_address_in: wp.array2d[int], + efc_id_in: wp.array2d[int], + njmax_in: int, + # In: + nv_padded: int, + condim_in: wp.array[int], + geom_in: wp.array[wp.vec2i], + flex_in: wp.array[wp.vec2i], + vert_in: wp.array[wp.vec2i], + pos_in: wp.array[wp.vec3], + frame_in: wp.array2d[wp.vec3], + friction_in: wp.array2d[float], + # Data out: + efc_J_out: wp.array3d[float], + efc_Jqvel_out: wp.array2d[float], + ): + worldid, dof_block_id, tid = wp.tid() - if not type_in[conid] & ContactType.CONSTRAINT: - 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: - contact_efc_address_out[conid, dimid] = -1 + dof_start = dof_block_id * wp.static(TILE_SIZE) + if dof_start >= nv_padded: return + cdof_tile = wp.tile_load(cdof_in[worldid], shape=TILE_SIZE, offset=dof_start, bounds_check=True) + qvel_tile = wp.tile_load(qvel_in[worldid], shape=TILE_SIZE, offset=dof_start, bounds_check=True) + + efcid_start = ne_in[worldid] + nf_in[worldid] + nl_in[worldid] + efcid_end = wp.min(nefc_in[worldid], njmax_in) + + prev_conid = int(-1) + condim = int(0) + + for efcid in range(efcid_start, efcid_end): + conid = efc_id_in[worldid, efcid] + + # Recompute per-contact data only when contact changes + if conid != prev_conid: + prev_conid = conid + condim = condim_in[conid] + + geom = geom_in[conid] + if geom[0] >= 0: + body1 = geom_bodyid[geom[0]] + else: + flex = flex_in[conid] + vert = vert_in[conid] + body1 = flex_vertbodyid[flex_vertadr[flex[0]] + vert[0]] + + if geom[1] >= 0: + body2 = geom_bodyid[geom[1]] + else: + flex = flex_in[conid] + vert = vert_in[conid] + body2 = flex_vertbodyid[flex_vertadr[flex[1]] + vert[1]] + + con_pos = pos_in[conid] + offset1 = con_pos - subtree_com_in[worldid, body_rootid[body1]] + offset2 = con_pos - subtree_com_in[worldid, body_rootid[body2]] + + affects1_tile = wp.tile_load(body_isdofancestor[body1], shape=TILE_SIZE, offset=dof_start, bounds_check=True) + affects2_tile = wp.tile_load(body_isdofancestor[body2], shape=TILE_SIZE, offset=dof_start, bounds_check=True) + + jacp1_tile = wp.tile_map(support._compute_jacp, cdof_tile, offset1, affects1_tile) + jacp2_tile = wp.tile_map(support._compute_jacp, cdof_tile, offset2, affects2_tile) + jacp_dif_tile = wp.tile_map(wp.sub, jacp2_tile, jacp1_tile) + + jacr1_tile = wp.tile_map(support._compute_jacr, cdof_tile, affects1_tile) + jacr2_tile = wp.tile_map(support._compute_jacr, cdof_tile, affects2_tile) + jacr_dif_tile = wp.tile_map(wp.sub, jacr2_tile, jacr1_tile) + + if not wp.static(IS_ELLIPTIC): + frame_0 = frame_in[conid, 0] + Ji_0p_tile = wp.tile_map(wp.dot, jacp_dif_tile, frame_0) + + if condim > 1: + Ji_0r_tile = wp.tile_map(wp.dot, jacr_dif_tile, frame_0) + frame_1 = frame_in[conid, 1] + Ji_1p_tile = wp.tile_map(wp.dot, jacp_dif_tile, frame_1) + Ji_1r_tile = wp.tile_map(wp.dot, jacr_dif_tile, frame_1) + frame_2 = frame_in[conid, 2] + Ji_2p_tile = wp.tile_map(wp.dot, jacp_dif_tile, frame_2) + Ji_2r_tile = wp.tile_map(wp.dot, jacr_dif_tile, frame_2) + + if wp.static(IS_ELLIPTIC): + dimid = efcid - contact_efc_address_in[conid, 0] + if dimid < 3: + frame_idx = dimid + else: + frame_idx = dimid - 3 + + frame_row = frame_in[conid, frame_idx] + + if dimid < 3: + J_tile = wp.tile_map(wp.dot, jacp_dif_tile, frame_row) + else: + J_tile = wp.tile_map(wp.dot, jacr_dif_tile, frame_row) + else: + J_tile = Ji_0p_tile + if condim > 1: + dimid = efcid - contact_efc_address_in[conid, 0] + dimid2 = dimid / 2 + 1 + frii = friction_in[conid, dimid2 - 1] + frii_sign = frii * (1.0 - 2.0 * float(dimid & 1)) + + if dimid2 == 1: + J_tile = wp.tile_map(wp.add, J_tile, wp.tile_map(wp.mul, Ji_1p_tile, frii_sign)) + elif dimid2 == 2: + J_tile = wp.tile_map(wp.add, J_tile, wp.tile_map(wp.mul, Ji_2p_tile, frii_sign)) + elif dimid2 == 3: + J_tile = wp.tile_map(wp.add, J_tile, wp.tile_map(wp.mul, Ji_0r_tile, frii_sign)) + elif dimid2 == 4: + J_tile = wp.tile_map(wp.add, J_tile, wp.tile_map(wp.mul, Ji_1r_tile, frii_sign)) + else: + J_tile = wp.tile_map(wp.add, J_tile, wp.tile_map(wp.mul, Ji_2r_tile, frii_sign)) + + wp.tile_store(efc_J_out[worldid, efcid], J_tile, offset=dof_start, bounds_check=True) + + Jqvel_tile = wp.tile_map(wp.mul, J_tile, qvel_tile) + Jqvel_sum = wp.tile_reduce(wp.add, Jqvel_tile) + if tid == 0: + wp.atomic_add(efc_Jqvel_out[worldid], efcid, wp.tile_extract(Jqvel_sum, 0)) + + return kernel + + +def _efc_contact_update(cone_type: types.ConeType): + IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC + + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + opt_timestep: wp.array[float], + opt_disableflags: int, + opt_impratio_invsqrt: wp.array[float], + body_invweight0: wp.array2d[wp.vec2], + geom_bodyid: wp.array[int], + flex_vertadr: wp.array[int], + flex_vertbodyid: wp.array[int], + # Data in: + contact_efc_address_in: wp.array2d[int], + efc_Jqvel_in: wp.array2d[float], + nacon_in: wp.array[int], + # In: + dist_in: wp.array[float], + condim_in: wp.array[int], + includemargin_in: wp.array[float], + worldid_in: wp.array[int], + geom_in: wp.array[wp.vec2i], + flex_in: wp.array[wp.vec2i], + vert_in: wp.array[wp.vec2i], + friction_in: wp.array[vec5], + solref_in: wp.array[wp.vec2], + solreffriction_in: wp.array[wp.vec2], + solimp_in: wp.array[vec5], + type_in: wp.array[int], + # Data out: + efc_type_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_pos_out: wp.array2d[float], + efc_margin_out: wp.array2d[float], + efc_D_out: wp.array2d[float], + efc_vel_out: wp.array2d[float], + efc_aref_out: wp.array2d[float], + efc_frictionloss_out: wp.array2d[float], + ): + conid, dimid = wp.tid() + + if conid >= nacon_in[0]: + return + + if not type_in[conid] & ContactType.CONSTRAINT: + return + + condim = condim_in[conid] + + if wp.static(IS_ELLIPTIC): + if dimid > condim - 1: + return + else: + if condim == 1 and dimid > 0: + return + elif condim > 1 and dimid >= 2 * (condim - 1): + return + + efcid = contact_efc_address_in[conid, dimid] + if efcid < 0: + return + + worldid = worldid_in[conid] timestep = opt_timestep[worldid % opt_timestep.shape[0]] impratio_invsqrt = opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] - contact_efc_address_out[conid, dimid] = efcid + + includemargin = includemargin_in[conid] + pos = dist_in[conid] - includemargin geom = geom_in[conid] + Jqvel = efc_Jqvel_in[worldid, efcid] if geom[0] >= 0: body1 = geom_bodyid[geom[0]] @@ -2040,145 +2204,43 @@ def _contact_elliptic( vert = vert_in[conid] body2 = flex_vertbodyid[flex_vertadr[flex[1]] + vert[1]] - con_pos = pos_in[conid] - frame = frame_in[conid] - - Jqvel = float(0.0) - - # skip fixed bodies - body1 = body_weldid[body1] - body2 = body_weldid[body2] - - da1 = int(body_dofadr[body1] + body_dofnum[body1] - 1) - da2 = int(body_dofadr[body2] + body_dofnum[body2] - 1) - - if is_sparse: - # count non-zeros - pda1 = da1 - pda2 = da2 - rownnz = int(0) - while pda1 >= 0 or pda2 >= 0: - da = wp.max(pda1, pda2) - # skip common dofs - if pda1 == da and pda2 == da: - break - if pda1 == da: - pda1 = dof_parentid[pda1] - if pda2 == da: - pda2 = dof_parentid[pda2] - rownnz += 1 - - # get rowadr - rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz) - if rowadr + rownnz > njmax_nnz_in: - return - efc_J_rowadr_out[worldid, efcid] = rowadr - efc_J_rownnz_out[worldid, efcid] = rownnz - - da = wp.max(da1, da2) - - if is_sparse: - nnz = int(0) - dofid = int(da) - else: - dofid = int(nv - 1) - - while True: - if is_sparse: - if nnz >= rownnz: - break - else: - if dofid < 0: - break - - if dofid == da: - # TODO(team): contact jacobian - jac1p, jac1r = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - con_pos, - body1, - dofid, - worldid, - ) - jac2p, jac2r = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - con_pos, - body2, - dofid, - worldid, - ) - - J = float(0.0) - 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 - - if is_sparse: - sparseid = rowadr + nnz - efc_J_colind_out[worldid, 0, sparseid] = dofid - efc_J_out[worldid, 0, sparseid] = J - nnz += 1 - else: - efc_J_out[worldid, efcid, dofid] = J - Jqvel += J * qvel_in[worldid, dofid] - if is_sparse and nnz >= rownnz: - break - - # Advance tree pointers and recompute da for next iteration - if da1 == da: - da1 = dof_parentid[da1] - if da2 == da: - da2 = dof_parentid[da2] - da = wp.max(da1, da2) - if is_sparse: - dofid = da - else: - dofid -= 1 - else: - if not is_sparse: - efc_J_out[worldid, efcid, dofid] = 0.0 - dofid -= 1 - body_invweight0_id = worldid % body_invweight0.shape[0] invweight = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, body2][0] ref = solref_in[conid] pos_aref = pos - if dimid > 0: - solreffriction = solreffriction_in[conid] + if wp.static(IS_ELLIPTIC): + if dimid > 0: + solreffriction = solreffriction_in[conid] - # non-normal directions use solreffriction (if non-zero) - if solreffriction[0] or solreffriction[1]: - ref = solreffriction + # non-normal directions use solreffriction (if non-zero) + if solreffriction[0] or solreffriction[1]: + ref = solreffriction - invweight = invweight * impratio_invsqrt * impratio_invsqrt - friction = friction_in[conid] + invweight = invweight * impratio_invsqrt * impratio_invsqrt + friction = friction_in[conid] - if dimid > 1: + if dimid > 1: + fri0 = friction[0] + frii = friction[dimid - 1] + fri = fri0 * fri0 / (frii * frii) + invweight *= fri + + pos_aref = 0.0 + else: + if condim > 1: + friction = friction_in[conid] fri0 = friction[0] - frii = friction[dimid - 1] - fri = fri0 * fri0 / (frii * frii) - invweight *= fri - - pos_aref = 0.0 + invweight = invweight + fri0 * fri0 * invweight + invweight = invweight * 2.0 * fri0 * fri0 * impratio_invsqrt * impratio_invsqrt if condim == 1: efc_type = ConstraintType.CONTACT_FRICTIONLESS - else: + elif wp.static(IS_ELLIPTIC): efc_type = ConstraintType.CONTACT_ELLIPTIC + else: + efc_type = ConstraintType.CONTACT_PYRAMIDAL _efc_row( opt_disableflags, @@ -2205,6 +2267,8 @@ def _contact_elliptic( efc_frictionloss_out, ) + return kernel + @event_scope def make_constraint(m: types.Model, d: types.Data): @@ -2243,6 +2307,7 @@ def make_constraint(m: types.Model, d: types.Data): m.eq_solimp, m.eq_data, m.is_sparse, + m.body_isdofancestor, m.eq_connect_adr, d.qvel, d.eq_active, @@ -2297,6 +2362,7 @@ def make_constraint(m: types.Model, d: types.Data): m.eq_solimp, m.eq_data, m.is_sparse, + m.body_isdofancestor, m.eq_wld_adr, d.qvel, d.eq_active, @@ -2433,6 +2499,7 @@ def make_constraint(m: types.Model, d: types.Data): m.eq_solimp, m.eq_flex_adr, d.qvel, + d.eq_active, d.flexedge_J, d.flexedge_length, d.njmax, @@ -2658,122 +2725,170 @@ def make_constraint(m: types.Model, d: types.Data): # contact if not (m.opt.disableflags & types.DisableBit.CONTACT): - if m.opt.cone == types.ConeType.PYRAMIDAL: + nmaxdim = int(m.nmaxpyramid) if m.opt.cone == types.ConeType.PYRAMIDAL else int(m.nmaxcondim) + + # Reinterpret to avoid unnecessary loads + contact_frame_2d = wp.array( + ptr=d.contact.frame.ptr, + dtype=wp.vec3, + shape=(d.naconmax, 3), + device=d.contact.frame.device, + copy=False, + ) + contact_friction_2d = wp.array( + ptr=d.contact.friction.ptr, + dtype=float, + shape=(d.naconmax, 5), + device=d.contact.friction.device, + copy=False, + ) + + wp.launch( + _efc_contact_init(m.opt.cone, m.is_sparse), + dim=d.naconmax, + inputs=[ + m.body_weldid, + m.body_dofnum, + m.body_dofadr, + m.dof_parentid, + m.geom_bodyid, + m.flex_vertadr, + m.flex_vertbodyid, + d.njmax, + d.njmax_nnz, + d.nacon, + d.contact.dist, + d.contact.dim, + d.contact.includemargin, + d.contact.worldid, + d.contact.geom, + d.contact.flex, + d.contact.vert, + d.contact.type, + ], + outputs=[ + d.nefc, + d.contact.efc_address, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + efc_nnz, + ], + ) + + if m.is_sparse: wp.launch( - _contact_pyramidal, - dim=(d.naconmax, m.nmaxpyramid), + _efc_contact_jac_sparse(m.opt.cone), + dim=(d.naconmax, nmaxdim), inputs=[ - m.nv, - m.opt.timestep, - m.opt.disableflags, - m.opt.impratio_invsqrt, m.body_parentid, m.body_rootid, m.body_weldid, m.body_dofnum, m.body_dofadr, - m.body_invweight0, m.dof_bodyid, m.dof_parentid, m.geom_bodyid, m.flex_vertadr, m.flex_vertbodyid, - m.is_sparse, + m.body_isdofancestor, d.qvel, d.subtree_com, d.cdof, - d.njmax, - d.njmax_nnz, + d.contact.efc_address, + d.efc.J_rownnz, + d.efc.J_rowadr, d.nacon, - d.contact.dist, d.contact.dim, - d.contact.includemargin, - d.contact.worldid, d.contact.geom, d.contact.flex, d.contact.vert, d.contact.pos, - d.contact.frame, - d.contact.friction, - d.contact.solref, - d.contact.solimp, - d.contact.type, + contact_frame_2d, + contact_friction_2d, + d.contact.worldid, ], outputs=[ - d.nefc, - d.contact.efc_address, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, d.efc.J_colind, d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - efc_nnz, + d.efc.Jqvel, ], ) - elif m.opt.cone == types.ConeType.ELLIPTIC: - wp.launch( - _contact_elliptic, - dim=(d.naconmax, m.nmaxcondim), + else: + d.efc.Jqvel.zero_() + tile_size = m.block_dim.contact_jac_tiled + n_dof_blocks = (m.nv_pad + tile_size - 1) // tile_size + + wp.launch_tiled( + _efc_contact_jac_dense(tile_size, m.opt.cone), + dim=(d.nworld, n_dof_blocks), inputs=[ - m.nv, - m.opt.timestep, - m.opt.disableflags, - m.opt.impratio_invsqrt, - m.body_parentid, m.body_rootid, - m.body_weldid, - m.body_dofnum, - m.body_dofadr, - m.body_invweight0, - m.dof_bodyid, - m.dof_parentid, m.geom_bodyid, m.flex_vertadr, m.flex_vertbodyid, - m.is_sparse, + m.body_isdofancestor, + d.ne, + d.nf, + d.nl, + d.nefc, d.qvel, d.subtree_com, d.cdof, + d.contact.efc_address, + d.efc.id, d.njmax, - d.njmax_nnz, - d.nacon, - d.contact.dist, + m.nv_pad, d.contact.dim, - d.contact.includemargin, - d.contact.worldid, d.contact.geom, d.contact.flex, d.contact.vert, d.contact.pos, - d.contact.frame, - d.contact.friction, - d.contact.solref, - d.contact.solreffriction, - d.contact.solimp, - d.contact.type, + contact_frame_2d, + contact_friction_2d, ], outputs=[ - d.nefc, - d.contact.efc_address, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - efc_nnz, + d.efc.Jqvel, ], + block_dim=tile_size, ) + + wp.launch( + _efc_contact_update(m.opt.cone), + dim=(d.naconmax, nmaxdim), + inputs=[ + m.opt.timestep, + m.opt.disableflags, + m.opt.impratio_invsqrt, + m.body_invweight0, + m.geom_bodyid, + m.flex_vertadr, + m.flex_vertbodyid, + d.contact.efc_address, + d.efc.Jqvel, + d.nacon, + d.contact.dist, + d.contact.dim, + d.contact.includemargin, + d.contact.worldid, + d.contact.geom, + d.contact.flex, + d.contact.vert, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + d.contact.type, + ], + outputs=[ + d.efc.type, + d.efc.id, + 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/derivative.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py index bdcd81cb..4a241f3b 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py @@ -15,7 +15,9 @@ import warp as wp +from mujoco.mjx.third_party.mujoco_warp._src import util_misc from mujoco.mjx.third_party.mujoco_warp._src.support import next_act +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 @@ -65,6 +67,28 @@ def _qderiv_actuator_passive_vel( if actuator_biastype[actid] == BiasType.AFFINE: bias = actuator_biasprm[actuator_biasprm_id, actid][2] + elif actuator_biastype[actid] == BiasType.DCMOTOR: + dynprm = actuator_dynprm[worldid % actuator_dynprm.shape[0], actid] + te = dynprm[0] + if te <= 0.0: + gainprm = actuator_gainprm[actuator_gainprm_id, actid] + R = gainprm[0] + K = gainprm[1] + + slots = util_misc.dcmotor_slots(dynprm, gainprm) + slot_Ta = slots[2] + + if slot_Ta >= 0: + adr = actuator_actadr[actid] + slot_Ta + T = act_in[worldid, adr] + alpha = gainprm[2] + T0 = gainprm[3] + Ta = dynprm[4] + R *= 1.0 + alpha * (T + Ta - T0) + + bias = -K * K / wp.max(MJ_MINVAL, R) + else: + bias = 0.0 else: bias = 0.0 @@ -228,8 +252,10 @@ def _qderiv_actuator_passive( opt_timestep: wp.array[float], opt_disableflags: int, dof_damping: wp.array2d[float], + dof_dampingpoly: wp.array2d[wp.vec2], is_sparse: bool, # Data in: + qvel_in: wp.array2d[float], qM_in: wp.array3d[float], # In: qMi: wp.array[int], @@ -249,7 +275,10 @@ def _qderiv_actuator_passive( qderiv = qDeriv_in[worldid, dofiid, dofjid] if not (opt_disableflags & DisableBit.DAMPER) and dofiid == dofjid: - qderiv -= dof_damping[worldid % dof_damping.shape[0], dofiid] + damping = dof_damping[worldid % dof_damping.shape[0], dofiid] + dpoly = dof_dampingpoly[worldid % dof_dampingpoly.shape[0], dofiid] + v = qvel_in[worldid, dofiid] + qderiv -= util_misc._poly_force_deriv(damping, dpoly, v, 1) qderiv *= opt_timestep[worldid % opt_timestep.shape[0]] @@ -272,9 +301,11 @@ def _qderiv_tendon_damping( ten_J_rowadr: wp.array[int], ten_J_colind: wp.array[int], tendon_damping: wp.array2d[float], + tendon_dampingpoly: wp.array2d[wp.vec2], is_sparse: bool, # Data in: ten_J_in: wp.array2d[float], + ten_velocity_in: wp.array2d[float], # In: qMi: wp.array[int], qMj: wp.array[int], @@ -289,7 +320,8 @@ def _qderiv_tendon_damping( tendon_damping_id = worldid % tendon_damping.shape[0] for tenid in range(ntendon): damping = tendon_damping[tendon_damping_id, tenid] - if damping == 0.0: + dpoly = tendon_dampingpoly[worldid % tendon_dampingpoly.shape[0], tenid] + if damping == 0.0 and dpoly[0] == 0.0 and dpoly[1] == 0.0: continue rownnz = ten_J_rownnz[tenid] @@ -305,7 +337,9 @@ def _qderiv_tendon_damping( Ji = ten_J_in[worldid, sparseid] if colind == dofjid: Jj = ten_J_in[worldid, sparseid] - qderiv -= Ji * Jj * damping + + v = ten_velocity_in[worldid, tenid] + qderiv -= Ji * Jj * util_misc._poly_force_deriv(damping, dpoly, v, 1) qderiv *= opt_timestep[worldid % opt_timestep.shape[0]] @@ -382,7 +416,9 @@ def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d[float]): m.opt.timestep, m.opt.disableflags, m.dof_damping, + m.dof_dampingpoly, m.is_sparse, + d.qvel, d.qM, qMi, qMj, @@ -405,8 +441,10 @@ def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d[float]): m.ten_J_rowadr, m.ten_J_colind, m.tendon_damping, + m.tendon_dampingpoly, m.is_sparse, d.ten_J, + d.ten_velocity, qMi, qMj, ], diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py index 55c7f357..c9b3c153 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py @@ -138,10 +138,13 @@ def _next_activation( actuator_actnum: wp.array[int], actuator_actlimited: wp.array[bool], actuator_dynprm: wp.array2d[vec10f], + actuator_gainprm: wp.array2d[vec10f], + actuator_biasprm: wp.array2d[vec10f], actuator_actrange: wp.array2d[wp.vec2], # Data in: act_in: wp.array2d[float], act_dot_in: wp.array2d[float], + actuator_velocity_in: wp.array2d[float], # In: act_dot_scale: float, limit: bool, @@ -152,20 +155,65 @@ def _next_activation( opt_timestep_id = worldid % opt_timestep.shape[0] actuator_dynprm_id = worldid % actuator_dynprm.shape[0] actuator_actrange_id = worldid % actuator_actrange.shape[0] + actuator_gainprm_id = worldid % actuator_gainprm.shape[0] + actuator_biasprm_id = worldid % actuator_biasprm.shape[0] + actadr = actuator_actadr[uid] actnum = actuator_actnum[uid] - for j in range(actadr, actadr + actnum): - act = next_act( - opt_timestep[opt_timestep_id], - actuator_dyntype[uid], - actuator_dynprm[actuator_dynprm_id, uid], - actuator_actrange[actuator_actrange_id, uid], - act_in[worldid, j], - act_dot_in[worldid, j], - act_dot_scale, - limit and actuator_actlimited[uid], - ) - act_out[worldid, j] = act + dyntype = actuator_dyntype[uid] + + if dyntype == DynType.DCMOTOR: + dynprm = actuator_dynprm[actuator_dynprm_id, uid] + gainprm = actuator_gainprm[actuator_gainprm_id, uid] + biasprm = actuator_biasprm[actuator_biasprm_id, uid] + slots = util_misc.dcmotor_slots(dynprm, gainprm) + + for j in range(actadr, actadr + actnum): + offset = j - actadr + act = act_in[worldid, j] + act_dot = act_dot_in[worldid, j] + + if offset == slots[4]: # current + R = gainprm[0] + te = wp.max(MJ_MINVAL, dynprm[0]) + act = act + act_dot * te * (1.0 - wp.exp(-opt_timestep[opt_timestep_id] / te)) + elif offset == slots[3]: # bristle + F_C = biasprm[3] + F_S = biasprm[4] + v_S = biasprm[5] + sigma0 = dynprm[5] + velocity = actuator_velocity_in[worldid, uid] + g = util_misc.lugre_stribeck(velocity, F_C, F_S, v_S) + + a = -sigma0 * wp.abs(velocity) / wp.max(MJ_MINVAL, g) + h = opt_timestep[opt_timestep_id] + exp_ah = wp.exp(a * h) + int_h = h + if wp.abs(a) > MJ_MINVAL: + int_h = (exp_ah - 1.0) / a + act = exp_ah * act + int_h * velocity + elif offset == slots[1]: # integral + act = act + act_dot * opt_timestep[opt_timestep_id] + Imax = dynprm[8] + if Imax > 0.0: + act = wp.clamp(act, -Imax, Imax) + else: # temperature and slew + act = act + act_dot * opt_timestep[opt_timestep_id] + + act_out[worldid, j] = act + else: + for j in range(actadr, actadr + actnum): + act = next_act( + opt_timestep[opt_timestep_id], + dyntype, + actuator_dynprm[actuator_dynprm_id, uid], + actuator_actrange[actuator_actrange_id, uid], + act_in[worldid, j], + act_dot_in[worldid, j], + act_dot_scale, + limit and actuator_actlimited[uid], + ) + act_out[worldid, j] = act @wp.kernel @@ -225,9 +273,12 @@ def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None) m.actuator_actnum, m.actuator_actlimited, m.actuator_dynprm, + m.actuator_gainprm, + m.actuator_biasprm, m.actuator_actrange, d.act, d.act_dot, + d.actuator_velocity, 1.0, True, ], @@ -274,12 +325,30 @@ def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None) wp.copy(d.qacc_warmstart, d.qacc) +@wp.kernel +def _compute_damping_deriv( + # Model: + dof_damping: wp.array2d[float], + dof_dampingpoly: wp.array2d[wp.vec2], + # Data in: + qvel_in: wp.array2d[float], + # Out: + deriv_out: wp.array2d[float], +): + worldid, tid = wp.tid() + damping = dof_damping[worldid % dof_damping.shape[0], tid] + dpoly = dof_dampingpoly[worldid % dof_dampingpoly.shape[0], tid] + v = qvel_in[worldid, tid] + deriv_out[worldid, tid] = util_misc._poly_force_deriv(damping, dpoly, v, 1) + + @wp.kernel def _euler_damp_qfrc_sparse( # Model: opt_timestep: wp.array[float], dof_Madr: wp.array[int], - dof_damping: wp.array2d[float], + # In: + damp_deriv: wp.array2d[float], # Out: qM_integration_out: wp.array3d[float], ): @@ -287,7 +356,7 @@ def _euler_damp_qfrc_sparse( timestep = opt_timestep[worldid % opt_timestep.shape[0]] adr = dof_Madr[tid] - qM_integration_out[worldid, 0, adr] += timestep * dof_damping[worldid % dof_damping.shape[0], tid] + qM_integration_out[worldid, 0, adr] += timestep * damp_deriv[worldid, tid] @cache_kernel @@ -296,11 +365,11 @@ def _tile_euler_dense(tile: TileSet): def euler_dense( # Model: opt_timestep: wp.array[float], - dof_damping: wp.array2d[float], # Data in: qM_in: wp.array3d[float], efc_Ma_in: wp.array2d[float], # In: + damp_deriv: wp.array2d[float], adr_in: wp.array[int], # Data out: qacc_out: wp.array2d[float], @@ -311,7 +380,7 @@ def _tile_euler_dense(tile: TileSet): 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 % dof_damping.shape[0]], shape=(TILE_SIZE,), offset=(dofid,)) + damping_tile = wp.tile_load(damp_deriv[worldid], shape=(TILE_SIZE,), offset=(dofid,)) damping_scaled = damping_tile * timestep qm_integration_tile = wp.tile_diag_add(M_tile, damping_scaled) @@ -329,6 +398,16 @@ def euler(m: Model, d: Data): # integrate damping implicitly if not (m.opt.disableflags & (DisableBit.EULERDAMP | DisableBit.DAMPER)): qacc = wp.empty((d.nworld, m.nv), dtype=float) + + # Compute damping derivative + damp_deriv = wp.empty((d.nworld, m.nv), dtype=float) + wp.launch( + _compute_damping_deriv, + dim=(d.nworld, m.nv), + inputs=[m.dof_damping, m.dof_dampingpoly, d.qvel], + outputs=[damp_deriv], + ) + if m.is_sparse: qM = wp.clone(d.qM) qLD = wp.empty((d.nworld, 1, m.nC), dtype=float) @@ -336,7 +415,7 @@ def euler(m: Model, d: Data): wp.launch( _euler_damp_qfrc_sparse, dim=(d.nworld, m.nv), - inputs=[m.opt.timestep, m.dof_Madr, m.dof_damping], + inputs=[m.opt.timestep, m.dof_Madr, damp_deriv], outputs=[qM], ) smooth.factor_solve_i(m, d, qM, qLD, qLDiagInv, qacc, d.efc.Ma) @@ -345,7 +424,7 @@ def euler(m: Model, d: Data): wp.launch_tiled( _tile_euler_dense(tile), dim=(d.nworld, tile.adr.size), - inputs=[m.opt.timestep, m.dof_damping, d.qM, d.efc.Ma, tile.adr], + inputs=[m.opt.timestep, d.qM, d.efc.Ma, damp_deriv, tile.adr], outputs=[qacc], block_dim=m.block_dim.euler_dense, ) @@ -390,9 +469,12 @@ def _rk_perturb_state( m.actuator_actnum, m.actuator_actlimited, m.actuator_dynprm, + m.actuator_gainprm, + m.actuator_biasprm, m.actuator_actrange, act_t0, d.act_dot, + d.actuator_velocity, scale, False, ], @@ -672,6 +754,96 @@ def _actuator_force( dynprm = actuator_dynprm[worldid % actuator_dynprm.shape[0], uid] act = act_in[worldid, act_last] act_dot = util_misc.muscle_dynamics(ctrl, act, dynprm) + elif dyntype == DynType.DCMOTOR: + gainprm = actuator_gainprm[worldid % actuator_gainprm.shape[0], uid] + slots = util_misc.dcmotor_slots(dynprm, gainprm) + adr = act_first + + act_dot = 0.0 + + # slew rate + if slots[0] >= 0: + u_prev = act_in[worldid, adr] + slew_s = dynprm[7] + slew = slew_s * opt_timestep[worldid % opt_timestep.shape[0]] + u_eff = wp.clamp(ctrl, u_prev - slew, u_prev + slew) + act_dot = (u_eff - u_prev) / opt_timestep[worldid % opt_timestep.shape[0]] + act_dot_out[worldid, adr] = act_dot + ctrl = u_eff + adr += 1 + + # integral + if slots[1] >= 0: + x_I = act_in[worldid, adr] + input_mode = int(gainprm[8]) + Imax = dynprm[8] + act_dot = ctrl + if input_mode == 1: + act_dot = ctrl - actuator_length_in[worldid, uid] + + if Imax > 0.0: + if x_I >= Imax: + act_dot = wp.min(act_dot, 0.0) + elif x_I <= -Imax: + act_dot = wp.max(act_dot, 0.0) + + act_dot_out[worldid, adr] = act_dot + adr += 1 + + # voltage + V = util_misc.dcmotor_voltage( + ctrl, + actuator_length_in[worldid, uid], + actuator_velocity_in[worldid, uid], + x_I, + gainprm, + ) + + # temperature + R = gainprm[0] + K = gainprm[1] + te = wp.max(MJ_MINVAL, dynprm[0]) + + if slots[2] >= 0: + RT = dynprm[2] + C = dynprm[3] + Ta = dynprm[4] + alpha = gainprm[2] + T0 = gainprm[3] + T = act_in[worldid, adr] + R_eff = R * (1.0 + alpha * (T + Ta - T0)) + + current = (V - K * actuator_velocity_in[worldid, uid]) / R_eff + if slots[4] >= 0: + current = act_in[worldid, act_last] + + act_dot = (R_eff * current * current - T / RT) / C + act_dot_out[worldid, adr] = act_dot + adr += 1 + R = R_eff + + # bristle + if slots[3] >= 0: + sigma0 = dynprm[5] + biasprm = actuator_biasprm[worldid % actuator_biasprm.shape[0], uid] + F_C = biasprm[3] + F_S = biasprm[4] + v_S = biasprm[5] + z = act_in[worldid, adr] + g = util_misc.lugre_stribeck(actuator_velocity_in[worldid, uid], F_C, F_S, v_S) + a = -sigma0 * wp.abs(actuator_velocity_in[worldid, uid]) / wp.max(MJ_MINVAL, g) + act_dot = a * z + actuator_velocity_in[worldid, uid] + act_dot_out[worldid, adr] = act_dot + adr += 1 + + # current + if slots[4] >= 0: + dimax = dynprm[1] + act_dot = (V / R - K / R * actuator_velocity_in[worldid, uid] - act_in[worldid, act_last]) / te + if dimax > 0.0: + act_dot = wp.clamp(act_dot, -dimax, dimax) + act_dot_out[worldid, act_last] = act_dot + elif dyntype == DynType.USER: act_dot = 0.0 # set by act_dyn_callback else: # DynType.NONE @@ -680,19 +852,54 @@ def _actuator_force( act_dot_out[worldid, act_last] = act_dot if actuator_actearly[uid]: - if dyntype == DynType.INTEGRATOR or dyntype == DynType.NONE: + if dyntype == DynType.INTEGRATOR or dyntype == DynType.NONE or dyntype == DynType.DCMOTOR: act = act_in[worldid, act_last] - ctrl_act = next_act( - opt_timestep[worldid % opt_timestep.shape[0]], - dyntype, - dynprm, - actuator_actrange[worldid % actuator_actrange.shape[0], uid], - act, - act_dot, - 1.0, - actuator_actlimited[uid], - ) + if dyntype == DynType.DCMOTOR: + gainprm = actuator_gainprm[worldid % actuator_gainprm.shape[0], uid] + slots = util_misc.dcmotor_slots(dynprm, gainprm) + offset = actuator_actnum[uid] - 1 + + if offset == slots[4]: # current + te = wp.max(MJ_MINVAL, dynprm[0]) + ctrl_act = act + act_dot * te * (1.0 - wp.exp(-opt_timestep[worldid % opt_timestep.shape[0]] / te)) + elif offset == slots[3]: # bristle + sigma0 = dynprm[5] + biasprm = actuator_biasprm[worldid % actuator_biasprm.shape[0], uid] + F_C = biasprm[3] + F_S = biasprm[4] + v_S = biasprm[5] + velocity = actuator_velocity_in[worldid, uid] + g = util_misc.lugre_stribeck(velocity, F_C, F_S, v_S) + a = -sigma0 * wp.abs(velocity) / wp.max(MJ_MINVAL, g) + h = opt_timestep[worldid % opt_timestep.shape[0]] + exp_ah = wp.exp(a * h) + int_h = h + if wp.abs(a) > MJ_MINVAL: + int_h = (exp_ah - 1.0) / a + ctrl_act = exp_ah * act + int_h * velocity + elif offset == slots[1]: # integral + ctrl_act = act + act_dot * opt_timestep[worldid % opt_timestep.shape[0]] + Imax = dynprm[8] + if Imax > 0.0: + ctrl_act = wp.clamp(ctrl_act, -Imax, Imax) + else: # temperature or slew or default + ctrl_act = act + act_dot * opt_timestep[worldid % opt_timestep.shape[0]] + + if actuator_actlimited[uid]: + actrange = actuator_actrange[worldid % actuator_actrange.shape[0], uid] + ctrl_act = wp.clamp(ctrl_act, actrange[0], actrange[1]) + else: + ctrl_act = next_act( + opt_timestep[worldid % opt_timestep.shape[0]], + dyntype, + dynprm, + actuator_actrange[worldid % actuator_actrange.shape[0], uid], + act, + act_dot, + 1.0, + actuator_actlimited[uid], + ) else: ctrl_act = act_in[worldid, act_last] @@ -712,6 +919,32 @@ def _actuator_force( acc0 = actuator_acc0[worldid % actuator_acc0.shape[0], uid] lengthrange = actuator_lengthrange[worldid % actuator_lengthrange.shape[0], uid] gain = util_misc.muscle_gain(length, velocity, lengthrange, acc0, gainprm) + elif gaintype == GainType.DCMOTOR: + R = gainprm[0] + K = gainprm[1] + te = dynprm[0] + + slots = util_misc.dcmotor_slots(dynprm, gainprm) + adr = act_first + + if slots[2] >= 0: + T = act_in[worldid, adr + slots[2]] + alpha = gainprm[2] + T0 = gainprm[3] + Ta = dynprm[4] + R *= 1.0 + alpha * (T + Ta - T0) + + gain = K if te > 0.0 else K / wp.max(MJ_MINVAL, R) + + if te <= 0.0: + input_mode = int(gainprm[8]) + if input_mode > 0: + x_I = 0.0 + if slots[1] >= 0: + x_I = act_in[worldid, adr + slots[1]] + ctrl_act = util_misc.dcmotor_voltage(ctrl, length, velocity, x_I, gainprm) + else: + ctrl_act = ctrl # GainType.USER: gain stays 0, modified by act_gain_callback # bias @@ -725,6 +958,10 @@ def _actuator_force( acc0 = actuator_acc0[worldid % actuator_acc0.shape[0], uid] lengthrange = actuator_lengthrange[worldid % actuator_lengthrange.shape[0], uid] bias = util_misc.muscle_bias(length, lengthrange, acc0, biasprm) + elif biastype == BiasType.DCMOTOR: + if dynprm[0] <= 0.0: + K = gainprm[1] + bias -= gain * K * velocity force = gain * ctrl_act + bias @@ -732,6 +969,25 @@ def _actuator_force( forcerange = actuator_forcerange[worldid % actuator_forcerange.shape[0], uid] force = wp.clamp(force, forcerange[0], forcerange[1]) + # add DC motor mechanical forces (not subject to current limits) + if biastype == BiasType.DCMOTOR: + # cogging torque + A = biasprm[0] + if A != 0.0: + Np = biasprm[1] + phi = biasprm[2] + force += A * wp.sin(Np * length + phi) + + # LuGre friction + sigma0 = dynprm[5] + if sigma0 > 0.0: + sigma1 = dynprm[6] + slots = util_misc.dcmotor_slots(dynprm, gainprm) + adr = act_first + slots[3] # slots[3] is bristle + z = act_in[worldid, adr] + z_dot = act_dot_out[worldid, adr] + force -= sigma0 * z + sigma1 * z_dot + actuator_force_out[worldid, uid] = force @@ -839,6 +1095,7 @@ def fwd_actuation(m: Model, d: Data): if not m.nu or (m.opt.disableflags & DisableBit.ACTUATION): d.act_dot.zero_() d.qfrc_actuator.zero_() + d.actuator_force.zero_() return wp.launch( @@ -1003,10 +1260,7 @@ def forward(m: Model, d: Data): @event_scope def step(m: Model, d: Data): """Advance simulation.""" - # TODO(team): mj_checkPos - # TODO(team): mj_checkVel forward(m, d) - # TODO(team): mj_checkAcc if m.opt.integrator == IntegratorType.EULER: euler(m, d) @@ -1022,8 +1276,6 @@ def step(m: Model, d: Data): def step1(m: Model, d: Data): """Advance simulation in two phases: before input is set by user.""" energy = m.opt.enableflags & EnableBit.ENERGY - # TODO(team): mj_checkPos - # TODO(team): mj_checkVel fwd_position(m, d) d.sensordata.zero_() sensor.sensor_pos(m, d) @@ -1053,7 +1305,6 @@ def step2(m: Model, d: Data): fwd_acceleration(m, d) solver.solve(m, d) sensor.sensor_acc(m, d) - # TODO(team): mj_checkAcc # integrate with Euler or implicitfast # TODO(team): implicit diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py index 065afde0..82b7cf4a 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py @@ -21,6 +21,7 @@ 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 import util_misc from mujoco.mjx.third_party.mujoco_warp._src.support import mul_m from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit @@ -36,14 +37,22 @@ def _qfrc_eulerdamp( # Model: opt_timestep: wp.array[float], dof_damping: wp.array2d[float], + dof_dampingpoly: wp.array2d[wp.vec2], # Data in: + qvel_in: wp.array2d[float], qacc_in: wp.array2d[float], # Out: qfrc_out: wp.array2d[float], ): worldid, dofid = wp.tid() timestep = opt_timestep[worldid % opt_timestep.shape[0]] - qfrc_out[worldid, dofid] += timestep * dof_damping[worldid % dof_damping.shape[0], dofid] * qacc_in[worldid, dofid] + + damping = dof_damping[worldid % dof_damping.shape[0], dofid] + dpoly = dof_dampingpoly[worldid % dof_dampingpoly.shape[0], dofid] + v = qvel_in[worldid, dofid] + + damp_deriv = util_misc._poly_force_deriv(damping, dpoly, v, 1) + qfrc_out[worldid, dofid] += timestep * damp_deriv * qacc_in[worldid, dofid] @wp.kernel @@ -91,11 +100,11 @@ def discrete_acc(m: Model, d: Data, qacc: wp.array2d[float]): # d.qM @ d.qacc support.mul_m(m, d, qfrc, d.qacc) - # qfrc += m.opt.timestep * m.dof_damping * d.qacc + # qfrc += m.opt.timestep * damp_deriv * d.qacc wp.launch( _qfrc_eulerdamp, dim=(d.nworld, m.nv), - inputs=[m.opt.timestep, m.dof_damping, d.qacc], + inputs=[m.opt.timestep, m.dof_damping, m.dof_dampingpoly, d.qvel, d.qacc], outputs=[qfrc], ) elif m.opt.integrator == IntegratorType.IMPLICITFAST: diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py index 56bbb18c..dc834e06 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py @@ -138,6 +138,10 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: if (mjm.sensor_plugin != -1).any(): raise NotImplementedError("Sensor plugins not supported.") + # array sizes may change in the future + if mujoco.mjNPOLY != 2: + warnings.warn(f"mujoco.mjNPOLY is {mujoco.mjNPOLY}, expected 2. Higher order polynomials may not be supported correctly.") + # 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: @@ -222,7 +226,10 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: m.nsensortaxel = mjm.mesh_vertnum[mjm.sensor_objid[mjm.sensor_type == mujoco.mjtSensor.mjSENS_TACTILE]].sum() m.nsensorcontact = (mjm.sensor_type == mujoco.mjtSensor.mjSENS_CONTACT).sum() m.nrangefinder = (mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER).sum() - m.nmaxcondim = np.concatenate(([0], mjm.geom_condim, mjm.pair_dim)).max() + condim_arrays = [np.array([0]), mjm.geom_condim, mjm.pair_dim] + if mjm.nflex > 0: + condim_arrays.append(mjm.flex_condim) + m.nmaxcondim = np.concatenate(condim_arrays).max() m.nmaxpyramid = np.maximum(1, 2 * (m.nmaxcondim - 1)) m.has_sdf_geom = (mjm.geom_type == mujoco.mjtGeom.mjGEOM_SDF).any() m.block_dim = types.BlockDim() @@ -266,6 +273,21 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: m.jnt_limited_ball_adr = np.nonzero(mjm.jnt_limited & (mjm.jnt_type == mujoco.mjtJoint.mjJNT_BALL))[0] m.dof_tri_row, m.dof_tri_col = np.tril_indices(mjm.nv) + # precompute body_isdofancestor: which DOFs affect each body + # TODO: Investigate alternative approach such as bitmap + body_isdofancestor = np.zeros((mjm.nbody, m.nv_pad), dtype=np.int32) + for bodyid in range(mjm.nbody): + b = bodyid + while b > 0 and mjm.body_dofnum[b] == 0: + b = mjm.body_parentid[b] + if mjm.body_dofnum[b] == 0: + continue + dofid = mjm.body_dofadr[b] + mjm.body_dofnum[b] - 1 + while dofid >= 0: + body_isdofancestor[bodyid, dofid] = 1 + dofid = mjm.dof_parentid[dofid] + m.body_isdofancestor = body_isdofancestor + # precalculated geom pairs filterparent = not (mjm.opt.disableflags & types.DisableBit.FILTERPARENT) @@ -918,7 +940,10 @@ def make_data( raise ValueError(f"nccdmax ({nccdmax}) must be <= nconmax ({nconmax})") sizes = dict({"*": 1}, **{f.name: getattr(mjm, f.name, None) for f in dataclasses.fields(types.Model) if f.type is int}) - sizes["nmaxcondim"] = np.concatenate(([0], mjm.geom_condim, mjm.pair_dim)).max() + condim_arrays = [np.array([0]), mjm.geom_condim, mjm.pair_dim] + if mjm.nflex > 0: + condim_arrays.append(mjm.flex_condim) + sizes["nmaxcondim"] = np.concatenate(condim_arrays).max() sizes["nmaxpyramid"] = np.maximum(1, 2 * (sizes["nmaxcondim"] - 1)) tile_size = types.TILE_SIZE_JTDAJ_SPARSE if is_sparse(mjm) else types.TILE_SIZE_JTDAJ_DENSE sizes["njmax_pad"], sizes["nv_pad"] = _get_padded_sizes(mjm.nv, njmax, is_sparse(mjm), tile_size) @@ -1090,7 +1115,10 @@ def put_data( raise ValueError(f"njmax overflow (njmax must be >= {mjd.nefc})") sizes = dict({"*": 1}, **{f.name: getattr(mjm, f.name, None) for f in dataclasses.fields(types.Model) if f.type is int}) - sizes["nmaxcondim"] = np.concatenate(([0], mjm.geom_condim, mjm.pair_dim)).max() + condim_arrays = [np.array([0]), mjm.geom_condim, mjm.pair_dim] + if mjm.nflex > 0: + condim_arrays.append(mjm.flex_condim) + sizes["nmaxcondim"] = np.concatenate(condim_arrays).max() sizes["nmaxpyramid"] = np.maximum(1, 2 * (sizes["nmaxcondim"] - 1)) tile_size = types.TILE_SIZE_JTDAJ_SPARSE if is_sparse(mjm) else types.TILE_SIZE_JTDAJ_DENSE sizes["njmax_pad"], sizes["nv_pad"] = _get_padded_sizes(mjm.nv, njmax, is_sparse(mjm), tile_size) @@ -1183,7 +1211,7 @@ def put_data( if mujoco.mj_isSparse(mjm): mujoco.mju_sparse2dense(mj_efc_J, mjd.efc_J, mjd.efc_J_rownnz, mjd.efc_J_rowadr, mjd.efc_J_colind) else: - mj_efc_J = mjd.efc_J.reshape((mjd.nefc, mjm.nv)) + mj_efc_J = mjd.efc_J.reshape((-1, mjm.nv))[: mjd.nefc] efc_J = np.zeros((nworld, sizes["njmax_pad"], sizes["nv_pad"]), dtype=float) efc_J[:, : mjd.nefc, : mjm.nv] = np.tile(mj_efc_J, (nworld, 1, 1)) efc.J = wp.array(efc_J, dtype=float) @@ -2659,6 +2687,7 @@ def create_render_context( cam_active: list[bool] | None = None, flex_render_smooth: bool = True, use_precomputed_rays: bool = True, + render_skybox: bool = False, ) -> types.RenderContext: """Creates a render context on device. @@ -2669,8 +2698,8 @@ def create_render_context( MuJoCo model values. render_rgb: Whether to render RGB images. If None, uses the MuJoCo model values. render_depth: Whether to render depth images. If None, uses the MuJoCo model values. - render_seg: Whether to render segmentation (per-pixel geom IDs). If None, - uses the MuJoCo model values. + render_seg: Whether to render segmentation (per-pixel object ID/type pairs). + If None, uses the MuJoCo model values. use_textures: Whether to use textures. use_shadows: Whether to use shadows. enabled_geom_groups: The geom groups to render. @@ -2679,6 +2708,8 @@ def create_render_context( flex_render_smooth: Whether to render flex meshes smoothly. use_precomputed_rays: Use precomputed rays instead of computing during rendering. When using domain randomization for camera intrinsics, set to False. + render_skybox: Whether to shade missed rays with the MuJoCo skybox texture. + Requires the model to contain a texture with type `mjTEXTURE_SKYBOX`. Returns: The render context containing rendering fields and output arrays on device. @@ -2737,27 +2768,37 @@ def create_render_context( flex_geom_flexid = [] flex_geom_edgeid = [] flex_bvh_id = np.full(nflex, 0, dtype=wp.uint64) - flex_group_root = np.zeros((nflex, nworld), dtype=int) + # Indexed later as [worldid, flexid]. + flex_group_root = np.full((nworld, nflex), -1, dtype=int) for f in range(nflex): if mjm.flex_dim[f] == 1: edge_adr = mjm.flex_edgeadr[f] flex_geom_flexid.extend([f] * mjm.flex_edgenum[f]) flex_geom_edgeid.extend([edge_adr + e for e in range(mjm.flex_edgenum[f])]) - flex_group_root[f] = np.zeros(nworld, dtype=int) else: flex_geom_flexid.append(f) flex_geom_edgeid.append(-1) fmesh, group_root = bvh.build_flex_bvh(mjm, mjd, nworld, f) flex_registry[f] = fmesh flex_bvh_id[f] = fmesh.id - flex_group_root[f] = group_root.numpy() + flex_group_root[:, f] = group_root.numpy() textures_registry = [] for i in range(mjm.ntex): textures_registry.append(render_util.create_warp_texture(mjm, i)) textures = wp.array(textures_registry, dtype=wp.Texture2D) + # Locate skybox texture + skybox_tex_ids = np.nonzero(mjm.tex_type == mujoco.mjtTexture.mjTEXTURE_SKYBOX)[0] if mjm.ntex else np.array([], dtype=int) + if render_skybox: + assert skybox_tex_ids.size > 0, "render_skybox=True but the model has no texture with type mjTEXTURE_SKYBOX" + skybox_tex_id = int(skybox_tex_ids[0]) + skybox_face_width = int(mjm.tex_width[skybox_tex_id]) + else: + skybox_tex_id = -1 + skybox_face_width = 1 + # Filter active cameras if cam_active is not None: assert len(cam_active) == mjm.ncam, f"cam_active must have length {mjm.ncam} (got {len(cam_active)})" @@ -2857,6 +2898,9 @@ def create_render_context( use_shadows=use_shadows, background_color=render_util.pack_rgba_to_uint32(0.1 * 255.0, 0.1 * 255.0, 0.2 * 255.0, 1.0 * 255.0), use_precomputed_rays=use_precomputed_rays, + render_skybox=render_skybox, + skybox_tex_id=skybox_tex_id, + skybox_face_width=skybox_face_width, bvh_ngeom=bvh_ngeom, enabled_geom_ids=wp.array(geom_enabled_idx, dtype=int), mesh_registry=mesh_registry, @@ -2892,7 +2936,7 @@ def create_render_context( depth_adr=wp.array(depth_adr, dtype=int), render_rgb=wp.array(render_rgb, dtype=bool), render_depth=wp.array(render_depth, dtype=bool), - seg_data=wp.zeros((nworld, max(si, 1)), dtype=int), + seg_data=wp.zeros((nworld, max(si, 1)), dtype=wp.vec2i), seg_adr=wp.array(seg_adr, dtype=int), render_seg=wp.array(render_seg, dtype=bool), znear=znear, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py index 4c76d1f8..31469590 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py @@ -17,6 +17,7 @@ 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 Data from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit @@ -76,7 +77,9 @@ def _spring_damper_dof_passive( jnt_qposadr: wp.array[int], jnt_dofadr: wp.array[int], jnt_stiffness: wp.array2d[float], + jnt_stiffnesspoly: wp.array2d[wp.vec2], dof_damping: wp.array2d[float], + dof_dampingpoly: wp.array2d[wp.vec2], # Data in: qpos_in: wp.array2d[float], qvel_in: wp.array2d[float], @@ -86,22 +89,37 @@ def _spring_damper_dof_passive( ): worldid, jntid = wp.tid() dofid = jnt_dofadr[jntid] + jnttype = jnt_type[jntid] stiffness = jnt_stiffness[worldid % jnt_stiffness.shape[0], jntid] + spoly = jnt_stiffnesspoly[worldid % jnt_stiffnesspoly.shape[0], jntid] damping = dof_damping[worldid % dof_damping.shape[0], dofid] + dpoly = dof_dampingpoly[worldid % dof_dampingpoly.shape[0], dofid] - has_stiffness = stiffness != 0.0 and not (opt_disableflags & DisableBit.SPRING) - has_damping = damping != 0.0 and not (opt_disableflags & DisableBit.DAMPER) + has_stiffness = (stiffness != 0.0 or spoly[0] != 0.0 or spoly[1] != 0.0) and not (opt_disableflags & DisableBit.SPRING) + has_damping = (damping != 0.0 or dpoly[0] != 0.0 or dpoly[1] != 0.0) and not (opt_disableflags & DisableBit.DAMPER) if not has_stiffness: - qfrc_spring_out[worldid, dofid] = 0.0 + if jnttype == JointType.FREE: + for i in range(6): + qfrc_spring_out[worldid, dofid + i] = 0.0 + elif jnttype == JointType.BALL: + for i in range(3): + qfrc_spring_out[worldid, dofid + i] = 0.0 + else: + qfrc_spring_out[worldid, dofid] = 0.0 if not has_damping: - qfrc_damper_out[worldid, dofid] = 0.0 + if jnttype == JointType.FREE: + for i in range(6): + qfrc_damper_out[worldid, dofid + i] = 0.0 + elif jnttype == JointType.BALL: + for i in range(3): + qfrc_damper_out[worldid, dofid + i] = 0.0 + else: + qfrc_damper_out[worldid, dofid] = 0.0 if not (has_stiffness or has_damping): return - - jnttype = jnt_type[jntid] qposid = jnt_qposadr[jntid] qpos_spring_id = worldid % qpos_spring.shape[0] @@ -113,9 +131,12 @@ def _spring_damper_dof_passive( qpos_in[worldid, qposid + 1] - qpos_spring[qpos_spring_id, qposid + 1], qpos_in[worldid, qposid + 2] - qpos_spring[qpos_spring_id, 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] + r = wp.length(dif) + k = util_misc._poly_force(stiffness, spoly, r, 0) + qfrc_spring_out[worldid, dofid + 0] = -k * dif[0] + qfrc_spring_out[worldid, dofid + 1] = -k * dif[1] + qfrc_spring_out[worldid, dofid + 2] = -k * dif[2] + rot = wp.quat( qpos_in[worldid, qposid + 3], qpos_in[worldid, qposid + 4], @@ -130,18 +151,18 @@ def _spring_damper_dof_passive( qpos_spring[qpos_spring_id, 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] + r_rot = wp.length(dif) + k_rot = util_misc._poly_force(stiffness, spoly, r_rot, 0) + qfrc_spring_out[worldid, dofid + 3] = -k_rot * dif[0] + qfrc_spring_out[worldid, dofid + 4] = -k_rot * dif[1] + qfrc_spring_out[worldid, dofid + 5] = -k_rot * 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] + for i in range(6): + v = qvel_in[worldid, dofid + i] + qfrc_damper_out[worldid, dofid + i] = -v * util_misc._poly_force(damping, dpoly, v, 1) + elif jnttype == JointType.BALL: # spring if has_stiffness: @@ -159,24 +180,28 @@ def _spring_damper_dof_passive( qpos_spring[qpos_spring_id, 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] + r = wp.length(dif) + k = util_misc._poly_force(stiffness, spoly, r, 0) + qfrc_spring_out[worldid, dofid + 0] = -k * dif[0] + qfrc_spring_out[worldid, dofid + 1] = -k * dif[1] + qfrc_spring_out[worldid, dofid + 2] = -k * 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] + for i in range(3): + v = qvel_in[worldid, dofid + i] + qfrc_damper_out[worldid, dofid + i] = -v * util_misc._poly_force(damping, dpoly, v, 1) + else: # mjJNT_SLIDE, mjJNT_HINGE # spring if has_stiffness: fdif = qpos_in[worldid, qposid] - qpos_spring[qpos_spring_id, qposid] - qfrc_spring_out[worldid, dofid] = -stiffness * fdif + qfrc_spring_out[worldid, dofid] = -fdif * util_misc._poly_force(stiffness, spoly, fdif, 0) # damper if has_damping: - qfrc_damper_out[worldid, dofid] = -damping * qvel_in[worldid, dofid] + v = qvel_in[worldid, dofid] + qfrc_damper_out[worldid, dofid] = -v * util_misc._poly_force(damping, dpoly, v, 1) @wp.kernel @@ -186,7 +211,9 @@ def _spring_damper_tendon_passive( ten_J_rowadr: wp.array[int], ten_J_colind: wp.array[int], tendon_stiffness: wp.array2d[float], + tendon_stiffnesspoly: wp.array2d[wp.vec2], tendon_damping: wp.array2d[float], + tendon_dampingpoly: wp.array2d[wp.vec2], tendon_lengthspring: wp.array2d[wp.vec2], # Data in: ten_J_in: wp.array2d[float], @@ -202,10 +229,12 @@ def _spring_damper_tendon_passive( worldid, tenid, dofid_sparse = wp.tid() stiffness = tendon_stiffness[worldid % tendon_stiffness.shape[0], tenid] + spoly = tendon_stiffnesspoly[worldid % tendon_stiffnesspoly.shape[0], tenid] damping = tendon_damping[worldid % tendon_damping.shape[0], tenid] + dpoly = tendon_dampingpoly[worldid % tendon_dampingpoly.shape[0], tenid] - has_stiffness = stiffness != 0.0 and not dsbl_spring - has_damping = damping != 0.0 and not dsbl_damper + has_stiffness = (stiffness != 0.0 or spoly[0] != 0.0 or spoly[1] != 0.0) and not dsbl_spring + has_damping = (damping != 0.0 or dpoly[0] != 0.0 or dpoly[1] != 0.0) and not dsbl_damper if not has_stiffness and not has_damping: return @@ -225,19 +254,16 @@ def _spring_damper_tendon_passive( 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 + x = wp.where(length > upper, length - upper, wp.where(length < lower, length - lower, 0.0)) + frc_spring = -x * util_misc._poly_force(stiffness, spoly, x, 0) # transform to joint torque wp.atomic_add(qfrc_spring_out[worldid], dofid, J * frc_spring) if has_damping: - # compute damper linear force along tendon - frc_damper = -damping * ten_velocity_in[worldid, tenid] + # compute damper force along tendon + v = ten_velocity_in[worldid, tenid] + frc_damper = -v * util_misc._poly_force(damping, dpoly, v, 1) # transform to joint torque wp.atomic_add(qfrc_damper_out[worldid], dofid, J * frc_damper) @@ -252,6 +278,7 @@ def _gravity_force( body_mass: wp.array2d[float], body_gravcomp: wp.array2d[float], dof_bodyid: wp.array[int], + body_isdofancestor: wp.array2d[int], # Data in: xipos_in: wp.array2d[wp.vec3], subtree_com_in: wp.array2d[wp.vec3], @@ -267,7 +294,9 @@ def _gravity_force( if gravcomp: force = -gravity * body_mass[worldid % body_mass.shape[0], bodyid] * gravcomp pos = xipos_in[worldid, bodyid] - jac, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos, bodyid, dofid, worldid) + jac, _ = support.jac_dof( + body_parentid, body_rootid, dof_bodyid, body_isdofancestor, subtree_com_in, cdof_in, pos, bodyid, dofid, worldid + ) wp.atomic_add(qfrc_gravcomp_out[worldid], dofid, wp.dot(jac, force)) @@ -715,9 +744,10 @@ def _flex_bending( force = wp.matrix(0.0, shape=(nvert, 3)) for i in range(nvert): for x in range(3): + acc = float(0.0) for j in range(nvert): - force[i, x] -= flex_bending[edgeid, 4 * i + j] * flexvert_xpos_in[worldid, v[j]][x] - force[i, x] -= flex_bending[edgeid, 16] * frc[i, x] + acc += flex_bending[edgeid, 4 * i + j] * flexvert_xpos_in[worldid, v[j]][x] + force[i, x] = -(acc + flex_bending[edgeid, 16] * frc[i, x]) for i in range(nvert): bodyid = flex_vertbodyid[v[i]] @@ -749,7 +779,9 @@ def passive(m: Model, d: Data): m.jnt_qposadr, m.jnt_dofadr, m.jnt_stiffness, + m.jnt_stiffnesspoly, m.dof_damping, + m.dof_dampingpoly, d.qpos, d.qvel, ], @@ -765,7 +797,9 @@ def passive(m: Model, d: Data): m.ten_J_rowadr, m.ten_J_colind, m.tendon_stiffness, + m.tendon_stiffnesspoly, m.tendon_damping, + m.tendon_dampingpoly, m.tendon_lengthspring, d.ten_J, d.ten_length, @@ -840,6 +874,7 @@ def passive(m: Model, d: Data): m.body_mass, m.body_gravcomp, m.dof_bodyid, + m.body_isdofancestor, d.xipos, d.subtree_com, d.cdof, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py index 371032e1..ba3b4cfe 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py @@ -34,6 +34,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL 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 ObjType from mujoco.mjx.third_party.mujoco_warp._src.types import RenderContext from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope @@ -84,6 +85,72 @@ def sample_texture( return wp.vec3(tex_color[0], tex_color[1], tex_color[2]) +@wp.func +def sample_skybox( + # In: + skybox_tex: wp.Texture2D, + face_width_inv: float, + ray_dir_world: wp.vec3, +) -> wp.vec3: + # MuJoCo maps a world-space direction to cube-map space by rotating 90° about X + # (see render_gl3.c: S=x, T=z, R=-y). Faces in tex_data are stacked vertically + # in OpenGL cube-face order: +X, -X, +Y, -Y, +Z, -Z. + rx = ray_dir_world[0] + ry = ray_dir_world[2] + rz = -ray_dir_world[1] + + arx = wp.abs(rx) + ary = wp.abs(ry) + arz = wp.abs(rz) + + face = int(0) + sc = float(0.0) + tc = float(0.0) + ma = float(1.0) + + if arx >= ary and arx >= arz: + ma = arx + if rx > 0.0: + face = 0 + sc = -rz + tc = -ry + else: + face = 1 + sc = rz + tc = -ry + elif ary >= arz: + ma = ary + if ry > 0.0: + face = 2 + sc = rx + tc = rz + else: + face = 3 + sc = rx + tc = -rz + else: + ma = arz + if rz > 0.0: + face = 4 + sc = rx + tc = -ry + else: + face = 5 + sc = -rx + tc = -ry + + s = (math.safe_div(sc, ma) + 1.0) * 0.5 + t = (math.safe_div(tc, ma) + 1.0) * 0.5 + + # Keep the linear filter from bleeding between adjacent faces in the vertical strip. + t_min = 0.5 * face_width_inv + t = wp.clamp(t, t_min, 1.0 - t_min) + + v = (float(face) + t) * wp.static(1.0 / 6.0) + color = wp.texture_sample(skybox_tex, wp.vec2(s, v), dtype=wp.vec4) + return wp.vec3(color[0], color[1], color[2]) + + # TODO: Investigate combining cast_ray and cast_ray_first_hit @wp.func def cast_ray( @@ -525,7 +592,7 @@ def render(m: Model, d: Data, rc: RenderContext): """ rc.rgb_data.fill_(rc.background_color) rc.depth_data.fill_(0.0) - rc.seg_data.fill_(-1) + rc.seg_data.fill_(wp.vec2i(-1, -1)) @wp.kernel(module="unique", enable_backward=False) def _render_megakernel( @@ -588,7 +655,7 @@ def render(m: Model, d: Data, rc: RenderContext): # Out: rgb_out: wp.array2d[wp.uint32], depth_out: wp.array2d[float], - seg_out: wp.array2d[int], + seg_out: wp.array2d[wp.vec2i], ): worldid, rayid = wp.tid() @@ -661,10 +728,25 @@ def render(m: Model, d: Data, rc: RenderContext): ) if render_seg[cam_idx] and geom_id != -1: - seg_out[worldid, seg_adr[cam_idx] + rayid_local] = geom_id + if geom_id == -2: + seg_out[worldid, seg_adr[cam_idx] + rayid_local] = wp.vec2i(mesh_id, int(ObjType.FLEX)) + else: + seg_out[worldid, seg_adr[cam_idx] + rayid_local] = wp.vec2i(geom_id, int(ObjType.GEOM)) # Early Out if geom_id == -1: + if wp.static(rc.render_skybox) and render_rgb[cam_idx]: + skybox_color = sample_skybox( + textures[wp.static(rc.skybox_tex_id)], + wp.static(1.0 / float(rc.skybox_face_width)), + ray_dir_world, + ) + rgb_out[worldid, rgb_adr[cam_idx] + rayid_local] = pack_rgba_to_uint32( + skybox_color[0] * 255.0, + skybox_color[1] * 255.0, + skybox_color[2] * 255.0, + 255.0, + ) return if render_depth[cam_idx]: diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render_util.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render_util.py index c59da005..b87c8b2e 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render_util.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render_util.py @@ -211,13 +211,13 @@ def get_depth(rc: RenderContext, camera_index: int, depth_scale: float, depth_ou @wp.kernel def _extract_seg_kernel( # In: - seg_data: wp.array2d[int], + seg_data: wp.array2d[wp.vec2i], seg_adr: wp.array[int], camera_index: int, # Out: - seg_out: wp.array3d[int], + seg_out: wp.array3d[wp.vec2i], ): - """Extract per-pixel geom IDs from the render context buffers for a given camera index.""" + """Extract per-pixel `(object_id, object_type)` pairs for a camera.""" worldid, pixelid = wp.tid() xid = pixelid % seg_out.shape[2] yid = pixelid // seg_out.shape[2] @@ -226,17 +226,18 @@ def _extract_seg_kernel( seg_out[worldid, yid, xid] = seg_data[worldid, seg_adr_offset + pixelid] -def get_segmentation(rc: RenderContext, camera_index: int, seg_out: wp.array3d[int]): +def get_segmentation(rc: RenderContext, camera_index: int, seg_out: wp.array3d[wp.vec2i]): """Get the segmentation data from the render context buffers for a given camera index. - Each pixel contains the MuJoCo geom ID of the geometry hit by the ray, -1 for - background, or -2 for flex bodies. + Each pixel stores MuJoCo-style `(object_id, object_type)` data. Background + pixels are `(-1, -1)`. Regular geometry hits are `(geom_id, mjOBJ_GEOM)`. + Flex hits are `(flex_id, mjOBJ_FLEX)`. Args: rc: The render context on device. camera_index: The index of the camera to get the segmentation data for. - seg_out: The output array to store the geom IDs in, with shape - (nworld, height, width). + seg_out: The output array to store segmentation data in, with shape + `(nworld, height, width)` and dtype `wp.vec2i`. """ wp.launch( _extract_seg_kernel, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py index b615d20c..381dfd66 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py @@ -44,6 +44,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import vec8 from mujoco.mjx.third_party.mujoco_warp._src.types import vec8i from mujoco.mjx.third_party.mujoco_warp._src.types import vec_pluginattr from mujoco.mjx.third_party.mujoco_warp._src.util_misc import inside_geom +from mujoco.mjx.third_party.mujoco_warp._src.util_misc import poly_potential 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 @@ -2735,6 +2736,7 @@ def _energy_pos_passive_joint( jnt_type: wp.array[int], jnt_qposadr: wp.array[int], jnt_stiffness: wp.array2d[float], + jnt_stiffnesspoly: wp.array2d[wp.vec2], # Data in: qpos_in: wp.array2d[float], # Data out: @@ -2743,8 +2745,9 @@ def _energy_pos_passive_joint( worldid, jntid = wp.tid() jnt_stiffness_id = worldid % jnt_stiffness.shape[0] stiffness = jnt_stiffness[jnt_stiffness_id, jntid] + spoly = jnt_stiffnesspoly[worldid % jnt_stiffnesspoly.shape[0], jntid] - if stiffness == 0.0: + if stiffness == 0.0 and spoly[0] == 0.0 and spoly[1] == 0.0: return padr = jnt_qposadr[jntid] @@ -2776,8 +2779,11 @@ def _energy_pos_passive_joint( dif1 = math.quat_sub(quat1, quat_spring) + r0 = wp.length(dif0) + r1 = wp.length(dif1) + energy = wp.vec2( - 0.5 * stiffness * (wp.dot(dif0, dif0) + wp.dot(dif1, dif1)), + poly_potential(stiffness, spoly, r0, 0) + poly_potential(stiffness, spoly, r1, 0), 0.0, ) @@ -2800,15 +2806,16 @@ def _energy_pos_passive_joint( ) dif = math.quat_sub(quat, quat_spring) + r = wp.length(dif) energy = wp.vec2( - 0.5 * stiffness * wp.dot(dif, dif), + poly_potential(stiffness, spoly, r, 0), 0.0, ) wp.atomic_add(energy_out, worldid, energy) elif jnttype == JointType.SLIDE or jnttype == JointType.HINGE: dif_ = qpos_in[worldid, padr] - qpos_spring[qpos_spring_id, padr] energy = wp.vec2( - 0.5 * stiffness * dif_ * dif_, + poly_potential(stiffness, spoly, dif_, 0), 0.0, ) wp.atomic_add(energy_out, worldid, energy) @@ -2818,6 +2825,7 @@ def _energy_pos_passive_joint( def _energy_pos_passive_tendon( # Model: tendon_stiffness: wp.array2d[float], + tendon_stiffnesspoly: wp.array2d[wp.vec2], tendon_lengthspring: wp.array2d[wp.vec2], # Data in: ten_length_in: wp.array2d[float], @@ -2828,8 +2836,9 @@ def _energy_pos_passive_tendon( tendon_stiffness_id = worldid % tendon_stiffness.shape[0] stiffness = tendon_stiffness[tendon_stiffness_id, tenid] + spoly = tendon_stiffnesspoly[worldid % tendon_stiffnesspoly.shape[0], tenid] - if stiffness == 0.0: + if stiffness == 0.0 and spoly[0] == 0.0 and spoly[1] == 0.0: return length = ten_length_in[worldid, tenid] @@ -2841,13 +2850,13 @@ def _energy_pos_passive_tendon( upper = lengthspring[1] if length > upper: - displacement = upper - length + x = length - upper elif length < lower: - displacement = lower - length + x = length - lower else: - displacement = 0.0 + x = 0.0 - energy = wp.vec2(0.5 * stiffness * displacement * displacement, 0.0) + energy = wp.vec2(poly_potential(stiffness, spoly, x, 0), 0.0) wp.atomic_add(energy_out, worldid, energy) @@ -2871,6 +2880,7 @@ def energy_pos(m: Model, d: Data): m.jnt_type, m.jnt_qposadr, m.jnt_stiffness, + m.jnt_stiffnesspoly, d.qpos, ], outputs=[d.energy], @@ -2883,6 +2893,7 @@ def energy_pos(m: Model, d: Data): dim=(d.nworld, m.ntendon), inputs=[ m.tendon_stiffness, + m.tendon_stiffnesspoly, m.tendon_lengthspring, d.ten_length, ], diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py index bdb90ecd..bda36c7c 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py @@ -1197,7 +1197,9 @@ def _cfrc( cfrc_int_out: wp.array2d[wp.spatial_vector], ): worldid, bodyid = wp.tid() - bodyid += 1 # skip world body + if bodyid == 0: + cfrc_int_out[worldid, 0] = wp.spatial_vector(0.0, 0.0, 0.0, 0.0, 0.0, 0.0) + return cacc = cacc_in[worldid, bodyid] cinert = cinert_in[worldid, bodyid] cvel = cvel_in[worldid, bodyid] @@ -1210,9 +1212,7 @@ def _cfrc( 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.launch(_cfrc, dim=[d.nworld, m.nbody], inputs=[d.cinert, d.cvel, d.cacc, d.cfrc_ext, flg_cfrc_ext], outputs=[d.cfrc_int]) @wp.kernel @@ -1983,6 +1983,9 @@ def _comvel_branch( cvel += cdof[dofid + 1] * qvel[dofid + 1] cvel += cdof[dofid + 2] * qvel[dofid + 2] + cdof_dot_out[worldid, dofid + 0] = wp.spatial_vector(0.0, 0.0, 0.0, 0.0, 0.0, 0.0) + cdof_dot_out[worldid, dofid + 1] = wp.spatial_vector(0.0, 0.0, 0.0, 0.0, 0.0, 0.0) + cdof_dot_out[worldid, dofid + 2] = wp.spatial_vector(0.0, 0.0, 0.0, 0.0, 0.0, 0.0) 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]) @@ -2061,6 +2064,7 @@ def _transmission( actuator_trnid: wp.array[wp.vec2i], actuator_gear: wp.array2d[wp.spatial_vector], actuator_cranklength: wp.array2d[float], + body_isdofancestor: wp.array2d[int], # Data in: qpos_in: wp.array2d[float], xquat_in: wp.array2d[wp.quat], @@ -2219,12 +2223,30 @@ def _transmission( # get Jacobians of axis(jacA) and vec(jac) jacp, jacr = support.jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos_idslider, site_bodyid[idslider], da, worldid + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + site_xpos_idslider, + site_bodyid[idslider], + da, + worldid, ) jacS = jacp jacA = wp.cross(jacr, axis) jac, _ = support.jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos_id, site_bodyid[id], da, worldid + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + site_xpos_id, + site_bodyid[id], + da, + worldid, ) jac -= jacS @@ -2313,6 +2335,7 @@ def _transmission( body_parentid, body_rootid, dof_bodyid, + body_isdofancestor, subtree_com_in, cdof_in, site_xpos_in[worldid, siteid], @@ -2419,10 +2442,28 @@ def _transmission( break jacp, jacr = support.jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos, site_bodyid[siteid], da, worldid + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + site_xpos, + site_bodyid[siteid], + da, + worldid, ) jacpref, jacrref = support.jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, ref_xpos, site_bodyid[refid], da, worldid + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + ref_xpos, + site_bodyid[refid], + da, + worldid, ) moment = float(0.0) @@ -2453,6 +2494,7 @@ def _transmission_body_moment( dof_bodyid: wp.array[int], geom_bodyid: wp.array[int], actuator_trnid: wp.array[wp.vec2i], + body_isdofancestor: wp.array2d[int], actuator_trntype_body_adr: wp.array[int], # Data in: subtree_com_in: wp.array2d[wp.vec3], @@ -2568,10 +2610,10 @@ def _transmission_body_moment( colind = dofid jacp1, _ = support.jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, contact_pos, b1, colind, worldid + body_parentid, body_rootid, dof_bodyid, body_isdofancestor, subtree_com_in, cdof_in, contact_pos, b1, colind, worldid ) jacp2, _ = support.jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, contact_pos, b2, colind, worldid + body_parentid, body_rootid, dof_bodyid, body_isdofancestor, subtree_com_in, cdof_in, contact_pos, b2, colind, worldid ) jacdif = jacp2 - jacp1 @@ -2635,6 +2677,7 @@ def transmission(m: Model, d: Data): m.actuator_trnid, m.actuator_gear, m.actuator_cranklength, + m.body_isdofancestor, d.qpos, d.xquat, d.site_xpos, @@ -2662,6 +2705,7 @@ def transmission(m: Model, d: Data): m.dof_bodyid, m.geom_bodyid, m.actuator_trnid, + m.body_isdofancestor, m.actuator_trntype_body_adr, d.subtree_com, d.cdof, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py index 3964107d..0c804716 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py @@ -2785,6 +2785,36 @@ def update_gradient_cholesky_blocked(tile_size: int, matrix_size: int): return kernel +def update_gradient_cholesky_blocked_skip_unchanged(tile_size: int, matrix_size: int): + """Blocked Cholesky that skips factorization when no constraints changed.""" + + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # In: + ctx_done_in: wp.array[bool], + ctx_grad_in: wp.array3d[float], + ctx_h_in: wp.array3d[float], + changed_count_in: wp.array[int], + ctx_hfactor: wp.array3d[float], + # Out: + ctx_Mgrad_out: wp.array3d[float], + ): + worldid = wp.tid() + TILE_SIZE = wp.static(tile_size) + + if ctx_done_in[worldid]: + return + + if changed_count_in[worldid] > 0: + wp.static(create_blocked_cholesky_func(TILE_SIZE))(ctx_h_in[worldid], matrix_size, ctx_hfactor[worldid]) + + wp.static(create_blocked_cholesky_solve_func(TILE_SIZE, matrix_size))( + ctx_hfactor[worldid], ctx_grad_in[worldid], matrix_size, ctx_Mgrad_out[worldid] + ) + + return kernel + + @wp.kernel def padding_h(nv: int, ctx_done_in: wp.array[bool], ctx_h_out: wp.array3d[float]): worldid, elementid = wp.tid() @@ -2796,8 +2826,12 @@ def padding_h(nv: int, ctx_done_in: wp.array[bool], ctx_h_out: wp.array3d[float] ctx_h_out[worldid, dofid, dofid] = 1.0 -def _cholesky_factorize_solve(m: types.Model, d: types.Data, ctx: SolverContext): - """Cholesky factorize ctx.h and solve for Mgrad.""" +def _cholesky_factorize_solve(m: types.Model, d: types.Data, ctx: SolverContext, skip_unchanged: bool = False): + """Cholesky factorize ctx.h and solve for Mgrad. + + If skip_unchanged is True (blocked path only), worlds where no constraints + changed reuse the cached factorization in hfactor instead of refactorizing. + """ if m.nv <= _BLOCK_CHOLESKY_DIM: wp.launch_tiled( update_gradient_cholesky(m.nv), @@ -2814,13 +2848,22 @@ def _cholesky_factorize_solve(m: types.Model, d: types.Data, ctx: SolverContext) outputs=[ctx.h], ) - wp.launch_tiled( - update_gradient_cholesky_blocked(types.TILE_SIZE_JTDAJ_DENSE, m.nv_pad), - dim=d.nworld, - inputs=[ctx.done, ctx.grad.reshape(shape=(d.nworld, ctx.grad.shape[1], 1)), ctx.h, ctx.hfactor], - outputs=[ctx.Mgrad.reshape(shape=(d.nworld, ctx.Mgrad.shape[1], 1))], - block_dim=m.block_dim.update_gradient_cholesky_blocked, - ) + if skip_unchanged: + wp.launch_tiled( + update_gradient_cholesky_blocked_skip_unchanged(types.TILE_SIZE_JTDAJ_DENSE, m.nv_pad), + dim=d.nworld, + inputs=[ctx.done, ctx.grad.reshape(shape=(d.nworld, ctx.grad.shape[1], 1)), ctx.h, ctx.changed_efc_count, ctx.hfactor], + outputs=[ctx.Mgrad.reshape(shape=(d.nworld, ctx.Mgrad.shape[1], 1))], + block_dim=m.block_dim.update_gradient_cholesky_blocked, + ) + else: + wp.launch_tiled( + update_gradient_cholesky_blocked(types.TILE_SIZE_JTDAJ_DENSE, m.nv_pad), + dim=d.nworld, + inputs=[ctx.done, ctx.grad.reshape(shape=(d.nworld, ctx.grad.shape[1], 1)), ctx.h, ctx.hfactor], + outputs=[ctx.Mgrad.reshape(shape=(d.nworld, ctx.Mgrad.shape[1], 1))], + block_dim=m.block_dim.update_gradient_cholesky_blocked, + ) @wp.kernel @@ -3055,7 +3098,7 @@ def _update_gradient_incremental(m: types.Model, d: types.Data, ctx: SolverConte outputs=[ctx.h], ) - _cholesky_factorize_solve(m, d, ctx) + _cholesky_factorize_solve(m, d, ctx, skip_unchanged=True) @wp.kernel diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py index f6b3fe9d..b42e9808 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py @@ -393,12 +393,29 @@ def transform_force(frc: wp.spatial_vector, offset: wp.vec3) -> wp.spatial_vecto return transform_force(force, torque, offset) +@wp.func +def _compute_jacp(cdof_clip: wp.spatial_vector, offset: wp.vec3, affect: int) -> wp.vec3: + if affect == 0: + return wp.vec3(0.0) + cdof_lin = wp.spatial_bottom(cdof_clip) + cdof_ang = wp.spatial_top(cdof_clip) + return cdof_lin + wp.cross(cdof_ang, offset) + + +@wp.func +def _compute_jacr(cdof_clip: wp.spatial_vector, affect: int) -> wp.vec3: + if affect == 0: + return wp.vec3(0.0) + return wp.spatial_top(cdof_clip) + + @wp.func def jac_dof( # Model: body_parentid: wp.array[int], body_rootid: wp.array[int], dof_bodyid: wp.array[int], + body_isdofancestor: wp.array2d[int], # Data in: subtree_com_in: wp.array2d[wp.vec3], cdof_in: wp.array2d[wp.spatial_vector], @@ -408,16 +425,7 @@ def jac_dof( 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: + if body_isdofancestor[bodyid, dofid] == 0: return wp.vec3(0.0), wp.vec3(0.0) offset = point - wp.vec3(subtree_com_in[worldid, body_rootid[bodyid]]) @@ -440,6 +448,7 @@ def _make_jac_kernel(has_jacp: bool, has_jacr: bool): body_parentid: wp.array[int], body_rootid: wp.array[int], dof_bodyid: wp.array[int], + body_isdofancestor: wp.array2d[int], # Data in: subtree_com_in: wp.array2d[wp.vec3], cdof_in: wp.array2d[wp.spatial_vector], @@ -453,7 +462,16 @@ def _make_jac_kernel(has_jacp: bool, has_jacr: bool): worldid, dofid = wp.tid() jacp_val, jacr_val = jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, point_in[worldid], bodyid_in[worldid], dofid, worldid + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + point_in[worldid], + bodyid_in[worldid], + dofid, + worldid, ) if wp.static(has_jacp): @@ -496,7 +514,7 @@ def jac( wp.launch( kernel, dim=(d.nworld, m.nv), - inputs=[m.body_parentid, m.body_rootid, m.dof_bodyid, d.subtree_com, d.cdof, point, body], + inputs=[m.body_parentid, m.body_rootid, m.dof_bodyid, m.body_isdofancestor, d.subtree_com, d.cdof, point, body], outputs=[jacp_arr, jacr_arr], ) @@ -510,6 +528,7 @@ def jac_dot_dof( jnt_dofadr: wp.array[int], dof_bodyid: wp.array[int], dof_jntid: wp.array[int], + body_isdofancestor: wp.array2d[int], # Data in: subtree_com_in: wp.array2d[wp.vec3], cdof_in: wp.array2d[wp.spatial_vector], @@ -521,16 +540,7 @@ def jac_dot_dof( 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: + if body_isdofancestor[bodyid, dofid] == 0: return wp.vec3(0.0), wp.vec3(0.0) com = subtree_com_in[worldid, body_rootid[bodyid]] diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py index bf4f1b6b..2395a711 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py @@ -67,6 +67,7 @@ class BlockDim: update_gradient_JTDAJ_sparse: int = 64 update_gradient_JTDAJ_dense: int = 96 linesearch_iterative: int = 32 + contact_jac_tiled: int = 32 # derivative qderiv_actuator_dense: int = 32 @@ -184,7 +185,6 @@ class DisableBit(enum.IntFlag): EULERDAMP: implicit damping for Euler integration NATIVECCD: native convex collision detection (ignored in MJWarp) ISLAND: constraint islands - MULTICCD: multiple CCD contact points """ CONSTRAINT = mujoco.mjtDisableBit.mjDSBL_CONSTRAINT @@ -214,6 +214,7 @@ class EnableBit(enum.IntFlag): Attributes: ENERGY: energy computation INVDISCRETE: discrete-time inverse dynamics + MULTICCD: multiple contacts with CCD """ ENERGY = mujoco.mjtEnableBit.mjENBL_ENERGY @@ -251,6 +252,7 @@ class DynType(enum.IntEnum): FILTEREXACT: linear filter: da/dt = (u-a) / tau, with exact integration MUSCLE: piece-wise linear filter with two time constants USER: user-defined dynamics via act_dyn_callback + DCMOTOR: DC motor dynamics """ NONE = mujoco.mjtDyn.mjDYN_NONE @@ -259,6 +261,7 @@ class DynType(enum.IntEnum): FILTEREXACT = mujoco.mjtDyn.mjDYN_FILTEREXACT MUSCLE = mujoco.mjtDyn.mjDYN_MUSCLE USER = mujoco.mjtDyn.mjDYN_USER + DCMOTOR = mujoco.mjtDyn.mjDYN_DCMOTOR class GainType(enum.IntEnum): @@ -269,12 +272,14 @@ class GainType(enum.IntEnum): AFFINE: const + kp*length + kv*velocity MUSCLE: muscle FLV curve computed by muscle_gain USER: user-defined gain via act_gain_callback + DCMOTOR: DC motor gain """ FIXED = mujoco.mjtGain.mjGAIN_FIXED AFFINE = mujoco.mjtGain.mjGAIN_AFFINE MUSCLE = mujoco.mjtGain.mjGAIN_MUSCLE USER = mujoco.mjtGain.mjGAIN_USER + DCMOTOR = mujoco.mjtGain.mjGAIN_DCMOTOR class BiasType(enum.IntEnum): @@ -285,12 +290,14 @@ class BiasType(enum.IntEnum): AFFINE: const + kp*length + kv*velocity MUSCLE: muscle passive force computed by muscle_bias USER: user-defined bias via act_bias_callback + DCMOTOR: DC motor back-EMF bias """ NONE = mujoco.mjtBias.mjBIAS_NONE AFFINE = mujoco.mjtBias.mjBIAS_AFFINE MUSCLE = mujoco.mjtBias.mjBIAS_MUSCLE USER = mujoco.mjtBias.mjBIAS_USER + DCMOTOR = mujoco.mjtBias.mjBIAS_DCMOTOR class JointType(enum.IntEnum): @@ -546,6 +553,7 @@ class ObjType(enum.IntEnum): BODY: body XBODY: body, used to access regular frame instead of i-frame GEOM: geom + FLEX: flex SITE: site CAMERA: camera """ @@ -554,6 +562,7 @@ class ObjType(enum.IntEnum): BODY = mujoco.mjtObj.mjOBJ_BODY XBODY = mujoco.mjtObj.mjOBJ_XBODY GEOM = mujoco.mjtObj.mjOBJ_GEOM + FLEX = mujoco.mjtObj.mjOBJ_FLEX SITE = mujoco.mjtObj.mjOBJ_SITE CAMERA = mujoco.mjtObj.mjOBJ_CAMERA @@ -646,6 +655,10 @@ class vec6f(wp.types.vector(length=6, dtype=float)): pass +class vec6i(wp.types.vector(length=6, dtype=int)): + pass + + class vec8f(wp.types.vector(length=8, dtype=float)): pass @@ -921,6 +934,7 @@ class Model: jnt_pos: local anchor position (*, njnt, 3) jnt_axis: local joint axis (*, njnt, 3) jnt_stiffness: stiffness coefficient (*, njnt) + jnt_stiffnesspoly: high-order stiffness coefficients (*, njnt, 2) jnt_range: joint limits (*, njnt, 2) jnt_actfrcrange: range of total actuator force (*, njnt, 2) jnt_margin: min distance for limit detection (*, njnt) @@ -934,6 +948,7 @@ class Model: dof_frictionloss: dof friction loss (*, nv) dof_armature: dof armature inertia/mass (*, nv) dof_damping: damping coefficient (*, nv) + dof_dampingpoly: high-order damping coefficients (*, nv, 2) dof_invweight0: diag. inverse inertia in qpos0 (*, nv) tree_bodynum: number of bodies in tree (incl. root) (ntree,) tree_dofadr: start address of tree's dofs (ntree,) @@ -992,8 +1007,13 @@ class Model: flex_contype: flex contact type (nflex,) flex_conaffinity: flex contact affinity (nflex,) flex_condim: contact dimensionality (1, 3, 4, 6) (nflex,) + flex_priority: geom contact priority (nflex,) + flex_solmix: mixing coef for solref/imp in geom pair (nflex,) + flex_solref: constraint solver reference: contact (nflex, mjNREF) + flex_solimp: constraint solver impedance: contact (nflex, mjNIMP) flex_friction: friction for (slide, spin, roll) (nflex, 3) flex_margin: detect contact if dist float: return dctrl / wp.max(MJ_MINVAL, tau) +@wp.func +def dcmotor_slots(dynprm: types.vec10, gainprm: types.vec10) -> types.vec6i: + """Compute activation slot layout for a DC motor actuator. + + Each DC motor can have up to 5 optional activation states. This function + determines which states are enabled (based on nonzero parameters) and + assigns each a contiguous slot offset in the activation array. + + Returns a vec6i where: + s[0]: slew rate — enabled when dynprm[7] > 0 (slew rate limit) + s[1]: integral — enabled when gainprm[5] > 0 (integral gain ki) + s[2]: temperature — enabled when dynprm[2] > 0 (thermal resistance RT) + s[3]: bristle — enabled when dynprm[5] > 0 (LuGre stiffness sigma0) + s[4]: current — enabled when dynprm[0] > 0 (electrical time const te) + s[5]: total number of active slots (num_slots) + + Enabled slots hold a contiguous offset (0, 1, 2, ...); disabled slots + are set to -1. + """ + s = types.vec6i(-1, -1, -1, -1, -1, 0) + num_slots = 0 + if dynprm[7] > 0.0: + s[0] = num_slots + num_slots += 1 + if gainprm[5] > 0.0: + s[1] = num_slots + num_slots += 1 + if dynprm[2] > 0.0: + s[2] = num_slots + num_slots += 1 + if dynprm[5] > 0.0: + s[3] = num_slots + num_slots += 1 + if dynprm[0] > 0.0: + s[4] = num_slots + num_slots += 1 + s[5] = num_slots + return s + + +@wp.func +def lugre_stribeck(velocity: float, F_C: float, F_S: float, v_S: float) -> float: + ratio = velocity / wp.max(MJ_MINVAL, v_S) + return F_C + (F_S - F_C) * wp.exp(-ratio * ratio) + + +@wp.func +def dcmotor_voltage(u: float, length: float, velocity: float, x_I: float, gainprm: types.vec10) -> float: + input_mode = int(gainprm[8]) + Vmax = gainprm[7] + voltage = 0.0 + + if input_mode > 0: + kp = gainprm[4] + ki = gainprm[5] + kd = gainprm[6] + + if input_mode == 1: + # position mode + voltage = kp * (u - length) + ki * x_I - kd * velocity + else: + # velocity mode + voltage = kp * (u - velocity) + ki * (x_I - length) + else: + voltage = u + + if Vmax > 0.0: + voltage = wp.clamp(voltage, -Vmax, Vmax) + + return voltage + + @wp.func def inside_geom(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, geomtype: int, point: wp.vec3) -> bool: """Return True if point is inside primitive geom, False otherwise.""" @@ -630,3 +703,34 @@ def inside_geom(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, geomtype: int, point return plocal[2] < 0.0 return False + + +@wp.func +def _poly_force(linear: float, poly: wp.vec2, x: float, flg_odd: int) -> float: + x_val = wp.where(flg_odd == 1, wp.abs(x), x) + res = linear + res += poly[0] * x_val + res += poly[1] * x_val * x_val + return res + + +@wp.func +def _poly_force_deriv(linear: float, poly: wp.vec2, x: float, flg_odd: int) -> float: + x_val = wp.where(flg_odd == 1, wp.abs(x), x) + res = linear + res += 2.0 * poly[0] * x_val + res += 3.0 * poly[1] * x_val * x_val + return res + + +@wp.func +def poly_potential(linear: float, poly: wp.vec2, x: float, flg_odd: int) -> float: + x_val = wp.where(flg_odd == 1, wp.abs(x), x) + x_val2 = x_val * x_val + x_val3 = x_val2 * x_val + x_val4 = x_val3 * x_val + + res = 0.5 * linear * x_val2 + res += poly[0] * wp.static(1.0 / 3.0) * x_val3 + res += poly[1] * 0.25 * x_val4 + return res diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml b/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml index f1c4bc4a..a5caa4b4 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name="mujoco-warp" -version = "3.6.0" +version = "3.8.0.1" # TODO(team): create a distribution list authors = [ {name = "Newton Developers", email = "mujoco@deepmind.com"}, @@ -28,7 +28,7 @@ requires-python = ">=3.10" dependencies = [ "absl-py", "etils[epath]", - "mujoco>=3.6.0", + "mujoco>=3.8.0", "numpy", "warp-lang>=1.12", ] @@ -55,9 +55,10 @@ dev = [ "ruff", "pygls>=1.0.0,<2.0.0", "lsprotocol>=2023.0.1,<2024.0.0", - "mujoco>=3.6.0.dev0", + "mujoco>=3.8.0.dev0", "warp-lang>=1.11.0.dev0", "mjviser>=0.0.10", + "pillow", ] # TODO(team): cpu and cuda JAX optional dependencies are temporary, remove after we land MJX:Warp cpu = [ @@ -70,6 +71,8 @@ cuda = [ [project.scripts] mjwarp-testspeed = "mujoco_warp.testspeed:main" mjwarp-viewer = "mujoco_warp.viewer:main" +mjwarp-record = "mujoco_warp.record:main" + [project.urls] Homepage = "https://github.com/google-deepmind/mujoco_warp" diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py b/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py index 8a825da2..bac30dd3 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py @@ -24,7 +24,6 @@ Example: import copy import enum import logging -import shutil import sys import time from typing import Sequence @@ -39,11 +38,6 @@ from etils import epath import mujoco.mjx.third_party.mujoco_warp as mjw -# mjwarp-viewer has priviledged access to a few internal methods -from mujoco.mjx.third_party.mujoco_warp._src.io import find_keys -from mujoco.mjx.third_party.mujoco_warp._src.io import make_trajectory -from mujoco.mjx.third_party.mujoco_warp._src.io import override_model - class EngineOptions(enum.IntEnum): """Engine option.""" @@ -52,16 +46,9 @@ class EngineOptions(enum.IntEnum): C = 1 -_CLEAR_WARP_CACHE = flags.DEFINE_bool("clear_warp_cache", False, "Clear warp caches (kernel, LTO, CUDA compute)") _ENGINE = flags.DEFINE_enum_class("engine", EngineOptions.WARP, EngineOptions, "Simulation engine") -_NCONMAX = flags.DEFINE_integer("nconmax", None, "Maximum number of contacts.") -_NJMAX = flags.DEFINE_integer("njmax", None, "Maximum number of constraints per world.") -_NJMAX_NNZ = flags.DEFINE_integer("njmax_nnz", None, "Maximum number of non-zeros in constraint Jacobian.") -_NCCDMAX = flags.DEFINE_integer("nccdmax", None, "Maximum number of CCD contacts per world.") -_OVERRIDE = flags.DEFINE_multi_string("override", [], "Model overrides (notation: foo.bar = baz)", short_name="o") -_KEYFRAME = flags.DEFINE_integer("keyframe", 0, "keyframe to initialize simulation.") -_DEVICE = flags.DEFINE_string("device", None, "override the default Warp device") -_REPLAY = flags.DEFINE_string("replay", None, "keyframe sequence to replay, keyframe name must prefix match") +from mujoco.mjx.third_party.mujoco_warp._src import cli + _VIEWER = flags.DEFINE_enum("viewer", "mujoco", ["mujoco", "viser"], "Viewer backend (mujoco native or mjviser web)") _VIEWER_GLOBAL_STATE = {"running": True, "step_once": False} @@ -75,26 +62,6 @@ def key_callback(key: int) -> None: _VIEWER_GLOBAL_STATE["step_once"] = True -def _load_model(path: epath.Path) -> mujoco.MjModel: - if not path.exists(): - resource_path = epath.resource_path("mjx") / "third_party/mujoco_warp" / path - if not resource_path.exists(): - raise FileNotFoundError(f"file not found: {path}\nalso tried: {resource_path}") - path = resource_path - - print(f"Loading model from: {path}...") - if path.suffix == ".mjb": - return mujoco.MjModel.from_binary_path(path.as_posix()) - - spec = mujoco.MjSpec.from_file(path.as_posix()) - # check if the file has any mujoco.sdf test plugins - if any(p.plugin_name.startswith("mujoco.sdf") for p in spec.plugins): - from mujoco.mjx.third_party.mujoco_warp.test_data.collision_sdf.utils import register_sdf_plugins as register_sdf_plugins - - register_sdf_plugins(mjw) - return spec.compile() - - def _compile_step(m, d): print("Compiling physics step...", end="", flush=True) start = time.time() @@ -177,20 +144,13 @@ def _main(argv: Sequence[str]) -> None: elif len(argv) > 2: raise app.UsageError("Too many command-line arguments.") - mjm = _load_model(epath.Path(argv[1])) - mjd = mujoco.MjData(mjm) - ctrls = None - if _REPLAY.value: - keys = find_keys(mjm, _REPLAY.value) - if not keys: - raise app.UsageError(f"Key prefix not find: {_REPLAY.value}") - ctrls = make_trajectory(mjm, keys) - mujoco.mj_resetDataKeyframe(mjm, mjd, keys[0]) - elif mjm.nkey > 0 and _KEYFRAME.value > -1: - mujoco.mj_resetDataKeyframe(mjm, mjd, _KEYFRAME.value) + wp.config.quiet = flags.FLAGS["verbosity"].value < 1 + wp.init() + + mjm = cli.load_model(epath.Path(argv[1])) + m, d, rc, ctrls = cli.init_structs(mjw.step, mjm) if _ENGINE.value == EngineOptions.C: - override_model(mjm, _OVERRIDE.value) print( f" nbody: {mjm.nbody} nv: {mjm.nv} ngeom: {mjm.ngeom} nu: {mjm.nu}\n" f" solver: {mujoco.mjtSolver(mjm.opt.solver).name} cone: {mujoco.mjtCone(mjm.opt.cone).name}" @@ -199,22 +159,8 @@ def _main(argv: Sequence[str]) -> None: ) print(f"MuJoCo C simulating with dt = {mjm.opt.timestep:.3f}...") else: - wp.config.quiet = flags.FLAGS["verbosity"].value < 1 - wp.init() - wp.set_device(_DEVICE.value) - if _CLEAR_WARP_CACHE.value: - wp.clear_kernel_cache() - wp.clear_lto_cache() - # Clear CUDA compute cache for truly cold start JIT - compute_cache = epath.Path("~/.nv/ComputeCache").expanduser() - if compute_cache.exists(): - shutil.rmtree(compute_cache) - compute_cache.mkdir() + wp.set_device(cli.DEVICE.value) - override_model(mjm, _OVERRIDE.value) - m = mjw.put_model(mjm) - override_model(m, _OVERRIDE.value) - d = mjw.put_data(mjm, mjd, nconmax=_NCONMAX.value, njmax=_NJMAX.value, njmax_nnz=_NJMAX_NNZ.value, nccdmax=_NCCDMAX.value) graph = _compile_step(m, d) if wp.get_device().is_cuda else None if graph is None: mjw.step(m, d) # warmup step @@ -238,6 +184,8 @@ def _main(argv: Sequence[str]) -> None: else: step_fn = _make_c_step_fn(ctrls) + mjd = mujoco.MjData(mjm) + mjw.get_data_into(mjd, mjm, d) if _VIEWER.value == "viser": _run_viser_viewer(mjm, mjd, step_fn) else: @@ -249,6 +197,8 @@ def main(): # pyproject bin scripts break this assumption, so manually set argv and docstring sys.argv[0] = "mujoco_warp.viewer" sys.modules["__main__"].__doc__ = __doc__ + # default to single world with no noise + flags.FLAGS.set_default("nworld", 1) app.run(_main) diff --git a/mjx/mujoco/mjx/warp/bvh.py b/mjx/mujoco/mjx/warp/bvh.py index edde9732..d80b12fd 100644 --- a/mjx/mujoco/mjx/warp/bvh.py +++ b/mjx/mujoco/mjx/warp/bvh.py @@ -48,7 +48,6 @@ _cb = mjwp_types.Callback( **{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init} ) - @ffi.format_args_for_warp def _refit_bvh_shim( # Model diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py index 530de314..dbecb9c1 100644 --- a/mjx/mujoco/mjx/warp/collision_driver.py +++ b/mjx/mujoco/mjx/warp/collision_driver.py @@ -61,11 +61,16 @@ def _collision_shim( flex_elemdataadr: wp.array[int], flex_elemnum: wp.array[int], flex_friction: wp.array[wp.vec3], + flex_gap: wp.array[float], flex_margin: wp.array[float], + flex_priority: wp.array[int], flex_radius: wp.array[float], flex_shell: wp.array[int], flex_shelldataadr: wp.array[int], flex_shellnum: wp.array[int], + flex_solimp: wp.array[mjwp_types.vec5], + flex_solmix: wp.array[float], + flex_solref: wp.array[wp.vec2], flex_vertadr: wp.array[int], flex_vertflexid: wp.array[int], geom_aabb: wp.array3d[wp.vec3], @@ -136,7 +141,6 @@ def _collision_shim( opt__ccd_iterations: int, opt__ccd_tolerance: wp.array[float], opt__disableflags: int, - opt__enableflags: int, opt__sdf_initpoints: int, opt__sdf_iterations: int, # Data @@ -179,11 +183,16 @@ def _collision_shim( _m.flex_elemdataadr = flex_elemdataadr _m.flex_elemnum = flex_elemnum _m.flex_friction = flex_friction + _m.flex_gap = flex_gap _m.flex_margin = flex_margin + _m.flex_priority = flex_priority _m.flex_radius = flex_radius _m.flex_shell = flex_shell _m.flex_shelldataadr = flex_shelldataadr _m.flex_shellnum = flex_shellnum + _m.flex_solimp = flex_solimp + _m.flex_solmix = flex_solmix + _m.flex_solref = flex_solref _m.flex_vertadr = flex_vertadr _m.flex_vertflexid = flex_vertflexid _m.geom_aabb = geom_aabb @@ -245,7 +254,6 @@ def _collision_shim( _m.opt.ccd_iterations = opt__ccd_iterations _m.opt.ccd_tolerance = opt__ccd_tolerance _m.opt.disableflags = opt__disableflags - _m.opt.enableflags = opt__enableflags _m.opt.sdf_initpoints = opt__sdf_initpoints _m.opt.sdf_iterations = opt__sdf_iterations _m.pair_dim = pair_dim @@ -366,11 +374,16 @@ def _collision_jax_impl(m: types.Model, d: types.Data): m._impl.flex_elemdataadr, m._impl.flex_elemnum, m._impl.flex_friction, + m._impl.flex_gap, m._impl.flex_margin, + m._impl.flex_priority, m._impl.flex_radius, m._impl.flex_shell, m._impl.flex_shelldataadr, m._impl.flex_shellnum, + m._impl.flex_solimp, + m._impl.flex_solmix, + m._impl.flex_solref, m.flex_vertadr, m._impl.flex_vertflexid, m.geom_aabb, @@ -441,7 +454,6 @@ def _collision_jax_impl(m: types.Model, d: types.Data): m.opt._impl.ccd_iterations, m.opt._impl.ccd_tolerance, m.opt.disableflags, - m.opt.enableflags, m.opt._impl.sdf_initpoints, m.opt._impl.sdf_iterations, d._impl.naccdmax, diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index 98459fbb..2f46f1f3 100644 --- a/mjx/mujoco/mjx/warp/forward.py +++ b/mjx/mujoco/mjx/warp/forward.py @@ -88,6 +88,7 @@ def _forward_shim( body_invweight0: wp.array2d[wp.vec2], body_ipos: wp.array2d[wp.vec3], body_iquat: wp.array2d[wp.quat], + body_isdofancestor: wp.array2d[int], body_jntadr: wp.array[int], body_jntnum: wp.array[int], body_mass: wp.array2d[float], @@ -116,6 +117,7 @@ def _forward_shim( dof_armature: wp.array2d[float], dof_bodyid: wp.array[int], dof_damping: wp.array2d[float], + dof_dampingpoly: wp.array2d[wp.vec2], dof_frictionloss: wp.array2d[float], dof_invweight0: wp.array2d[float], dof_jntid: wp.array[int], @@ -155,11 +157,16 @@ def _forward_shim( flex_elemedgeadr: wp.array[int], flex_elemnum: wp.array[int], flex_friction: wp.array[wp.vec3], + flex_gap: wp.array[float], flex_margin: wp.array[float], + flex_priority: wp.array[int], flex_radius: wp.array[float], flex_shell: wp.array[int], flex_shelldataadr: wp.array[int], flex_shellnum: wp.array[int], + flex_solimp: wp.array[mjwp_types.vec5], + flex_solmix: wp.array[float], + flex_solref: wp.array[wp.vec2], flex_stiffness: wp.array2d[float], flex_vert: wp.array[wp.vec3], flex_vertadr: wp.array[int], @@ -218,6 +225,7 @@ def _forward_shim( jnt_solimp: wp.array2d[mjwp_types.vec5], jnt_solref: wp.array2d[wp.vec2], jnt_stiffness: wp.array2d[float], + jnt_stiffnesspoly: wp.array2d[wp.vec2], jnt_type: wp.array[int], light_bodyid: wp.array[int], light_dir: wp.array2d[wp.vec3], @@ -352,6 +360,7 @@ def _forward_shim( tendon_adr: wp.array[int], tendon_armature: wp.array2d[float], tendon_damping: wp.array2d[float], + tendon_dampingpoly: wp.array2d[wp.vec2], tendon_frictionloss: wp.array2d[float], tendon_geom_adr: wp.array[int], tendon_invweight0: wp.array2d[float], @@ -368,6 +377,7 @@ def _forward_shim( tendon_solref_fri: wp.array2d[wp.vec2], tendon_solref_lim: wp.array2d[wp.vec2], tendon_stiffness: wp.array2d[float], + tendon_stiffnesspoly: wp.array2d[wp.vec2], wrap_geom_adr: wp.array[int], wrap_jnt_adr: wp.array[int], wrap_objid: wp.array[int], @@ -509,6 +519,7 @@ def _forward_shim( efc__J_colind: wp.array3d[int], efc__J_rowadr: wp.array2d[int], efc__J_rownnz: wp.array2d[int], + efc__Jqvel: wp.array2d[float], efc__Ma: wp.array2d[float], efc__aref: wp.array2d[float], efc__force: wp.array2d[float], @@ -562,6 +573,7 @@ def _forward_shim( _m.body_invweight0 = body_invweight0 _m.body_ipos = body_ipos _m.body_iquat = body_iquat + _m.body_isdofancestor = body_isdofancestor _m.body_jntadr = body_jntadr _m.body_jntnum = body_jntnum _m.body_mass = body_mass @@ -590,6 +602,7 @@ def _forward_shim( _m.dof_armature = dof_armature _m.dof_bodyid = dof_bodyid _m.dof_damping = dof_damping + _m.dof_dampingpoly = dof_dampingpoly _m.dof_frictionloss = dof_frictionloss _m.dof_invweight0 = dof_invweight0 _m.dof_jntid = dof_jntid @@ -629,11 +642,16 @@ def _forward_shim( _m.flex_elemedgeadr = flex_elemedgeadr _m.flex_elemnum = flex_elemnum _m.flex_friction = flex_friction + _m.flex_gap = flex_gap _m.flex_margin = flex_margin + _m.flex_priority = flex_priority _m.flex_radius = flex_radius _m.flex_shell = flex_shell _m.flex_shelldataadr = flex_shelldataadr _m.flex_shellnum = flex_shellnum + _m.flex_solimp = flex_solimp + _m.flex_solmix = flex_solmix + _m.flex_solref = flex_solref _m.flex_stiffness = flex_stiffness _m.flex_vert = flex_vert _m.flex_vertadr = flex_vertadr @@ -692,6 +710,7 @@ def _forward_shim( _m.jnt_solimp = jnt_solimp _m.jnt_solref = jnt_solref _m.jnt_stiffness = jnt_stiffness + _m.jnt_stiffnesspoly = jnt_stiffnesspoly _m.jnt_type = jnt_type _m.light_bodyid = light_bodyid _m.light_dir = light_dir @@ -853,6 +872,7 @@ def _forward_shim( _m.tendon_adr = tendon_adr _m.tendon_armature = tendon_armature _m.tendon_damping = tendon_damping + _m.tendon_dampingpoly = tendon_dampingpoly _m.tendon_frictionloss = tendon_frictionloss _m.tendon_geom_adr = tendon_geom_adr _m.tendon_invweight0 = tendon_invweight0 @@ -869,6 +889,7 @@ def _forward_shim( _m.tendon_solref_fri = tendon_solref_fri _m.tendon_solref_lim = tendon_solref_lim _m.tendon_stiffness = tendon_stiffness + _m.tendon_stiffnesspoly = tendon_stiffnesspoly _m.wrap_geom_adr = wrap_geom_adr _m.wrap_jnt_adr = wrap_jnt_adr _m.wrap_objid = wrap_objid @@ -914,6 +935,7 @@ def _forward_shim( _d.efc.J_colind = efc__J_colind _d.efc.J_rowadr = efc__J_rowadr _d.efc.J_rownnz = efc__J_rownnz + _d.efc.Jqvel = efc__Jqvel _d.efc.Ma = efc__Ma _d.efc.aref = efc__aref _d.efc.force = efc__force @@ -1090,6 +1112,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'efc__J_colind': d._impl.efc__J_colind.shape, 'efc__J_rowadr': d._impl.efc__J_rowadr.shape, 'efc__J_rownnz': d._impl.efc__J_rownnz.shape, + 'efc__Jqvel': d._impl.efc__Jqvel.shape, 'efc__Ma': d._impl.efc__Ma.shape, 'efc__aref': d._impl.efc__aref.shape, 'efc__force': d._impl.efc__force.shape, @@ -1103,7 +1126,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): } jf = ffi.jax_callable_variadic_tuple( _forward_shim, - num_outputs=102, + num_outputs=103, output_dims=output_dims, vmap_method=None, in_out_argnames=set([ @@ -1199,6 +1222,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'efc__J_colind', 'efc__J_rowadr', 'efc__J_rownnz', + 'efc__Jqvel', 'efc__Ma', 'efc__aref', 'efc__force', @@ -1249,6 +1273,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'cvel', 'dof_armature', 'dof_damping', + 'dof_dampingpoly', 'dof_frictionloss', 'dof_invweight0', 'dof_solimp', @@ -1281,6 +1306,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'jnt_solimp', 'jnt_solref', 'jnt_stiffness', + 'jnt_stiffnesspoly', 'light_dir', 'light_dir0', 'light_pos', @@ -1328,6 +1354,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'tendon_actfrcrange', 'tendon_armature', 'tendon_damping', + 'tendon_dampingpoly', 'tendon_frictionloss', 'tendon_invweight0', 'tendon_length0', @@ -1339,6 +1366,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'tendon_solref_fri', 'tendon_solref_lim', 'tendon_stiffness', + 'tendon_stiffnesspoly', 'time', 'xanchor', 'xaxis', @@ -1425,6 +1453,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.body_invweight0, m.body_ipos, m.body_iquat, + m._impl.body_isdofancestor, m.body_jntadr, m.body_jntnum, m.body_mass, @@ -1453,6 +1482,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.dof_armature, m.dof_bodyid, m.dof_damping, + m.dof_dampingpoly, m.dof_frictionloss, m.dof_invweight0, m.dof_jntid, @@ -1492,11 +1522,16 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.flex_elemedgeadr, m._impl.flex_elemnum, m._impl.flex_friction, + m._impl.flex_gap, m._impl.flex_margin, + m._impl.flex_priority, m._impl.flex_radius, m._impl.flex_shell, m._impl.flex_shelldataadr, m._impl.flex_shellnum, + m._impl.flex_solimp, + m._impl.flex_solmix, + m._impl.flex_solref, m._impl.flex_stiffness, m._impl.flex_vert, m.flex_vertadr, @@ -1555,6 +1590,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.jnt_solimp, m.jnt_solref, m.jnt_stiffness, + m.jnt_stiffnesspoly, m.jnt_type, m._impl.light_bodyid, m.light_dir, @@ -1689,6 +1725,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.tendon_adr, m.tendon_armature, m.tendon_damping, + m.tendon_dampingpoly, m.tendon_frictionloss, m._impl.tendon_geom_adr, m.tendon_invweight0, @@ -1705,6 +1742,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.tendon_solref_fri, m.tendon_solref_lim, m.tendon_stiffness, + m.tendon_stiffnesspoly, m._impl.wrap_geom_adr, m._impl.wrap_jnt_adr, m.wrap_objid, @@ -1845,6 +1883,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.efc__J_colind, d._impl.efc__J_rowadr, d._impl.efc__J_rownnz, + d._impl.efc__Jqvel, d._impl.efc__Ma, d._impl.efc__aref, d._impl.efc__force, @@ -1949,16 +1988,17 @@ def _forward_jax_impl(m: types.Model, d: types.Data): '_impl.efc__J_colind': out[89], '_impl.efc__J_rowadr': out[90], '_impl.efc__J_rownnz': out[91], - '_impl.efc__Ma': out[92], - '_impl.efc__aref': out[93], - '_impl.efc__force': out[94], - '_impl.efc__frictionloss': out[95], - '_impl.efc__id': out[96], - '_impl.efc__margin': out[97], - '_impl.efc__pos': out[98], - '_impl.efc__state': out[99], - '_impl.efc__type': out[100], - '_impl.efc__vel': out[101], + '_impl.efc__Jqvel': out[92], + '_impl.efc__Ma': out[93], + '_impl.efc__aref': out[94], + '_impl.efc__force': out[95], + '_impl.efc__frictionloss': out[96], + '_impl.efc__id': out[97], + '_impl.efc__margin': out[98], + '_impl.efc__pos': out[99], + '_impl.efc__state': out[100], + '_impl.efc__type': out[101], + '_impl.efc__vel': out[102], }) return d @@ -2017,6 +2057,7 @@ def _step_shim( body_invweight0: wp.array2d[wp.vec2], body_ipos: wp.array2d[wp.vec3], body_iquat: wp.array2d[wp.quat], + body_isdofancestor: wp.array2d[int], body_jntadr: wp.array[int], body_jntnum: wp.array[int], body_mass: wp.array2d[float], @@ -2045,6 +2086,7 @@ def _step_shim( dof_armature: wp.array2d[float], dof_bodyid: wp.array[int], dof_damping: wp.array2d[float], + dof_dampingpoly: wp.array2d[wp.vec2], dof_frictionloss: wp.array2d[float], dof_invweight0: wp.array2d[float], dof_jntid: wp.array[int], @@ -2084,11 +2126,16 @@ def _step_shim( flex_elemedgeadr: wp.array[int], flex_elemnum: wp.array[int], flex_friction: wp.array[wp.vec3], + flex_gap: wp.array[float], flex_margin: wp.array[float], + flex_priority: wp.array[int], flex_radius: wp.array[float], flex_shell: wp.array[int], flex_shelldataadr: wp.array[int], flex_shellnum: wp.array[int], + flex_solimp: wp.array[mjwp_types.vec5], + flex_solmix: wp.array[float], + flex_solref: wp.array[wp.vec2], flex_stiffness: wp.array2d[float], flex_vert: wp.array[wp.vec3], flex_vertadr: wp.array[int], @@ -2147,6 +2194,7 @@ def _step_shim( jnt_solimp: wp.array2d[mjwp_types.vec5], jnt_solref: wp.array2d[wp.vec2], jnt_stiffness: wp.array2d[float], + jnt_stiffnesspoly: wp.array2d[wp.vec2], jnt_type: wp.array[int], light_bodyid: wp.array[int], light_dir: wp.array2d[wp.vec3], @@ -2282,6 +2330,7 @@ def _step_shim( tendon_adr: wp.array[int], tendon_armature: wp.array2d[float], tendon_damping: wp.array2d[float], + tendon_dampingpoly: wp.array2d[wp.vec2], tendon_frictionloss: wp.array2d[float], tendon_geom_adr: wp.array[int], tendon_invweight0: wp.array2d[float], @@ -2298,6 +2347,7 @@ def _step_shim( tendon_solref_fri: wp.array2d[wp.vec2], tendon_solref_lim: wp.array2d[wp.vec2], tendon_stiffness: wp.array2d[float], + tendon_stiffnesspoly: wp.array2d[wp.vec2], wrap_geom_adr: wp.array[int], wrap_jnt_adr: wp.array[int], wrap_objid: wp.array[int], @@ -2440,6 +2490,7 @@ def _step_shim( efc__J_colind: wp.array3d[int], efc__J_rowadr: wp.array2d[int], efc__J_rownnz: wp.array2d[int], + efc__Jqvel: wp.array2d[float], efc__Ma: wp.array2d[float], efc__aref: wp.array2d[float], efc__force: wp.array2d[float], @@ -2493,6 +2544,7 @@ def _step_shim( _m.body_invweight0 = body_invweight0 _m.body_ipos = body_ipos _m.body_iquat = body_iquat + _m.body_isdofancestor = body_isdofancestor _m.body_jntadr = body_jntadr _m.body_jntnum = body_jntnum _m.body_mass = body_mass @@ -2521,6 +2573,7 @@ def _step_shim( _m.dof_armature = dof_armature _m.dof_bodyid = dof_bodyid _m.dof_damping = dof_damping + _m.dof_dampingpoly = dof_dampingpoly _m.dof_frictionloss = dof_frictionloss _m.dof_invweight0 = dof_invweight0 _m.dof_jntid = dof_jntid @@ -2560,11 +2613,16 @@ def _step_shim( _m.flex_elemedgeadr = flex_elemedgeadr _m.flex_elemnum = flex_elemnum _m.flex_friction = flex_friction + _m.flex_gap = flex_gap _m.flex_margin = flex_margin + _m.flex_priority = flex_priority _m.flex_radius = flex_radius _m.flex_shell = flex_shell _m.flex_shelldataadr = flex_shelldataadr _m.flex_shellnum = flex_shellnum + _m.flex_solimp = flex_solimp + _m.flex_solmix = flex_solmix + _m.flex_solref = flex_solref _m.flex_stiffness = flex_stiffness _m.flex_vert = flex_vert _m.flex_vertadr = flex_vertadr @@ -2623,6 +2681,7 @@ def _step_shim( _m.jnt_solimp = jnt_solimp _m.jnt_solref = jnt_solref _m.jnt_stiffness = jnt_stiffness + _m.jnt_stiffnesspoly = jnt_stiffnesspoly _m.jnt_type = jnt_type _m.light_bodyid = light_bodyid _m.light_dir = light_dir @@ -2786,6 +2845,7 @@ def _step_shim( _m.tendon_adr = tendon_adr _m.tendon_armature = tendon_armature _m.tendon_damping = tendon_damping + _m.tendon_dampingpoly = tendon_dampingpoly _m.tendon_frictionloss = tendon_frictionloss _m.tendon_geom_adr = tendon_geom_adr _m.tendon_invweight0 = tendon_invweight0 @@ -2802,6 +2862,7 @@ def _step_shim( _m.tendon_solref_fri = tendon_solref_fri _m.tendon_solref_lim = tendon_solref_lim _m.tendon_stiffness = tendon_stiffness + _m.tendon_stiffnesspoly = tendon_stiffnesspoly _m.wrap_geom_adr = wrap_geom_adr _m.wrap_jnt_adr = wrap_jnt_adr _m.wrap_objid = wrap_objid @@ -2847,6 +2908,7 @@ def _step_shim( _d.efc.J_colind = efc__J_colind _d.efc.J_rowadr = efc__J_rowadr _d.efc.J_rownnz = efc__J_rownnz + _d.efc.Jqvel = efc__Jqvel _d.efc.Ma = efc__Ma _d.efc.aref = efc__aref _d.efc.force = efc__force @@ -3027,6 +3089,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'efc__J_colind': d._impl.efc__J_colind.shape, 'efc__J_rowadr': d._impl.efc__J_rowadr.shape, 'efc__J_rownnz': d._impl.efc__J_rownnz.shape, + 'efc__Jqvel': d._impl.efc__Jqvel.shape, 'efc__Ma': d._impl.efc__Ma.shape, 'efc__aref': d._impl.efc__aref.shape, 'efc__force': d._impl.efc__force.shape, @@ -3040,7 +3103,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): } jf = ffi.jax_callable_variadic_tuple( _step_shim, - num_outputs=106, + num_outputs=107, output_dims=output_dims, vmap_method=None, in_out_argnames=set([ @@ -3140,6 +3203,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'efc__J_colind', 'efc__J_rowadr', 'efc__J_rownnz', + 'efc__Jqvel', 'efc__Ma', 'efc__aref', 'efc__force', @@ -3190,6 +3254,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'cvel', 'dof_armature', 'dof_damping', + 'dof_dampingpoly', 'dof_frictionloss', 'dof_invweight0', 'dof_solimp', @@ -3222,6 +3287,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'jnt_solimp', 'jnt_solref', 'jnt_stiffness', + 'jnt_stiffnesspoly', 'light_dir', 'light_dir0', 'light_pos', @@ -3269,6 +3335,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'tendon_actfrcrange', 'tendon_armature', 'tendon_damping', + 'tendon_dampingpoly', 'tendon_frictionloss', 'tendon_invweight0', 'tendon_length0', @@ -3280,6 +3347,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'tendon_solref_fri', 'tendon_solref_lim', 'tendon_stiffness', + 'tendon_stiffnesspoly', 'time', 'xanchor', 'xaxis', @@ -3370,6 +3438,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.body_invweight0, m.body_ipos, m.body_iquat, + m._impl.body_isdofancestor, m.body_jntadr, m.body_jntnum, m.body_mass, @@ -3398,6 +3467,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.dof_armature, m.dof_bodyid, m.dof_damping, + m.dof_dampingpoly, m.dof_frictionloss, m.dof_invweight0, m.dof_jntid, @@ -3437,11 +3507,16 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.flex_elemedgeadr, m._impl.flex_elemnum, m._impl.flex_friction, + m._impl.flex_gap, m._impl.flex_margin, + m._impl.flex_priority, m._impl.flex_radius, m._impl.flex_shell, m._impl.flex_shelldataadr, m._impl.flex_shellnum, + m._impl.flex_solimp, + m._impl.flex_solmix, + m._impl.flex_solref, m._impl.flex_stiffness, m._impl.flex_vert, m.flex_vertadr, @@ -3500,6 +3575,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.jnt_solimp, m.jnt_solref, m.jnt_stiffness, + m.jnt_stiffnesspoly, m.jnt_type, m._impl.light_bodyid, m.light_dir, @@ -3635,6 +3711,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.tendon_adr, m.tendon_armature, m.tendon_damping, + m.tendon_dampingpoly, m.tendon_frictionloss, m._impl.tendon_geom_adr, m.tendon_invweight0, @@ -3651,6 +3728,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.tendon_solref_fri, m.tendon_solref_lim, m.tendon_stiffness, + m.tendon_stiffnesspoly, m._impl.wrap_geom_adr, m._impl.wrap_jnt_adr, m.wrap_objid, @@ -3792,6 +3870,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.efc__J_colind, d._impl.efc__J_rowadr, d._impl.efc__J_rownnz, + d._impl.efc__Jqvel, d._impl.efc__Ma, d._impl.efc__aref, d._impl.efc__force, @@ -3900,16 +3979,17 @@ def _step_jax_impl(m: types.Model, d: types.Data): '_impl.efc__J_colind': out[93], '_impl.efc__J_rowadr': out[94], '_impl.efc__J_rownnz': out[95], - '_impl.efc__Ma': out[96], - '_impl.efc__aref': out[97], - '_impl.efc__force': out[98], - '_impl.efc__frictionloss': out[99], - '_impl.efc__id': out[100], - '_impl.efc__margin': out[101], - '_impl.efc__pos': out[102], - '_impl.efc__state': out[103], - '_impl.efc__type': out[104], - '_impl.efc__vel': out[105], + '_impl.efc__Jqvel': out[96], + '_impl.efc__Ma': out[97], + '_impl.efc__aref': out[98], + '_impl.efc__force': out[99], + '_impl.efc__frictionloss': out[100], + '_impl.efc__id': out[101], + '_impl.efc__margin': out[102], + '_impl.efc__pos': out[103], + '_impl.efc__state': out[104], + '_impl.efc__type': out[105], + '_impl.efc__vel': out[106], }) return d diff --git a/mjx/mujoco/mjx/warp/render.py b/mjx/mujoco/mjx/warp/render.py index c97003f0..57de37f9 100644 --- a/mjx/mujoco/mjx/warp/render.py +++ b/mjx/mujoco/mjx/warp/render.py @@ -48,7 +48,6 @@ _cb = mjwp_types.Callback( **{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init} ) - @ffi.format_args_for_warp def _render_shim( # Model diff --git a/mjx/mujoco/mjx/warp/smooth.py b/mjx/mujoco/mjx/warp/smooth.py index 020b9392..ce765d8b 100644 --- a/mjx/mujoco/mjx/warp/smooth.py +++ b/mjx/mujoco/mjx/warp/smooth.py @@ -46,7 +46,6 @@ _cb = mjwp_types.Callback( **{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init} ) - @ffi.format_args_for_warp def _kinematics_shim( # Model diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index 14e416a8..81e77ec8 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -76,6 +76,7 @@ class BlockDim: cholesky_factorize: int cholesky_factorize_solve: int cholesky_solve: int + contact_jac_tiled: int contact_sort: int energy_vel_kinetic: int euler_dense: int @@ -129,6 +130,7 @@ class ModelWarp(PyTreeNode): body_branch_start: np.ndarray body_branches: np.ndarray body_fluid_ellipsoid: np.ndarray + body_isdofancestor: np.ndarray body_tree: Tuple[np.ndarray, ...] callback: Callback cam_projection: np.ndarray @@ -158,11 +160,16 @@ class ModelWarp(PyTreeNode): flex_elemedgeadr: np.ndarray flex_elemnum: np.ndarray flex_friction: np.ndarray + flex_gap: np.ndarray flex_margin: np.ndarray + flex_priority: np.ndarray flex_radius: np.ndarray flex_shell: np.ndarray flex_shelldataadr: np.ndarray flex_shellnum: np.ndarray + flex_solimp: np.ndarray + flex_solmix: np.ndarray + flex_solref: np.ndarray flex_stiffness: np.ndarray flex_vert: np.ndarray flex_vertbodyid: np.ndarray @@ -302,6 +309,7 @@ class DataWarp(PyTreeNode): efc__J_colind: jax.Array efc__J_rowadr: jax.Array efc__J_rownnz: jax.Array + efc__Jqvel: jax.Array efc__Ma: jax.Array efc__aref: jax.Array efc__force: jax.Array @@ -442,6 +450,7 @@ _NDIM = { 'efc__J_colind': 3, 'efc__J_rowadr': 2, 'efc__J_rownnz': 2, + 'efc__Jqvel': 2, 'efc__Ma': 2, 'efc__aref': 2, 'efc__force': 2, @@ -554,6 +563,7 @@ _NDIM = { 'block_dim__cholesky_factorize': 0, 'block_dim__cholesky_factorize_solve': 0, 'block_dim__cholesky_solve': 0, + 'block_dim__contact_jac_tiled': 0, 'block_dim__contact_sort': 0, 'block_dim__energy_vel_kinetic': 0, 'block_dim__euler_dense': 0, @@ -580,6 +590,7 @@ _NDIM = { 'body_invweight0': 3, 'body_ipos': 3, 'body_iquat': 3, + 'body_isdofancestor': 2, 'body_jntadr': 1, 'body_jntnum': 1, 'body_mass': 2, @@ -610,6 +621,7 @@ _NDIM = { 'dof_armature': 2, 'dof_bodyid': 1, 'dof_damping': 2, + 'dof_dampingpoly': 3, 'dof_frictionloss': 2, 'dof_invweight0': 2, 'dof_jntid': 1, @@ -651,11 +663,16 @@ _NDIM = { 'flex_elemedgeadr': 1, 'flex_elemnum': 1, 'flex_friction': 2, + 'flex_gap': 1, 'flex_margin': 1, + 'flex_priority': 1, 'flex_radius': 1, 'flex_shell': 1, 'flex_shelldataadr': 1, 'flex_shellnum': 1, + 'flex_solimp': 2, + 'flex_solmix': 1, + 'flex_solref': 2, 'flex_stiffness': 2, 'flex_vert': 2, 'flex_vertadr': 1, @@ -715,6 +732,7 @@ _NDIM = { 'jnt_solimp': 3, 'jnt_solref': 3, 'jnt_stiffness': 2, + 'jnt_stiffnesspoly': 3, 'jnt_type': 1, 'light_active': 2, 'light_bodyid': 1, @@ -910,6 +928,7 @@ _NDIM = { 'tendon_adr': 1, 'tendon_armature': 2, 'tendon_damping': 2, + 'tendon_dampingpoly': 3, 'tendon_frictionloss': 2, 'tendon_geom_adr': 1, 'tendon_invweight0': 2, @@ -927,6 +946,7 @@ _NDIM = { 'tendon_solref_fri': 3, 'tendon_solref_lim': 3, 'tendon_stiffness': 2, + 'tendon_stiffnesspoly': 3, 'tree_bodynum': 1, 'tree_dofadr': 1, 'tree_dofnum': 1, @@ -1008,6 +1028,7 @@ _BATCH_DIM = { 'efc__J_colind': True, 'efc__J_rowadr': True, 'efc__J_rownnz': True, + 'efc__Jqvel': True, 'efc__Ma': True, 'efc__aref': True, 'efc__force': True, @@ -1120,6 +1141,7 @@ _BATCH_DIM = { 'block_dim__cholesky_factorize': False, 'block_dim__cholesky_factorize_solve': False, 'block_dim__cholesky_solve': False, + 'block_dim__contact_jac_tiled': False, 'block_dim__contact_sort': False, 'block_dim__energy_vel_kinetic': False, 'block_dim__euler_dense': False, @@ -1146,6 +1168,7 @@ _BATCH_DIM = { 'body_invweight0': True, 'body_ipos': True, 'body_iquat': True, + 'body_isdofancestor': False, 'body_jntadr': False, 'body_jntnum': False, 'body_mass': True, @@ -1176,6 +1199,7 @@ _BATCH_DIM = { 'dof_armature': True, 'dof_bodyid': False, 'dof_damping': True, + 'dof_dampingpoly': True, 'dof_frictionloss': True, 'dof_invweight0': True, 'dof_jntid': False, @@ -1217,11 +1241,16 @@ _BATCH_DIM = { 'flex_elemedgeadr': False, 'flex_elemnum': False, 'flex_friction': False, + 'flex_gap': False, 'flex_margin': False, + 'flex_priority': False, 'flex_radius': False, 'flex_shell': False, 'flex_shelldataadr': False, 'flex_shellnum': False, + 'flex_solimp': False, + 'flex_solmix': False, + 'flex_solref': False, 'flex_stiffness': False, 'flex_vert': False, 'flex_vertadr': False, @@ -1281,6 +1310,7 @@ _BATCH_DIM = { 'jnt_solimp': True, 'jnt_solref': True, 'jnt_stiffness': True, + 'jnt_stiffnesspoly': True, 'jnt_type': False, 'light_active': True, 'light_bodyid': False, @@ -1476,6 +1506,7 @@ _BATCH_DIM = { 'tendon_adr': False, 'tendon_armature': True, 'tendon_damping': True, + 'tendon_dampingpoly': True, 'tendon_frictionloss': True, 'tendon_geom_adr': False, 'tendon_invweight0': True, @@ -1493,6 +1524,7 @@ _BATCH_DIM = { 'tendon_solref_fri': True, 'tendon_solref_lim': True, 'tendon_stiffness': True, + 'tendon_stiffnesspoly': True, 'tree_bodynum': False, 'tree_dofadr': False, 'tree_dofnum': False,