Import google-deepmind/mujoco_warp from GitHub.

PiperOrigin-RevId: 941234486
Change-Id: Ie5d8cc0c9975e66f9a81bb6a86e33f640a5d19e3
This commit is contained in:
Taylor Howell
2026-07-01 12:25:55 -07:00
committed by Copybara-Service
parent bb80b55ae1
commit 1f6cf4035c
37 changed files with 10160 additions and 6959 deletions
+25 -13
View File
@@ -174,9 +174,6 @@ def _strip_weak_type(tree):
def _wp_to_np_type(wp_field: Any, name: str = '') -> Any:
"""Converts a warp type to an MJX compatible numpy type."""
if hasattr(wp_field, '_is_batched'):
wp_field.strides = wp_field.strides[1:]
wp_field.shape = wp_field.shape[1:]
# warp scalars
wp_dtype = type(wp_field)
if wp_dtype in wp._src.types.warp_type_to_np_dtype:
@@ -202,7 +199,13 @@ def _wp_to_np_type(wp_field: Any, name: str = '') -> Any:
wp_field[0], mjwp_types.TileSet
):
return tuple(
mjxw.types.TileSet(wp_field[i].adr.numpy(), wp_field[i].size)
mjxw.types.TileSet(
wp_field[i].adr.numpy(),
wp_field[i].size,
wp_field[i].elemid.numpy()
if wp_field[i].elemid is not None
else np.zeros(0, dtype=np.int32),
)
for i in range(len(wp_field))
)
if isinstance(wp_field, mjwp_types.BlockDim):
@@ -272,11 +275,15 @@ def _put_option(
if impl == types.Impl.WARP:
if not mjxw.mjwp_io.ENABLE_ISLANDS:
fields['disableflags'] = types.DisableBit(
fields['disableflags'] | mjwp_types.DisableBit.ISLAND
)
impl_fields = {k: _wp_to_np_type(v) for k, v in impl_fields.items()}
impl_fields = {
k: (
v_np.reshape(v_np.shape[1:])
if isinstance(v_np := _wp_to_np_type(v, k), np.ndarray)
and mjxw.types._BATCH_DIM['Option'].get(k, False) # pylint: disable=protected-access
else v_np
)
for k, v in impl_fields.items()
}
return types.Option(**fields, _impl=mjxw.types.OptionWarp(**impl_fields))
raise NotImplementedError(f'Unsupported implementation: {impl}')
@@ -460,6 +467,8 @@ def _put_model_warp(
if not hasattr(mw, k) or k in ('stat', 'opt'):
continue
field = _wp_to_np_type(getattr(mw, k), k)
if mjxw.types._BATCH_DIM['Model'].get(k, False): # pylint: disable=protected-access
field = field.reshape(field.shape[1:])
if k == 'geom_dataid' and field.ndim > 1:
# Batched geom_dataid is not supported in MJX.
field = field[0]
@@ -468,6 +477,8 @@ def _put_model_warp(
impl_fields = {}
for k in mjxw.types.ModelWarp.__annotations__.keys():
field = _wp_to_np_type(getattr(mw, k), k)
if mjxw.types._BATCH_DIM['Model'].get(k, False): # pylint: disable=protected-access
field = field.reshape(field.shape[1:])
impl_fields[k] = field
model = types.Model(
@@ -1298,18 +1309,19 @@ def _get_data_into_warp(
'ten_J',
'flexedge_J',
'M',
'map_efc2iefc', 'map_iefc2efc'
):
continue
if field.name.startswith('efc_'):
continue
# Skip island fields; host MjData arena memory cannot be reallocated from Python if not stepped.
if field.name.startswith('island_'):
continue
if isinstance(value, np.ndarray) and value.shape:
result_field = getattr(result_i, field.name)
if result_field.shape != value.shape:
# When ENABLE_ISLANDS is False, mujoco_warp allocates island fields with width 0,
# while the host MjData sizes to nv/ntree. Skip to prevent mismatch.
if value.size == 0 and not mjxw.mjwp_io.ENABLE_ISLANDS:
continue
raise ValueError(
f'Input field {field.name} has shape {value.shape}, but output'
f' has shape {result_field.shape}'
+9 -7
View File
@@ -550,7 +550,15 @@ class DataIOTest(parameterized.TestCase):
elif impl == 'warp':
qm = np.zeros((m.nv, m.nv), dtype=np.float64)
mujoco.mju_sym2dense(qm, d.M, m.M_rownnz, m.M_rowadr, m.M_colind)
np.testing.assert_allclose(dx._impl.M[:m.nv, :m.nv], qm)
warp_M = np.zeros((m.nv, m.nv))
mujoco.mju_sym2dense(
warp_M,
np.array(dx._impl.M),
m.M_rownnz,
m.M_rowadr,
m.M_colind,
)
np.testing.assert_allclose(warp_M, qm)
# TODO(taylorhowell): test efc__J
np.testing.assert_allclose(dx._impl.efc__aref[:3], d.efc_aref[:3])
@@ -811,12 +819,6 @@ class DataIOTest(parameterized.TestCase):
dx = mjx.make_data(m, impl='warp')
mjx.get_data_into(d, mx, dx)
# Island data is not populated when ENABLE_ISLANDS is False, so the host
# MjData island fields should be left at their default (nv,)/(ntree,)
# shapes rather than triggering a shape mismatch.
self.assertEqual(d.dof_island.shape, (m.nv,))
self.assertEqual(d.tree_island.shape, (m.ntree,))
@parameterized.parameters(('jax',))
def test_get_data_into_wrong_shape(self, impl):
"""Tests that get_data_into throwsif input and output shapes don't match."""
+11 -2
View File
@@ -327,7 +327,12 @@ def _jax_shim_fn(
f'Unknown param source: {mjwarp_field_info[arg].param_source}'
)
if arg in field_usage.data_out_fields:
# Only treat as output if it is a Warp array (excludes scalars like
# naconmax).
if (
arg in field_usage.data_out_fields
and 'array' in mjwarp_field_info[arg].expected_type
):
# all out fields are in_out, since JAX already allocated them
in_out_argnames.append(f"'{arg}'")
num_outputs += 1
@@ -344,7 +349,11 @@ def _jax_shim_fn(
if field_usage.render_context_in_caller:
jax_args.append('ctx.key')
needs_dummy_output = not field_usage.data_out_fields
# If there are no Warp array outputs, we need a dummy output for JAX FFI.
needs_dummy_output = not any(
'array' in mjwarp_field_info[f].expected_type
for f in field_usage.data_out_fields
)
if needs_dummy_output and fn_name != 'render':
num_outputs = 1
if field_usage.render_context_in_caller:
+31 -6
View File
@@ -181,11 +181,13 @@ def _get_annotations_recursive(
def _build_new_class_body_ast(
keys: Set[str],
keys: typing.Iterable[str],
cls_name: str,
target_annotations: Dict[str, Any],
defaults: Optional[Dict[str, Any]] = None,
shape_property: Optional[str] = None,
add_docstring: bool = True,
sort_keys: bool = True,
) -> List[ast.AST]:
"""Builds the list of AST nodes for the new class body."""
new_body_nodes: List[ast.AST] = []
@@ -194,16 +196,29 @@ def _build_new_class_body_ast(
docstring = f'Derived fields from {cls_name}.'
new_body_nodes.append(ast.Expr(value=ast.Constant(value=docstring)))
# Sort keys alphabetically before creating AST nodes
sorted_keys = sorted(list(keys))
for key in sorted_keys:
if sort_keys:
maybe_sorted_keys = sorted(list(keys))
else:
maybe_sorted_keys = list(keys)
for key in maybe_sorted_keys:
annotation_node = _get_target_annotation_node(key, target_annotations)
value_node = None
if defaults and key in defaults:
value_node = ast.Constant(value=defaults[key])
if value_node and value_node.value is None:
annotation_node = _ast_parse_type(
f'typing.Optional[{ast.unparse(annotation_node)}]'
)
new_body_nodes.append(
ast.AnnAssign(
target=ast.Name(id=key, ctx=ast.Store()),
annotation=annotation_node,
simple=1, # No value assignment
value=value_node,
simple=1,
)
)
@@ -316,11 +331,21 @@ _NESTED_DATACLASS_MANUAL_METHODS = {
def write_nested_dataclass(target_fpath: epath.Path, cls: Any):
fields = dataclasses.fields(cls)
keys = [f.name for f in fields]
defaults = {
f.name: f.default
for f in fields
if f.default is not dataclasses.MISSING
}
new_class_body = _build_new_class_body_ast(
set(cls.__annotations__.keys()),
keys,
cls.__name__,
dict(cls.__annotations__),
defaults=defaults,
add_docstring=False,
sort_keys=False,
)
cls_str = '\n'.join([' ' + ast.unparse(node) for node in new_class_body])
cls_str = cls_str.replace('jax.Array', 'np.ndarray')
+15 -3
View File
@@ -63,6 +63,10 @@ def _get_imported_module_fpaths(fpath: epath.Path) -> Dict[str, str]:
all_resolved_fpaths = {}
for fully_qualified_name, alias in all_imported_names:
# only trace mujoco and mujoco warp
if not (fully_qualified_name.startswith('mujoco_warp') or
fully_qualified_name.startswith('mujoco')):
continue
fpath = _resolve_module_name_to_fpath(fully_qualified_name)
if fpath:
all_resolved_fpaths[alias] = fpath
@@ -181,9 +185,17 @@ class _FunctionFieldUsageVisitor(ast.NodeVisitor):
self.add_field_usage(in_node, False)
return
for arg in node.args:
if isinstance(arg, ast.Attribute):
self.add_field_usage(arg, is_output=True)
# do not trace array creation
is_input_only_wp_fn = (
len(parts) == 2
and parts[0] == 'wp'
and parts[1] in ('empty', 'zeros', 'array', 'clone')
)
if not is_input_only_wp_fn:
for arg in node.args:
if isinstance(arg, ast.Attribute):
self.add_field_usage(arg, is_output=True)
called_fn_name = '.'.join(parts[1:])
key = (hash(self._current_fpath), called_fn_name)
+1
View File
@@ -61,6 +61,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.io import reset_data as reset_data
from mujoco.mjx.third_party.mujoco_warp._src.io import set_const as set_const
from mujoco.mjx.third_party.mujoco_warp._src.io import set_const_0 as set_const_0
from mujoco.mjx.third_party.mujoco_warp._src.io import set_const_fixed as set_const_fixed
from mujoco.mjx.third_party.mujoco_warp._src.io import set_const_spring as set_const_spring
from mujoco.mjx.third_party.mujoco_warp._src.io import set_length_range as set_length_range
from mujoco.mjx.third_party.mujoco_warp._src.island import island as island
from mujoco.mjx.third_party.mujoco_warp._src.passive import passive as passive
+16 -9
View File
@@ -25,7 +25,6 @@ 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 io
from mujoco.mjx.third_party.mujoco_warp._src import warp_util
from mujoco.mjx.third_party.mujoco_warp._src.io import load_trajectory
from mujoco.mjx.third_party.mujoco_warp._src.io import override_model
@@ -43,13 +42,13 @@ KEYFRAME = flags.DEFINE_integer("keyframe", 0, "keyframe to initialize simulatio
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)")
ENABLE_ISLANDS = flags.DEFINE_bool(
"enable_islands",
False,
"Enable constraint islands solver",
)
NVMAX = flags.DEFINE_integer("nvmax", None, "maximum active DOFs per world")
INIT_ASLEEP = flags.DEFINE_bool(
"init_asleep", False, "initialize all trees as asleep before simulation (requires sleep enabled)"
)
DEVICE = flags.DEFINE_string("device", None, "override the default Warp device")
REPLAY = flags.DEFINE_string("replay", None, "NPZ file with ctrl sequence to replay")
@@ -141,8 +140,6 @@ def init_structs(
fn: Callable[..., None], mjm: mujoco.MjModel
) -> Tuple[mjw.Model, mjw.Data, mjw.RenderContext | None, list[np.ndarray] | None]:
"""Initialize device structs."""
io.ENABLE_ISLANDS = ENABLE_ISLANDS.value
mjd = mujoco.MjData(mjm)
ctrls = None
if REPLAY.value:
@@ -158,8 +155,18 @@ def init_structs(
m = mjw.put_model(mjm)
if OVERRIDE.value:
override_model(m, OVERRIDE.value)
if INIT_ASLEEP.value:
mjd.tree_asleep[:] = np.arange(mjm.ntree, dtype=np.int32)
d = mjw.put_data(
mjm, mjd, nworld=NWORLD.value, nconmax=NCONMAX.value, njmax=NJMAX.value, njmax_nnz=NJMAX_NNZ.value, nccdmax=NCCDMAX.value
mjm,
mjd,
nworld=NWORLD.value,
nconmax=NCONMAX.value,
njmax=NJMAX.value,
njmax_nnz=NJMAX_NNZ.value,
nccdmax=NCCDMAX.value,
nvmax=NVMAX.value,
)
if mjw.RenderContext not in get_type_hints(fn).values():
+162 -129
View File
@@ -37,11 +37,11 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAX_EPAFACES
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAX_EPAHORIZON
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXCONPAIR
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL
from mujoco.mjx.third_party.mujoco_warp._src.types import NEW_GAP_SEMANTICS
from mujoco.mjx.third_party.mujoco_warp._src.types import Data
from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
from mujoco.mjx.third_party.mujoco_warp._src.types import Model
from mujoco.mjx.third_party.mujoco_warp._src.types import OverflowType
from mujoco.mjx.third_party.mujoco_warp._src.types import mat43
from mujoco.mjx.third_party.mujoco_warp._src.types import mat63
from mujoco.mjx.third_party.mujoco_warp._src.types import vec5
@@ -164,6 +164,7 @@ def ccd_hfield_kernel_builder(
gjk_iterations: int,
epa_iterations: int,
geomgeomid: int,
warn_overflow: bool,
):
"""Kernel builder for heightfield CCD collisions (no multiccd args)."""
@@ -243,6 +244,8 @@ def ccd_hfield_kernel_builder(
contact_type_out: wp.array[int],
contact_geomcollisionid_out: wp.array[int],
nacon_out: wp.array[int],
# Data out:
overflow_out: wp.array[int],
):
collisionid = wp.tid()
if collisionid >= ncollision_in[0]:
@@ -281,7 +284,9 @@ def ccd_hfield_kernel_builder(
ccdid = wp.atomic_add(nccd_in, wp.static(geomgeomid), 1)
if ccdid >= naccdmax_in:
wp.printf("CCD overflow - please increase naccdmax to %u\n", ccdid)
if wp.static(warn_overflow):
wp.printf("CCD overflow - please increase naccdmax to %u\n", ccdid)
wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.CCD))
return
_, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params(
@@ -416,10 +421,12 @@ def ccd_hfield_kernel_builder(
# add both triangles from this cell
for i in range(2):
if count >= MJ_MAXCONPAIR:
wp.printf(
"height field collision overflow, number of collisions >= %u - please adjust resolution: \n decrease the number of hfield rows/cols or modify size of colliding geom\n",
MJ_MAXCONPAIR,
)
if wp.static(warn_overflow):
wp.printf(
"height field collision overflow, number of collisions >= %u - please adjust resolution: \n decrease the number of hfield rows/cols or modify size of colliding geom\n",
MJ_MAXCONPAIR,
)
wp.atomic_or(overflow_out, worldid, OverflowType.HFIELD)
continue
# add vert
@@ -700,6 +707,10 @@ def ccd_hfield_kernel_builder(
return ccd_hfield_kernel
_CCD_OVERSUBSCRIBE_WAVES = 4
_CCD_MIN_BLOCKS = 2
@cache_kernel
def ccd_kernel_builder(
geomtype1: int,
@@ -708,6 +719,8 @@ def ccd_kernel_builder(
epa_iterations: int,
use_multiccd: bool,
geomgeomid: int,
block_dim: int,
warn_overflow: bool,
):
"""Kernel builder for non-heightfield CCD collisions (no hfield args)."""
@@ -767,6 +780,8 @@ def ccd_kernel_builder(
contact_type_out: wp.array[int],
contact_geomcollisionid_out: wp.array[int],
nacon_out: wp.array[int],
# Data out:
overflow_out: wp.array[int],
) -> int:
points = mat43()
witness1 = mat43()
@@ -777,10 +792,7 @@ def ccd_kernel_builder(
if is_collision_sensor:
cutoff = 1.0e32
else:
if wp.static(NEW_GAP_SEMANTICS):
cutoff = gap
else:
cutoff = 0.0
cutoff = gap
needs_epa, dist, ncollision, w1, w2, gjk_result, geom1, geom2 = gjk_phase(
opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]],
cutoff,
@@ -799,7 +811,9 @@ def ccd_kernel_builder(
if needs_epa:
ccdid = wp.atomic_add(nccd_in, geomgeomid, 1)
if ccdid >= naccdmax_in:
wp.printf("CCD overflow - please increase naccdmax to %u\n", ccdid)
if wp.static(warn_overflow):
wp.printf("CCD overflow - please increase naccdmax to %u\n", ccdid)
wp.atomic_or(overflow_out, worldid, OverflowType.CCD)
return 0
dist, ncollision, w1, w2, multiccd_idx = epa_phase(
opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]],
@@ -817,12 +831,8 @@ def ccd_kernel_builder(
epa_horizon_in[ccdid],
)
if wp.static(NEW_GAP_SEMANTICS):
if dist >= gap and pairid[1] == -1:
return 0
else:
if dist >= 0.0 and pairid[1] == -1:
return 0
if dist >= gap and pairid[1] == -1:
return 0
# CCD operates on margin-inflated shapes (support() inflates each geom by
# 0.5 * margin). The returned dist is therefore relative to the inflated
@@ -918,7 +928,7 @@ def ccd_kernel_builder(
return nactive
# runs convex collision on a set of geom pairs to recover contact info (non-heightfield)
@wp.kernel(module="unique", enable_backward=False)
@wp.kernel(module="unique", enable_backward=False, launch_bounds=(block_dim, _CCD_MIN_BLOCKS))
def ccd_kernel(
# Model:
opt_ccd_tolerance: wp.array[float],
@@ -961,6 +971,7 @@ def ccd_kernel_builder(
naccdmax_in: int,
ncollision_in: wp.array[int],
# In:
grid_stride_in: int,
collision_pair_in: wp.array[wp.vec2i],
collision_pairid_in: wp.array[wp.vec2i],
collision_worldid_in: wp.array[int],
@@ -998,122 +1009,131 @@ def ccd_kernel_builder(
contact_type_out: wp.array[int],
contact_geomcollisionid_out: wp.array[int],
nacon_out: wp.array[int],
# Data out:
overflow_out: wp.array[int],
):
collisionid = wp.tid()
if collisionid >= ncollision_in[0]:
return
tid = wp.tid()
for collisionid in range(tid, ncollision_in[0], grid_stride_in):
geoms = collision_pair_in[collisionid]
g1 = geoms[0]
g2 = geoms[1]
geoms = collision_pair_in[collisionid]
g1 = geoms[0]
g2 = geoms[1]
if geom_type[g1] != geomtype1 or geom_type[g2] != geomtype2:
continue
if geom_type[g1] != geomtype1 or geom_type[g2] != geomtype2:
return
worldid = collision_worldid_in[collisionid]
worldid = collision_worldid_in[collisionid]
_, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params(
geom_condim,
geom_priority,
geom_solmix,
geom_solref,
geom_solimp,
geom_friction,
geom_margin,
geom_gap,
pair_dim,
pair_solref,
pair_solreffriction,
pair_solimp,
pair_margin,
pair_gap,
pair_friction,
collision_pair_in,
collision_pairid_in,
collisionid,
worldid,
)
_, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params(
geom_condim,
geom_priority,
geom_solmix,
geom_solref,
geom_solimp,
geom_friction,
geom_margin,
geom_gap,
pair_dim,
pair_solref,
pair_solreffriction,
pair_solimp,
pair_margin,
pair_gap,
pair_friction,
collision_pair_in,
collision_pairid_in,
collisionid,
worldid,
)
geom1, geom2 = geom_collision_pair(
geom_type,
geom_dataid,
geom_size,
mesh_vertadr,
mesh_vertnum,
mesh_graphadr,
mesh_vert,
mesh_graph,
mesh_polynum,
mesh_polyadr,
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
mesh_polyvert,
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in,
geom_xmat_in,
geoms,
worldid,
)
geom1, geom2 = geom_collision_pair(
geom_type,
geom_dataid,
geom_size,
mesh_vertadr,
mesh_vertnum,
mesh_graphadr,
mesh_vert,
mesh_graph,
mesh_polynum,
mesh_polyadr,
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
mesh_polyvert,
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in,
geom_xmat_in,
geoms,
worldid,
)
eval_ccd_write_contact(
opt_ccd_tolerance,
naconmax_in,
naccdmax_in,
epa_vert_in,
epa_vert_index_in,
epa_face_in,
epa_pr_in,
epa_norm2_in,
epa_horizon_in,
multiccd_polygon_in,
multiccd_clipped_in,
multiccd_pnormal_in,
multiccd_pdist_in,
multiccd_idx1_in,
multiccd_idx2_in,
multiccd_n1_in,
multiccd_n2_in,
multiccd_endvert_in,
multiccd_face1_in,
multiccd_face2_in,
geom1,
geom2,
geoms,
worldid,
nccd_in,
margin,
gap,
condim,
friction,
solref,
solreffriction,
solimp,
geom1.pos,
geom2.pos,
collision_pairid_in[collisionid],
contact_dist_out,
contact_pos_out,
contact_frame_out,
contact_includemargin_out,
contact_friction_out,
contact_solref_out,
contact_solreffriction_out,
contact_solimp_out,
contact_dim_out,
contact_geom_out,
contact_efc_address_out,
contact_worldid_out,
contact_type_out,
contact_geomcollisionid_out,
nacon_out,
)
eval_ccd_write_contact(
opt_ccd_tolerance,
naconmax_in,
naccdmax_in,
epa_vert_in,
epa_vert_index_in,
epa_face_in,
epa_pr_in,
epa_norm2_in,
epa_horizon_in,
multiccd_polygon_in,
multiccd_clipped_in,
multiccd_pnormal_in,
multiccd_pdist_in,
multiccd_idx1_in,
multiccd_idx2_in,
multiccd_n1_in,
multiccd_n2_in,
multiccd_endvert_in,
multiccd_face1_in,
multiccd_face2_in,
geom1,
geom2,
geoms,
worldid,
nccd_in,
margin,
gap,
condim,
friction,
solref,
solreffriction,
solimp,
geom1.pos,
geom2.pos,
collision_pairid_in[collisionid],
contact_dist_out,
contact_pos_out,
contact_frame_out,
contact_includemargin_out,
contact_friction_out,
contact_solref_out,
contact_solreffriction_out,
contact_solimp_out,
contact_dim_out,
contact_geom_out,
contact_efc_address_out,
contact_worldid_out,
contact_type_out,
contact_geomcollisionid_out,
nacon_out,
overflow_out,
)
return ccd_kernel
def _ccd_grid_size(kernel, naconmax: int) -> int:
# Grid-stride launch width for the CCD kernel: a few device waves, capped at the contact
# capacity. The kernel strides over the actual candidate count, so we avoid launching one
# (mostly idle) thread per naconmax slot.
block_size, min_grid_size = wp.get_suggested_block_size(kernel)
return max(1, min(naconmax, _CCD_OVERSUBSCRIBE_WAVES * block_size * min_grid_size))
@event_scope
def convex_narrowphase(m: Model, d: Data, ctx: CollisionContext, collision_table: list[tuple[GeomType, GeomType]]):
"""Runs narrowphase collision detection for convex geom pairs.
@@ -1205,7 +1225,7 @@ def convex_narrowphase(m: Model, d: Data, ctx: CollisionContext, collision_table
count, geomgeomid = _pair_count(g1, g2)
if (g1 == GeomType.HFIELD or g2 == GeomType.HFIELD) and count:
wp.launch(
ccd_hfield_kernel_builder(g1, g2, m.opt.ccd_iterations, epa_iterations, geomgeomid),
ccd_hfield_kernel_builder(g1, g2, m.opt.ccd_iterations, epa_iterations, geomgeomid, bool(m.opt.warn_overflow)),
dim=d.naconmax,
inputs=[
m.opt.ccd_tolerance,
@@ -1263,7 +1283,7 @@ def convex_narrowphase(m: Model, d: Data, ctx: CollisionContext, collision_table
epa_horizon,
nccd,
],
outputs=contact_outputs,
outputs=contact_outputs + [d.overflow],
)
# Allocate multiccd arrays only for non-heightfield collisions
@@ -1296,9 +1316,21 @@ def convex_narrowphase(m: Model, d: Data, ctx: CollisionContext, collision_table
g2 = geom_pair[1].value
count, geomgeomid = _pair_count(g1, g2)
if g1 != GeomType.HFIELD and g2 != GeomType.HFIELD and count:
ccd_k = ccd_kernel_builder(
g1,
g2,
m.opt.ccd_iterations,
epa_iterations,
use_multiccd,
geomgeomid,
m.block_dim.convex_ccd,
bool(m.opt.warn_overflow),
)
ccd_grid = _ccd_grid_size(ccd_k, d.naconmax)
wp.launch(
ccd_kernel_builder(g1, g2, m.opt.ccd_iterations, epa_iterations, use_multiccd, geomgeomid),
dim=d.naconmax,
ccd_k,
dim=ccd_grid,
block_dim=m.block_dim.convex_ccd,
inputs=[
m.opt.ccd_tolerance,
m.geom_type,
@@ -1338,6 +1370,7 @@ def convex_narrowphase(m: Model, d: Data, ctx: CollisionContext, collision_table
d.naconmax,
d.naccdmax,
d.ncollision,
ccd_grid,
ctx.collision_pair,
ctx.collision_pairid,
ctx.collision_worldid,
@@ -1360,5 +1393,5 @@ def convex_narrowphase(m: Model, d: Data, ctx: CollisionContext, collision_table
multiccd_face2,
nccd,
],
outputs=contact_outputs,
outputs=contact_outputs + [d.overflow],
)
@@ -23,7 +23,6 @@ import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src.math import safe_div
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 NEW_GAP_SEMANTICS
from mujoco.mjx.third_party.mujoco_warp._src.types import ContactType
from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
from mujoco.mjx.third_party.mujoco_warp._src.types import mat63
@@ -219,10 +218,7 @@ def write_contact(
contact_frame_out[cid] = frame_in
contact_geom_out[cid] = geoms_in
contact_worldid_out[cid] = worldid_in
if wp.static(NEW_GAP_SEMANTICS):
includemargin = margin_in
else:
includemargin = margin_in - gap_in
includemargin = margin_in
contact_includemargin_out[cid] = includemargin
contact_dim_out[cid] = condim_in
contact_friction_out[cid] = friction_in
+200 -117
View File
@@ -20,7 +20,7 @@ import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src.collision_convex import convex_narrowphase
from mujoco.mjx.third_party.mujoco_warp._src.collision_core import CollisionContext
from mujoco.mjx.third_party.mujoco_warp._src.collision_core import create_collision_context
from mujoco.mjx.third_party.mujoco_warp._src.collision_flex import flex_narrowphase
from mujoco.mjx.third_party.mujoco_warp._src.collision_flex import flex_collision
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import primitive_narrowphase
from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import sdf_narrowphase
from mujoco.mjx.third_party.mujoco_warp._src.math import upper_tri_index
@@ -79,24 +79,18 @@ MJ_COLLISION_TABLE = {
}
@cache_kernel
def _zero_nacon_ncollision(enable_sleep: bool = False):
@wp.kernel(module="unique", enable_backward=False)
def zero_nacon_ncollision(
# In:
skip_in: wp.array[int],
# Data out:
nacon_out: wp.array[int],
ncollision_out: wp.array[int],
):
ncollision_out[0] = 0
if wp.static(enable_sleep):
if skip_in[0] != 0:
nacon_out[0] = 0
else:
nacon_out[0] = 0
return zero_nacon_ncollision
# TODO(team): Implement narrowphase flex collision support for:
# - HFIELD
# - ELLIPSOID
# - SDF
MJ_FLEX_COLLISION_TABLE = {
(GeomType.PLANE, GeomType.FLEX): CollisionType.PRIMITIVE,
(GeomType.SPHERE, GeomType.FLEX): CollisionType.PRIMITIVE,
(GeomType.CAPSULE, GeomType.FLEX): CollisionType.PRIMITIVE,
(GeomType.BOX, GeomType.FLEX): CollisionType.PRIMITIVE,
(GeomType.CYLINDER, GeomType.FLEX): CollisionType.PRIMITIVE,
(GeomType.MESH, GeomType.FLEX): CollisionType.CONVEX,
}
@wp.func
@@ -388,7 +382,7 @@ def _binary_search(values: wp.array[Any], value: Any, lower: int, upper: int) ->
@cache_kernel
def _sap_project(opt_broadphase: int, enable_sleep: bool = False):
def _sap_project(opt_broadphase: int):
@wp.kernel(module="unique", enable_backward=False)
def sap_project(
# Model:
@@ -401,7 +395,6 @@ def _sap_project(opt_broadphase: int, enable_sleep: bool = False):
nworld_in: int,
# In:
direction_in: wp.vec3,
skip_in: wp.array[int],
# Out:
projection_lower_out: wp.array2d[float],
projection_upper_out: wp.array2d[float],
@@ -410,10 +403,6 @@ def _sap_project(opt_broadphase: int, enable_sleep: bool = False):
):
worldid, geomid = wp.tid()
if wp.static(enable_sleep):
if skip_in[0] == 0:
return
xpos = geom_xpos_in[worldid, geomid]
rbound = geom_rbound[worldid % geom_rbound.shape[0], geomid]
@@ -441,44 +430,40 @@ def _sap_project(opt_broadphase: int, enable_sleep: bool = False):
return sap_project
@cache_kernel
def _sap_range(enable_sleep: bool = False):
@wp.kernel(module="unique", enable_backward=False)
def sap_range(
# Model:
ngeom: int,
# In:
projection_lower_in: wp.array2d[float],
projection_upper_in: wp.array2d[float],
sort_index_in: wp.array2d[int],
skip_in: wp.array[int],
# Out:
range_out: wp.array2d[int],
):
worldid, geomid = wp.tid()
@wp.kernel
def _sap_range(
# Model:
ngeom: int,
# In:
projection_lower_in: wp.array2d[float],
projection_upper_in: wp.array2d[float],
sort_index_in: wp.array2d[int],
# Out:
range_out: wp.array2d[int],
):
worldid, geomid = wp.tid()
if wp.static(enable_sleep):
if skip_in[0] == 0:
range_out[worldid, geomid] = 0
return
# current bounding geom
idx = sort_index_in[worldid, geomid]
# current bounding geom
idx = sort_index_in[worldid, geomid]
upper = projection_upper_in[worldid, idx]
upper = projection_upper_in[worldid, idx]
limit = _binary_search(projection_lower_in[worldid], upper, geomid + 1, ngeom)
limit = wp.min(ngeom - 1, limit)
limit = _binary_search(projection_lower_in[worldid], upper, geomid + 1, ngeom)
limit = wp.min(ngeom - 1, limit)
# range of geoms for the sweep and prune process
range_out[worldid, geomid] = limit - geomid
return sap_range
# range of geoms for the sweep and prune process
range_out[worldid, geomid] = limit - geomid
@cache_kernel
def _sap_broadphase(
opt_broadphase_filter: int, ngeom_aabb: int, ngeom_rbound: int, ngeom_margin: int, ngeom_gap: int, enable_sleep: bool = False
opt_broadphase_filter: int,
ngeom_aabb: int,
ngeom_rbound: int,
ngeom_margin: int,
ngeom_gap: int,
enable_sleep: bool = False,
incremental: bool = False,
):
@wp.kernel(module="unique", enable_backward=False)
def kernel(
@@ -501,7 +486,7 @@ def _sap_broadphase(
sort_index_in: wp.array2d[int],
cumulative_sum_in: wp.array[int],
nsweep_in: int,
skip_in: wp.array[int],
body_awake_prev_in: wp.array2d[int],
# Data out:
ncollision_out: wp.array[int],
# Out:
@@ -511,10 +496,6 @@ def _sap_broadphase(
):
worldgeomid = wp.tid()
if wp.static(enable_sleep):
if skip_in[0] == 0:
return
nworldgeom = nworld_in * ngeom
nworkpackages = cumulative_sum_in[nworldgeom - 1]
@@ -555,6 +536,18 @@ def _sap_broadphase(
if (s1 == SleepState.ASLEEP and s2 == SleepState.STATIC) or (s2 == SleepState.ASLEEP and s1 == SleepState.STATIC):
continue
if wp.static(incremental):
# only re-emit pairs that pass 1 skipped (see _nxn_broadphase for rationale)
p1 = body_awake_prev_in[worldid, b1]
p2 = body_awake_prev_in[worldid, b2]
skipped_pass1 = (
(p1 == SleepState.ASLEEP and p2 == SleepState.ASLEEP)
or (p1 == SleepState.ASLEEP and p2 == SleepState.STATIC)
or (p2 == SleepState.ASLEEP and p1 == SleepState.STATIC)
)
if not skipped_pass1:
continue
if (
wp.static(_broadphase_filter(opt_broadphase_filter, ngeom_aabb, ngeom_rbound, ngeom_margin, ngeom_gap))(
geom_aabb, geom_rbound, geom_margin, geom_gap, geom_xpos_in, geom_xmat_in, geom1, geom2, worldid
@@ -606,7 +599,12 @@ def _segmented_sort(tile_size: int):
@event_scope
def sap_broadphase(m: Model, d: Data, ctx: CollisionContext, skip: Optional[wp.array] = None):
def sap_broadphase(
m: Model,
d: Data,
ctx: CollisionContext,
awake_prev: Optional[wp.array] = None,
):
"""Runs broadphase collision detection using a sweep-and-prune (SAP) algorithm.
This method is more efficient than the N-squared approach for large numbers of
@@ -622,9 +620,15 @@ def sap_broadphase(m: Model, d: Data, ctx: CollisionContext, skip: Optional[wp.a
- `SAP_TILE`: Uses a tile-based sort.
- `SAP_SEGMENTED`: Uses a segmented sort.
Unlike `nxn_broadphase`, SAP cannot be wrapped in a CUDA graph conditional to skip the
incremental sleeping pass: its sort/scan utilities allocate scratch internally, which is
not allowed inside a conditional body. The incremental filter in the sweep kernel still
restricts the emitted pairs to those involving newly-awakened bodies.
"""
nworldgeom = d.nworld * m.ngeom
skip_in = skip if skip is not None else wp.ones(1, dtype=int)
incremental = awake_prev is not None
awake_prev_in = awake_prev if awake_prev is not None else d.body_awake
enable_sleep = bool(m.opt.enableflags & EnableBit.SLEEP)
# TODO(team): direction
@@ -641,9 +645,9 @@ def sap_broadphase(m: Model, d: Data, ctx: CollisionContext, skip: Optional[wp.a
segmented_index = wp.empty(d.nworld + 1 if m.opt.broadphase == BroadphaseType.SAP_SEGMENTED else 0, dtype=int)
wp.launch(
kernel=_sap_project(m.opt.broadphase, enable_sleep),
kernel=_sap_project(m.opt.broadphase),
dim=(d.nworld, m.ngeom),
inputs=[m.ngeom, m.geom_rbound, m.geom_margin, m.geom_gap, d.geom_xpos, d.nworld, direction, skip_in],
inputs=[m.ngeom, m.geom_rbound, m.geom_margin, m.geom_gap, d.geom_xpos, d.nworld, direction],
outputs=[
projection_lower.reshape((-1, m.ngeom)),
projection_upper,
@@ -666,9 +670,9 @@ def sap_broadphase(m: Model, d: Data, ctx: CollisionContext, skip: Optional[wp.a
)
wp.launch(
kernel=_sap_range(enable_sleep),
kernel=_sap_range,
dim=(d.nworld, m.ngeom),
inputs=[m.ngeom, projection_lower.reshape((-1, m.ngeom)), projection_upper, sort_index.reshape((-1, m.ngeom)), skip_in],
inputs=[m.ngeom, projection_lower.reshape((-1, m.ngeom)), projection_upper, sort_index.reshape((-1, m.ngeom))],
outputs=[range_],
)
@@ -686,6 +690,7 @@ def sap_broadphase(m: Model, d: Data, ctx: CollisionContext, skip: Optional[wp.a
m.geom_margin.shape[0],
m.geom_gap.shape[0],
enable_sleep,
incremental,
),
dim=nsweep,
inputs=[
@@ -705,7 +710,7 @@ def sap_broadphase(m: Model, d: Data, ctx: CollisionContext, skip: Optional[wp.a
sort_index.reshape((-1, m.ngeom)),
cumulative_sum.reshape(-1),
nsweep,
skip_in,
awake_prev_in,
],
outputs=[d.ncollision, ctx.collision_pair, ctx.collision_pairid, ctx.collision_worldid],
)
@@ -713,7 +718,13 @@ def sap_broadphase(m: Model, d: Data, ctx: CollisionContext, skip: Optional[wp.a
@cache_kernel
def _nxn_broadphase(
opt_broadphase_filter: int, ngeom_aabb: int, ngeom_rbound: int, ngeom_margin: int, ngeom_gap: int, enable_sleep: bool = False
opt_broadphase_filter: int,
ngeom_aabb: int,
ngeom_rbound: int,
ngeom_margin: int,
ngeom_gap: int,
enable_sleep: bool = False,
incremental: bool = False,
):
@wp.kernel(module="unique", enable_backward=False)
def kernel(
@@ -732,7 +743,7 @@ def _nxn_broadphase(
body_awake_in: wp.array2d[int],
naconmax_in: int,
# In:
skip_in: wp.array[int],
body_awake_prev_in: wp.array2d[int],
# Data out:
ncollision_out: wp.array[int],
# Out:
@@ -742,10 +753,6 @@ def _nxn_broadphase(
):
worldid, elementid = wp.tid()
if wp.static(enable_sleep):
if skip_in[0] == 0:
return
geom = nxn_geom_pair[elementid]
geom1 = geom[0]
geom2 = geom[1]
@@ -760,6 +767,21 @@ def _nxn_broadphase(
if (s1 == SleepState.ASLEEP and s2 == SleepState.STATIC) or (s2 == SleepState.ASLEEP and s1 == SleepState.STATIC):
return
if wp.static(incremental):
# On the incremental pass we only want pairs that were *not* already emitted in pass 1.
# Pass 1 skipped a pair iff both bodies were asleep, or one asleep and one static. Such a
# pair becomes relevant now only because one of its bodies was newly awakened, so re-emit
# it. Every other pair was handled in pass 1 and its contact geometry is unchanged.
p1 = body_awake_prev_in[worldid, b1]
p2 = body_awake_prev_in[worldid, b2]
skipped_pass1 = (
(p1 == SleepState.ASLEEP and p2 == SleepState.ASLEEP)
or (p1 == SleepState.ASLEEP and p2 == SleepState.STATIC)
or (p2 == SleepState.ASLEEP and p1 == SleepState.STATIC)
)
if not skipped_pass1:
return
if (
wp.static(_broadphase_filter(opt_broadphase_filter, ngeom_aabb, ngeom_rbound, ngeom_margin, ngeom_gap))(
geom_aabb, geom_rbound, geom_margin, geom_gap, geom_xpos_in, geom_xmat_in, geom1, geom2, worldid
@@ -783,8 +805,28 @@ def _nxn_broadphase(
return kernel
@wp.kernel
def _any_awake_changed(
# Data in:
body_awake_in: wp.array2d[int],
# In:
awake_prev_in: wp.array2d[int],
# Out:
changed_out: wp.array[int],
):
worldid, bodyid = wp.tid()
# benign race: every thread that fires writes the same value
if awake_prev_in[worldid, bodyid] != body_awake_in[worldid, bodyid]:
changed_out[0] = 1
@event_scope
def nxn_broadphase(m: Model, d: Data, ctx: CollisionContext, skip: Optional[wp.array] = None):
def nxn_broadphase(
m: Model,
d: Data,
ctx: CollisionContext,
awake_prev: Optional[wp.array] = None,
):
"""Runs broadphase collision detection using a brute-force N-squared approach.
This function iterates through a pre-filtered list of all possible geometry pairs and
@@ -796,46 +838,70 @@ def nxn_broadphase(m: Model, d: Data, ctx: CollisionContext, skip: Optional[wp.a
The initial list of pairs is filtered at model creation time to exclude pairs based on
`contype`/`conaffinity`, parent-child relationships, and explicit `<exclude>` tags.
Passing ``awake_prev`` runs the incremental sleeping pass: only pairs involving a newly-awakened
body are emitted. When graph conditionals are available the launch is wrapped in one gated on
whether any body woke since pass 1 (``awake_prev != body_awake``), so the broadphase is skipped
wholesale on steps where nothing woke; otherwise it runs unconditionally and the per-pair filter
restricts the emitted pairs.
"""
enable_sleep = bool(m.opt.enableflags & EnableBit.SLEEP)
skip_in = skip if skip is not None else wp.ones(1, dtype=int)
wp.launch(
_nxn_broadphase(
m.opt.broadphase_filter,
m.geom_aabb.shape[0],
m.geom_rbound.shape[0],
m.geom_margin.shape[0],
m.geom_gap.shape[0],
enable_sleep,
),
dim=(d.nworld, m.nxn_geom_pair_filtered.shape[0]),
inputs=[
m.geom_type,
m.geom_bodyid,
m.geom_aabb,
m.geom_rbound,
m.geom_margin,
m.geom_gap,
m.nxn_geom_pair_filtered,
m.nxn_pairid_filtered,
d.geom_xpos,
d.geom_xmat,
d.body_awake,
d.naconmax,
skip_in,
],
outputs=[
d.ncollision,
ctx.collision_pair,
ctx.collision_pairid,
ctx.collision_worldid,
],
)
incremental = awake_prev is not None
awake_prev_in = awake_prev if awake_prev is not None else d.body_awake
# On the incremental pass, skip the broadphase wholesale via a graph conditional when nothing woke
# since pass 1. A wake is exactly a body whose state changed between awake_prev and body_awake
# (nothing sleeps between the passes), so the condition is derived here rather than threaded in.
cond = None
if incremental and m.opt.graph_conditional:
cond = wp.zeros(1, dtype=int)
wp.launch(_any_awake_changed, dim=(d.nworld, m.nbody), inputs=[d.body_awake, awake_prev], outputs=[cond])
def _launch():
wp.launch(
_nxn_broadphase(
m.opt.broadphase_filter,
m.geom_aabb.shape[0],
m.geom_rbound.shape[0],
m.geom_margin.shape[0],
m.geom_gap.shape[0],
enable_sleep,
incremental,
),
dim=(d.nworld, m.nxn_geom_pair_filtered.shape[0]),
inputs=[
m.geom_type,
m.geom_bodyid,
m.geom_aabb,
m.geom_rbound,
m.geom_margin,
m.geom_gap,
m.nxn_geom_pair_filtered,
m.nxn_pairid_filtered,
d.geom_xpos,
d.geom_xmat,
d.body_awake,
d.naconmax,
awake_prev_in,
],
outputs=[
d.ncollision,
ctx.collision_pair,
ctx.collision_pairid,
ctx.collision_worldid,
],
)
if cond is not None:
wp.capture_if(cond, on_true=_launch)
else:
_launch()
def _narrowphase(m: Model, d: Data, ctx: CollisionContext):
collision_table = MJ_COLLISION_TABLE
if m.opt.disableflags & DisableBit.NATIVECCD:
collision_table = collision_table.copy()
collision_table[(GeomType.BOX, GeomType.BOX)] = CollisionType.PRIMITIVE
convex_pairs = [key for key, value in collision_table.items() if value == CollisionType.CONVEX]
@@ -849,12 +915,13 @@ def _narrowphase(m: Model, d: Data, ctx: CollisionContext):
if m.has_sdf_geom:
sdf_narrowphase(m, d, ctx)
if m.nflex > 0:
flex_narrowphase(m, d)
@event_scope
def collision(m: Model, d: Data, skip: Optional[wp.array] = None):
def collision(
m: Model,
d: Data,
awake_prev: Optional[wp.array] = None,
):
"""Runs the full collision detection pipeline.
This function orchestrates the broadphase and narrowphase collision detection stages. It
@@ -870,6 +937,10 @@ def collision(m: Model, d: Data, skip: Optional[wp.array] = None):
This function will do nothing except zero out arrays if collision detection is disabled
via `m.opt.disableflags` or if `d.nacon` is 0.
Passing `awake_prev` (the awake state snapshotted before the post-collision wake) runs the
incremental sleeping pass: contacts are appended to the existing buffer and only pairs involving
a newly-awakened body are emitted.
"""
if d.naconmax == 0 or m.opt.disableflags & (DisableBit.CONSTRAINT | DisableBit.CONTACT):
d.nacon.zero_()
@@ -877,18 +948,30 @@ def collision(m: Model, d: Data, skip: Optional[wp.array] = None):
# TODO(team): create context outside collision?
ctx = create_collision_context(d.naconmax)
skip_in = skip if skip is not None else wp.ones(1, dtype=int)
enable_sleep = bool(m.opt.enableflags & EnableBit.SLEEP)
# zero counters
wp.launch(_zero_nacon_ncollision(enable_sleep), dim=1, inputs=[skip_in], outputs=[d.nacon, d.ncollision])
incremental = awake_prev is not None
# The incremental sleeping pass appends newly-awakened contacts to the pass-1 buffer, so nacon is
# preserved and only ncollision (the broadphase pair counter) is reset. nxn_broadphase also skips
# its launch wholesale via a graph conditional when nothing woke; pre-zeroing ncollision here
# keeps the narrowphase below a no-op in that case. SAP cannot use a conditional (its sort/scan
# allocate internally), so it relies on the per-pair incremental filter instead.
d.ncollision.zero_()
if not incremental:
d.nacon.zero_()
if m.opt.broadphase == BroadphaseType.NXN:
nxn_broadphase(m, d, ctx, skip)
nxn_broadphase(m, d, ctx, awake_prev)
else:
sap_broadphase(m, d, ctx, skip)
sap_broadphase(m, d, ctx, awake_prev)
_narrowphase(m, d, ctx)
# Flex collision is not sleeping-aware: pass 1 emits every flex contact regardless of awake state,
# so the incremental pass has nothing to add (and re-running it would duplicate those contacts).
# It therefore only runs on the full pass.
if m.nflex > 0 and not incremental:
flex_collision(m, d, ctx)
if m.callback.contactfilter:
m.callback.contactfilter(m, d)
File diff suppressed because it is too large Load Diff
+63 -18
View File
@@ -28,8 +28,19 @@ wp.set_module_options({"enable_backward": False})
FLOAT_MIN = -1e30
FLOAT_MAX = 1e30
MINVAL = 1e-15
MIN_DIST = 1e-10
MINVAL2 = 1e-30
MAXVAL = 1e15
MAXVAL2 = 1e30
# polytope minimal interior distance to origin
MIN_DIST2 = 1e-10
MIN_DIST3 = 1e-10
MIN_DIST4 = 1e-17
# minimal tolerance for EPA
MIN_EPATOL = 1e-7
FACE_TOL = wp.static(math.cos(0.0016))
EDGE_TOL = wp.static(math.sin(0.0016))
@@ -185,6 +196,19 @@ def support(geom: Geom, geomtype: int, dir: wp.vec3) -> SupportPoint:
if dist > max_dist:
max_dist = dist
sp.point = vert
elif geomtype == GeomType.TRIANGLE:
t1 = geom.rot[0, :]
t2 = geom.rot[1, :]
t3 = geom.rot[2, :]
d1 = wp.dot(t1, dir)
d2 = wp.dot(t2, dir)
d3 = wp.dot(t3, dir)
if d1 > d2 and d1 > d3:
sp.point = t1
elif d2 > d3:
sp.point = t2
else:
sp.point = t3
if geom.margin > 0.0:
sp.point += dir * (0.5 * geom.margin)
@@ -984,27 +1008,27 @@ def _polytope2(
_epa_support(pt, 4, geom1, geom2, geomtype1, geomtype2, d3 / wp.norm_l2(d3))
# build hexahedron
if _attach_face(pt, 0, 0, 2, 3) < MIN_DIST:
if _attach_face(pt, 0, 0, 2, 3) < MIN_DIST2:
pt.status = -1
return pt, _replace_simplex3(pt, 0, 2, 3)
if _attach_face(pt, 1, 0, 4, 2) < MIN_DIST:
if _attach_face(pt, 1, 0, 4, 2) < MIN_DIST2:
pt.status = -1
return pt, _replace_simplex3(pt, 0, 4, 2)
if _attach_face(pt, 2, 0, 3, 4) < MIN_DIST:
if _attach_face(pt, 2, 0, 3, 4) < MIN_DIST2:
pt.status = -1
return pt, _replace_simplex3(pt, 0, 3, 4)
if _attach_face(pt, 3, 1, 3, 2) < MIN_DIST:
if _attach_face(pt, 3, 1, 3, 2) < MIN_DIST2:
pt.status = -1
return pt, _replace_simplex3(pt, 1, 3, 2)
if _attach_face(pt, 4, 1, 2, 4) < MIN_DIST:
if _attach_face(pt, 4, 1, 2, 4) < MIN_DIST2:
pt.status = -1
return pt, _replace_simplex3(pt, 1, 2, 4)
if _attach_face(pt, 5, 1, 4, 3) < MIN_DIST:
if _attach_face(pt, 5, 1, 4, 3) < MIN_DIST2:
pt.status = -1
return pt, _replace_simplex3(pt, 1, 4, 3)
@@ -1085,22 +1109,22 @@ def _polytope3(
return pt
# create hexahedron for EPA
if _attach_face(pt, 0, 4, 0, 1) < MIN_DIST:
if _attach_face(pt, 0, 4, 0, 1) < MIN_DIST3:
pt.status = 6
return pt
if _attach_face(pt, 1, 4, 2, 0) < MIN_DIST:
if _attach_face(pt, 1, 4, 2, 0) < MIN_DIST3:
pt.status = 7
return pt
if _attach_face(pt, 2, 4, 1, 2) < MIN_DIST:
if _attach_face(pt, 2, 4, 1, 2) < MIN_DIST3:
pt.status = 8
return pt
if _attach_face(pt, 3, 3, 1, 0) < MIN_DIST:
if _attach_face(pt, 3, 3, 1, 0) < MIN_DIST3:
pt.status = 9
return pt
if _attach_face(pt, 4, 3, 0, 2) < MIN_DIST:
if _attach_face(pt, 4, 3, 0, 2) < MIN_DIST3:
pt.status = 10
return pt
if _attach_face(pt, 5, 3, 2, 1) < MIN_DIST:
if _attach_face(pt, 5, 3, 2, 1) < MIN_DIST3:
pt.status = 11
return pt
@@ -1141,19 +1165,19 @@ def _polytope4(
pt.vert_index[7] = simplex_index2[3]
# if the origin is on a face, replace the 3-simplex with a 2-simplex
if _attach_face(pt, 0, 0, 1, 2) < MIN_DIST:
if _attach_face(pt, 0, 0, 1, 2) < MIN_DIST4:
pt.status = -1
return pt, _replace_simplex3(pt, 0, 1, 2)
if _attach_face(pt, 1, 0, 3, 1) < MIN_DIST:
if _attach_face(pt, 1, 0, 3, 1) < MIN_DIST4:
pt.status = -1
return pt, _replace_simplex3(pt, 0, 3, 1)
if _attach_face(pt, 2, 0, 2, 3) < MIN_DIST:
if _attach_face(pt, 2, 0, 2, 3) < MIN_DIST4:
pt.status = -1
return pt, _replace_simplex3(pt, 0, 2, 3)
if _attach_face(pt, 3, 3, 2, 1) < MIN_DIST:
if _attach_face(pt, 3, 3, 2, 1) < MIN_DIST4:
pt.status = -1
return pt, _replace_simplex3(pt, 3, 2, 1)
@@ -1215,7 +1239,7 @@ def _epa(
upper2 = FLOAT_MAX
idx = int(-1)
pidx = int(-1)
epsilon = wp.where(is_discrete, 1e-15, tolerance)
epsilon = wp.where(is_discrete, MIN_EPATOL, tolerance)
nvalid = pt.nface # number of potential faces for expanding the polytope
# the face vertices are encoded in 10-bits that index the vertex array,
@@ -1895,6 +1919,27 @@ def _polygon_clip(
if npolygon < 1:
return 0, witness1, witness2
# if the face is an edge, remove potential duplicates
if nface2 == 2 and npolygon > 2:
best1 = int(0)
best2 = int(1)
max_d = float(0.0)
for i in range(npolygon):
polygon_out_i = polygon_out[i]
for j in range(i + 1, npolygon):
diff = polygon_out[j] - polygon_out_i
d2 = wp.dot(diff, diff)
if d2 > max_d:
max_d = d2
best1 = i
best2 = j
witness2[0] = polygon_out[best1]
witness1[0] = witness2[0] - dir
witness2[1] = polygon_out[best2]
witness1[1] = witness2[1] - dir
return 2, witness1, witness2
if npolygon > 4:
quad = _polygon_quad(polygon_out, npolygon)
for i in range(4):
@@ -1988,6 +1988,33 @@ def cylinder_triangle(
nrm2 = nrm
cnt += 1
# edge-vs-segment checks
for edge_idx in range(3):
if cnt >= 2:
break
u = t1 if edge_idx == 0 else (t2 if edge_idx == 1 else t3)
v = t2 if edge_idx == 0 else (t3 if edge_idx == 1 else t1)
closest_axis, closest_edge = closest_segment_to_segment_points(p1, p2, u, v)
diff = closest_edge - closest_axis
dist_raw = wp.length(diff)
if dist_raw < cylinder_radius + tri_radius:
if dist_raw > MJ_MINVAL:
nrm = diff / dist_raw
d = dist_raw - cylinder_radius - tri_radius
p = (closest_axis + closest_edge + nrm * (cylinder_radius - tri_radius)) * 0.5
if cnt == 0:
dist1 = d
pos1 = p
nrm1 = nrm
else:
dist2 = d
pos2 = p
nrm2 = nrm
cnt += 1
return (
wp.vec2(dist1, dist2),
mat23f(pos1[0], pos1[1], pos1[2], pos2[0], pos2[1], pos2[2]),
File diff suppressed because it is too large Load Diff
+629 -112
View File
@@ -17,6 +17,8 @@ import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src import math
from mujoco.mjx.third_party.mujoco_warp._src import util_misc
from mujoco.mjx.third_party.mujoco_warp._src.passive import ellipsoid_max_moment
from mujoco.mjx.third_party.mujoco_warp._src.passive import geom_semiaxes
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
@@ -24,6 +26,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import Data
from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
from mujoco.mjx.third_party.mujoco_warp._src.types import DynType
from mujoco.mjx.third_party.mujoco_warp._src.types import GainType
from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType
from mujoco.mjx.third_party.mujoco_warp._src.types import Model
from mujoco.mjx.third_party.mujoco_warp._src.types import vec10
from mujoco.mjx.third_party.mujoco_warp._src.types import vec10f
@@ -172,58 +175,6 @@ def _nonzero_mask(x: float) -> float:
return 0.0
@wp.kernel
def _qderiv_actuator_passive_actuation_dense(
# Model:
nu: int,
# Data in:
moment_rownnz_in: wp.array2d[int],
moment_rowadr_in: wp.array2d[int],
moment_colind_in: wp.array2d[int],
actuator_moment_in: wp.array2d[float],
# In:
vel_in: wp.array2d[float],
Mi: wp.array[int],
Mj: wp.array[int],
# Out:
qDeriv_out: wp.array3d[float],
):
worldid, elemid = wp.tid()
dofiid = Mi[elemid]
dofjid = Mj[elemid]
qderiv_contrib = float(0.0)
for actid in range(nu):
vel = vel_in[worldid, actid]
if vel == 0.0:
continue
# TODO(team): restructure sparse version for better parallelism?
moment_i = float(0.0)
moment_j = float(0.0)
rownnz = moment_rownnz_in[worldid, actid]
rowadr = moment_rowadr_in[worldid, actid]
for i in range(rownnz):
sparseid = rowadr + i
colind = moment_colind_in[worldid, sparseid]
if colind == dofiid:
moment_i = actuator_moment_in[worldid, sparseid]
if colind == dofjid:
moment_j = actuator_moment_in[worldid, sparseid]
if moment_i != 0.0 and moment_j != 0.0:
break
if moment_i == 0 and moment_j == 0:
continue
qderiv_contrib += moment_i * moment_j * vel
qDeriv_out[worldid, dofiid, dofjid] = qderiv_contrib
if dofiid != dofjid:
qDeriv_out[worldid, dofjid, dofiid] = qderiv_contrib
@wp.kernel
def _qderiv_actuator_passive_actuation_sparse(
# Model:
@@ -236,7 +187,7 @@ def _qderiv_actuator_passive_actuation_sparse(
# In:
vel_in: wp.array2d[float],
# Out:
qDeriv_out: wp.array3d[float],
qDeriv_out: wp.array2d[float],
):
worldid, actid = wp.tid()
@@ -264,7 +215,7 @@ def _qderiv_actuator_passive_actuation_sparse(
elemid = M_elemid[dofi, dofj]
if elemid >= 0:
contrib = moment_i * moment_j * vel
wp.atomic_add(qDeriv_out[worldid, 0], elemid, contrib)
wp.atomic_add(qDeriv_out[worldid], elemid, contrib)
@wp.kernel
@@ -274,32 +225,28 @@ def _qderiv_actuator_passive(
opt_disableflags: int,
dof_damping: wp.array2d[float],
dof_dampingpoly: wp.array2d[wp.vec2],
is_sparse: bool,
M_elemid: wp.array2d[int],
# Data in:
qvel_in: wp.array2d[float],
M_in: wp.array3d[float],
M_in: wp.array2d[float],
# In:
Mi: wp.array[int],
Mj: wp.array[int],
qDeriv_in: wp.array3d[float],
qDeriv_in: wp.array2d[float],
# Out:
qDeriv_out: wp.array3d[float],
qDeriv_out: wp.array2d[float],
):
worldid, elemid = wp.tid()
dofiid = Mi[elemid]
dofjid = Mj[elemid]
# Off-pattern (dofiid, dofjid) pairs have no CSR entry (madr < 0).
madr = M_elemid[dofiid, dofjid]
if madr < 0:
return
if is_sparse:
if madr >= 0:
qderiv = qDeriv_in[worldid, 0, madr]
else:
qderiv = 0.0
else:
qderiv = qDeriv_in[worldid, dofiid, dofjid]
qderiv = qDeriv_in[worldid, madr]
if not (opt_disableflags & DisableBit.DAMPER) and dofiid == dofjid:
damping = dof_damping[worldid % dof_damping.shape[0], dofiid]
@@ -309,14 +256,7 @@ def _qderiv_actuator_passive(
qderiv *= opt_timestep[worldid % opt_timestep.shape[0]]
if is_sparse:
if madr >= 0:
qDeriv_out[worldid, 0, madr] = M_in[worldid, 0, madr] - qderiv
else:
M = M_in[worldid, dofiid, dofjid] - qderiv
qDeriv_out[worldid, dofiid, dofjid] = M
if dofiid != dofjid:
qDeriv_out[worldid, dofjid, dofiid] = M
qDeriv_out[worldid, madr] = M_in[worldid, madr] - qderiv
# TODO(team): improve performance with tile operations?
@@ -330,7 +270,6 @@ def _qderiv_tendon_damping(
ten_J_colind: wp.array[int],
tendon_damping: wp.array2d[float],
tendon_dampingpoly: wp.array2d[wp.vec2],
is_sparse: bool,
M_elemid: wp.array2d[int],
# Data in:
ten_J_in: wp.array2d[float],
@@ -339,12 +278,17 @@ def _qderiv_tendon_damping(
Mi: wp.array[int],
Mj: wp.array[int],
# Out:
qDeriv_out: wp.array3d[float],
qDeriv_out: wp.array2d[float],
):
worldid, elemid = wp.tid()
dofiid = Mi[elemid]
dofjid = Mj[elemid]
# Off-pattern (dofiid, dofjid) pairs have no CSR entry (madr < 0).
madr = M_elemid[dofiid, dofjid]
if madr < 0:
return
qderiv = float(0.0)
tendon_damping_id = worldid % tendon_damping.shape[0]
for tenid in range(ntendon):
@@ -372,15 +316,7 @@ def _qderiv_tendon_damping(
qderiv *= opt_timestep[worldid % opt_timestep.shape[0]]
madr = M_elemid[dofiid, dofjid]
if is_sparse:
if madr >= 0:
qDeriv_out[worldid, 0, madr] -= qderiv
else:
qDeriv_out[worldid, dofiid, dofjid] -= qderiv
if dofiid != dofjid:
qDeriv_out[worldid, dofjid, dofiid] -= qderiv
qDeriv_out[worldid, madr] -= qderiv
@wp.kernel
@@ -557,7 +493,7 @@ def deriv_rne_body2jnt_sparse(
Dcfrcbody_in: wp.array3d[wp.spatial_vector],
flg_subtract: bool,
# Out:
qDeriv_out: wp.array3d[float],
qDeriv_out: wp.array2d[float],
):
"""Project body-space RNE derivatives into joint-space qDeriv (sparse)."""
worldid, elemid = wp.tid()
@@ -571,12 +507,12 @@ def deriv_rne_body2jnt_sparse(
term = wp.dot(cdof_in[worldid, i], dcfrc)
if flg_subtract:
wp.atomic_sub(qDeriv_out[worldid, 0], elemid, dt * term)
wp.atomic_sub(qDeriv_out[worldid], elemid, dt * term)
else:
wp.atomic_add(qDeriv_out[worldid, 0], elemid, dt * term)
wp.atomic_add(qDeriv_out[worldid], elemid, dt * term)
def deriv_rne_vel(m: Model, d: Data, out: wp.array3d[float], flg_subtract: bool = False):
def deriv_rne_vel(m: Model, d: Data, out: wp.array2d[float], flg_subtract: bool = False):
"""Compute RNE velocity derivatives and add/subtract from the output.
Implements the analytical derivative of inverse-dynamics Coriolis/centrifugal
@@ -585,7 +521,7 @@ def deriv_rne_vel(m: Model, d: Data, out: wp.array3d[float], flg_subtract: bool
Args:
m: The model (device).
d: The data (device).
out: D-structure output array (nworld, 1, nD) to accumulate RNE terms into.
out: D-structure output array (nworld, nD) to accumulate RNE terms into.
flg_subtract: If True, subtract the RNE derivatives from output instead of adding them.
"""
# TODO(team): consider caching these allocations
@@ -649,6 +585,535 @@ def deriv_rne_vel(m: Model, d: Data, out: wp.array3d[float], flg_subtract: bool
)
@wp.func
def _deriv_ellipsoid_fluid(
# Model:
opt_integrator: int,
geom_type: wp.array[int],
geom_size: wp.array2d[wp.vec3],
geom_fluid: wp.array2d[float],
# Data in:
xipos_in: wp.array2d[wp.vec3],
geom_xpos_in: wp.array2d[wp.vec3],
geom_xmat_in: wp.array2d[wp.mat33],
subtree_com_in: wp.array2d[wp.vec3],
cvel_in: wp.array2d[wp.spatial_vector],
# In:
worldid: int,
bodyid: int,
rootid: int,
geomadr: int,
geomnum: int,
cdof_i: wp.spatial_vector,
cdof_j: wp.spatial_vector,
wind: wp.vec3,
density: float,
viscosity: float,
) -> float:
"""Compute one body's ellipsoid fluid derivative contribution for a DOF pair.
Returns the scalar J_i^T @ B @ J_j contribution accumulated across geoms.
"""
is_implicitfast = opt_integrator == IntegratorType.IMPLICITFAST
# Body kinematics
xipos = xipos_in[worldid, bodyid]
cvel = cvel_in[worldid, bodyid]
ang_global = wp.spatial_top(cvel)
lin_global = wp.spatial_bottom(cvel)
subtree_root = subtree_com_in[worldid, rootid]
lin_com = lin_global - wp.cross(xipos - subtree_root, ang_global)
qderiv_contrib = float(0.0)
cdof_ang_i = wp.vec3(cdof_i[0], cdof_i[1], cdof_i[2])
cdof_lin_i = wp.vec3(cdof_i[3], cdof_i[4], cdof_i[5])
cdof_ang_j = wp.vec3(cdof_j[0], cdof_j[1], cdof_j[2])
cdof_lin_j = wp.vec3(cdof_j[3], cdof_j[4], cdof_j[5])
for g in range(geomnum):
geomid = geomadr + g
coef = geom_fluid[geomid, 0]
if coef <= 0.0:
continue
size = geom_size[worldid % geom_size.shape[0], geomid]
semiaxes = geom_semiaxes(size, geom_type[geomid])
geom_rot = geom_xmat_in[worldid, geomid]
geom_rotT = wp.transpose(geom_rot)
geom_pos = geom_xpos_in[worldid, geomid]
# compute local velocity
lin_point = lin_com + wp.cross(ang_global, geom_pos - xipos)
l_ang = geom_rotT @ ang_global
l_lin = geom_rotT @ lin_point
if wind[0] != 0.0 or wind[1] != 0.0 or wind[2] != 0.0:
l_lin -= geom_rotT @ wind
ang_vel = l_ang
lin_vel = l_lin
# read fluid coefficients
blunt_drag_coef = geom_fluid[geomid, 1]
slender_drag_coef = geom_fluid[geomid, 2]
ang_drag_coef = geom_fluid[geomid, 3]
kutta_lift_coef = geom_fluid[geomid, 4]
magnus_lift_coef = geom_fluid[geomid, 5]
virtual_mass = wp.vec3(geom_fluid[geomid, 6], geom_fluid[geomid, 7], geom_fluid[geomid, 8])
virtual_inertia = wp.vec3(geom_fluid[geomid, 9], geom_fluid[geomid, 10], geom_fluid[geomid, 11])
# ===== Build 6x6 B matrix as four 3x3 quadrants =====
# B = [[B00, B01], [B10, B11]] where rows are [ang; lin], cols are [ang; lin]
B00 = wp.mat33(0.0) # torque wrt ang_vel
B01 = wp.mat33(0.0) # torque wrt lin_vel
B10 = wp.mat33(0.0) # force wrt ang_vel
B11 = wp.mat33(0.0) # force wrt lin_vel
if density > 0.0:
# --- added mass forces ---
density_vm = density * virtual_mass
density_vi = density * virtual_inertia
virtual_lin_mom = wp.cw_mul(density_vm, lin_vel)
virtual_ang_mom = wp.cw_mul(density_vi, ang_vel)
# torque += cross(virtual_ang_mom, ang_vel) -> B00
B00 += wp.skew(virtual_ang_mom) - wp.skew(ang_vel) @ wp.diag(density_vi)
# torque += cross(virtual_lin_mom, lin_vel) -> B01
B01 += wp.skew(virtual_lin_mom) - wp.skew(lin_vel) @ wp.diag(density_vm)
# force += cross(virtual_lin_mom, ang_vel) -> B10
B10 += wp.skew(virtual_lin_mom)
# Da = d/d(vlm via lin) = _skew_neg(ang), scaled by density*vm -> B11
B11 += -wp.skew(ang_vel) @ wp.diag(density_vm)
# --- Magnus force: force += magnus_coef * cross(ang_vel, lin_vel) ---
volume = wp.static(4.0 / 3.0 * wp.pi) * semiaxes[0] * semiaxes[1] * semiaxes[2]
magnus_coef = magnus_lift_coef * density * volume
B10 -= wp.skew(lin_vel) * magnus_coef
B11 += wp.skew(ang_vel) * magnus_coef
# --- Kutta lift (3x3 -> B11) ---
a = (semiaxes[1] * semiaxes[2]) * (semiaxes[1] * semiaxes[2])
b = (semiaxes[2] * semiaxes[0]) * (semiaxes[2] * semiaxes[0])
c = (semiaxes[0] * semiaxes[1]) * (semiaxes[0] * semiaxes[1])
aa = a * a
bb = b * b
cc = c * c
x = lin_vel[0]
y = lin_vel[1]
z = lin_vel[2]
xx = x * x
yy = y * y
zz = z * z
xy = x * y
yz = y * z
xz = x * z
proj_denom = aa * xx + bb * yy + cc * zz
proj_num = a * xx + b * yy + c * zz
norm2 = xx + yy + zz
df_denom = wp.pi * kutta_lift_coef * density / wp.max(MJ_MINVAL, wp.sqrt(proj_denom * proj_num * norm2))
dfx_coef = yy * (a - b) + zz * (a - c)
dfy_coef = xx * (b - a) + zz * (b - c)
dfz_coef = xx * (c - a) + yy * (c - b)
proj_term = proj_num / wp.max(MJ_MINVAL, proj_denom)
cos_term = proj_num / wp.max(MJ_MINVAL, norm2)
D = wp.skew(wp.vec3(b - c, c - a, a - b)) * (2.0 * proj_num)
df_coef = wp.vec3(dfx_coef, dfy_coef, dfz_coef)
inner_term = wp.vec3(
aa * proj_term - a + cos_term,
bb * proj_term - b + cos_term,
cc * proj_term - c + cos_term,
)
D += wp.outer(df_coef, inner_term)
V = wp.diag(lin_vel)
D = V @ D @ V - wp.diag(df_coef * proj_num)
D *= df_denom
B11 += D
# --- viscous drag (3x3 -> B11) ---
d_max = wp.max(wp.max(semiaxes[0], semiaxes[1]), semiaxes[2])
d_min = wp.min(wp.min(semiaxes[0], semiaxes[1]), semiaxes[2])
d_mid = semiaxes[0] + semiaxes[1] + semiaxes[2] - d_max - d_min
eq_sphere_D = wp.static(2.0 / 3.0) * (semiaxes[0] + semiaxes[1] + semiaxes[2])
A_max = wp.pi * d_max * d_mid
A_proj = wp.pi * wp.sqrt(proj_denom / wp.max(MJ_MINVAL, proj_num))
norm = wp.sqrt(xx + yy + zz)
inv_norm = 1.0 / wp.max(MJ_MINVAL, norm)
lin_coef = viscosity * wp.static(3.0 * wp.pi) * eq_sphere_D
quad_coef = density * (A_proj * blunt_drag_coef + slender_drag_coef * (A_max - A_proj))
Aproj_coef = density * norm * (blunt_drag_coef - slender_drag_coef)
dA_coef = wp.pi / wp.max(MJ_MINVAL, wp.sqrt(proj_num * proj_num * proj_num * proj_denom))
dAproj_dv = wp.vec3(
Aproj_coef * dA_coef * a * x * (b * yy * (a - b) + c * zz * (a - c)),
Aproj_coef * dA_coef * b * y * (a * xx * (b - a) + c * zz * (b - c)),
Aproj_coef * dA_coef * c * z * (a * xx * (c - a) + b * yy * (c - b)),
)
inner = wp.length_sq(lin_vel)
D = (wp.outer(lin_vel, lin_vel) + wp.diag(wp.vec3(inner))) * (-quad_coef * inv_norm)
D -= wp.outer(lin_vel, dAproj_dv)
D -= wp.diag(wp.vec3(lin_coef))
B11 += D
# --- viscous torque (3x3 -> B00) ---
lin_visc_torq_coef = wp.pi * eq_sphere_D * eq_sphere_D * eq_sphere_D
I_max = wp.static(8.0 / 15.0 * wp.pi) * d_mid * d_max * d_max * d_max * d_max
II = wp.vec3(
ellipsoid_max_moment(semiaxes, 0),
ellipsoid_max_moment(semiaxes, 1),
ellipsoid_max_moment(semiaxes, 2),
)
mom_coef = wp.vec3(
ang_drag_coef * II[0] + slender_drag_coef * (I_max - II[0]),
ang_drag_coef * II[1] + slender_drag_coef * (I_max - II[1]),
ang_drag_coef * II[2] + slender_drag_coef * (I_max - II[2]),
)
mom_visc = wp.cw_mul(ang_vel, mom_coef)
norm_mom = wp.length(mom_visc)
density_scaled = density / wp.max(MJ_MINVAL, norm_mom)
mom_sq = -density_scaled * wp.cw_mul(wp.cw_mul(ang_vel, mom_coef), mom_coef)
torq_lin_coef = viscosity * lin_visc_torq_coef
diag_val = wp.dot(ang_vel, mom_sq) - torq_lin_coef
D = wp.outer(ang_vel, mom_sq) + wp.diag(wp.vec3(diag_val))
B00 += D
# symmetrize for implicitfast
if is_implicitfast:
B00 = 0.5 * (B00 + wp.transpose(B00))
B11 = 0.5 * (B11 + wp.transpose(B11))
B01_sym = 0.5 * (B01 + wp.transpose(B10))
B01 = B01_sym
B10 = wp.transpose(B01_sym)
# --- Jacobian transformation: J_i^T @ B @ J_j ---
offset = geom_pos - subtree_root
jac_p_i = cdof_lin_i + wp.cross(cdof_ang_i, offset)
la_i = geom_rotT @ cdof_ang_i
ll_i = geom_rotT @ jac_p_i
jac_p_j = cdof_lin_j + wp.cross(cdof_ang_j, offset)
la_j = geom_rotT @ cdof_ang_j
ll_j = geom_rotT @ jac_p_j
# B @ J_j = [B00 @ la_j + B01 @ ll_j; B10 @ la_j + B11 @ ll_j]
Bj_ang = B00 @ la_j + B01 @ ll_j
Bj_lin = B10 @ la_j + B11 @ ll_j
# J_i^T @ (B @ J_j) = la_i . Bj_ang + ll_i . Bj_lin
qderiv_contrib += wp.dot(la_i, Bj_ang) + wp.dot(ll_i, Bj_lin)
return qderiv_contrib
@wp.kernel
def _qderiv_ellipsoid_fluid(
# Model:
opt_timestep: wp.array[float],
opt_wind: wp.array[wp.vec3],
opt_density: wp.array[float],
opt_viscosity: wp.array[float],
opt_integrator: int,
body_parentid: wp.array[int],
body_rootid: wp.array[int],
body_geomnum: wp.array[int],
body_geomadr: wp.array[int],
dof_bodyid: wp.array[int],
geom_type: wp.array[int],
geom_size: wp.array2d[wp.vec3],
geom_fluid: wp.array2d[float],
body_fluid_ellipsoid_adr: wp.array[int],
body_isdofancestor: wp.array2d[int],
M_elemid: wp.array2d[int],
# Data in:
xipos_in: wp.array2d[wp.vec3],
geom_xpos_in: wp.array2d[wp.vec3],
geom_xmat_in: wp.array2d[wp.mat33],
subtree_com_in: wp.array2d[wp.vec3],
cdof_in: wp.array2d[wp.spatial_vector],
cvel_in: wp.array2d[wp.spatial_vector],
# In:
Mi: wp.array[int],
Mj: wp.array[int],
# Out:
qDeriv_out: wp.array2d[float],
):
"""Compute ellipsoid fluid force derivative contribution to qDeriv.
Parallelized over (world, fluid_body, elem). For each fluid body and DOF
pair, computes the 6x6 derivative matrix B in local geom frame via
_deriv_ellipsoid_fluid and accumulates J_i^T @ B @ J_j into qDeriv.
"""
worldid, fluid_idx, elemid = wp.tid()
bodyid = body_fluid_ellipsoid_adr[fluid_idx]
dofiid = Mi[elemid]
dofjid = Mj[elemid]
madr = M_elemid[dofiid, dofjid]
if madr < 0:
return
# dofiid is the "deeper" DOF (Mi >= Mj in tree ordering).
# Any body that has dofiid in its chain also has dofjid.
bodyid_i = dof_bodyid[dofiid]
if bodyid_i == 0:
return
if body_isdofancestor[bodyid, dofiid] == 0:
return
wind = opt_wind[worldid % opt_wind.shape[0]]
density = opt_density[worldid % opt_density.shape[0]]
viscosity = opt_viscosity[worldid % opt_viscosity.shape[0]]
timestep = opt_timestep[worldid % opt_timestep.shape[0]]
if density <= 0.0 and viscosity <= 0.0:
return
cdof_i = cdof_in[worldid, dofiid]
cdof_j = cdof_in[worldid, dofjid]
contrib = _deriv_ellipsoid_fluid(
opt_integrator,
geom_type,
geom_size,
geom_fluid,
xipos_in,
geom_xpos_in,
geom_xmat_in,
subtree_com_in,
cvel_in,
worldid,
bodyid,
body_rootid[bodyid],
body_geomadr[bodyid],
body_geomnum[bodyid],
cdof_i,
cdof_j,
wind,
density,
viscosity,
)
contrib *= timestep
if contrib != 0.0:
wp.atomic_add(qDeriv_out[worldid], madr, -contrib)
@wp.func
def _deriv_box_fluid(
# Model:
opt_integrator: int,
body_mass: wp.array2d[float],
body_inertia: wp.array2d[wp.vec3],
# In:
worldid: int,
bodyid: int,
lvel: wp.spatial_vector,
density: float,
viscosity: float,
) -> wp.spatial_matrix:
B = wp.spatial_matrix(0.0)
mass = body_mass[worldid % body_mass.shape[0], bodyid]
inertia = body_inertia[worldid % body_inertia.shape[0], bodyid]
scl = 6.0 / mass
# Equivalent inertia box
box = wp.vec3(
wp.sqrt(wp.max(MJ_MINVAL, inertia[1] + inertia[2] - inertia[0]) * scl),
wp.sqrt(wp.max(MJ_MINVAL, inertia[0] + inertia[2] - inertia[1]) * scl),
wp.sqrt(wp.max(MJ_MINVAL, inertia[0] + inertia[1] - inertia[2]) * scl),
)
# Viscous force and torque
if viscosity > 0.0:
diam = (box[0] + box[1] + box[2]) * wp.static(1.0 / 3.0)
# Rotational viscosity
visc_rot = -wp.pi * diam * diam * diam * viscosity
B[0, 0] += visc_rot
B[1, 1] += visc_rot
B[2, 2] += visc_rot
# Translational viscosity
visc_lin = wp.static(-3.0 * wp.pi) * diam * viscosity
B[3, 3] += visc_lin
B[4, 4] += visc_lin
B[5, 5] += visc_lin
# Lift and drag force and torque
if density > 0.0:
term0 = box[1] * box[1] * box[1] * box[1] + box[2] * box[2] * box[2] * box[2]
term1 = box[0] * box[0] * box[0] * box[0] + box[2] * box[2] * box[2] * box[2]
term2 = box[0] * box[0] * box[0] * box[0] + box[1] * box[1] * box[1] * box[1]
inv_32 = wp.static(1.0 / 32.0)
B[0, 0] -= density * box[0] * term0 * wp.abs(lvel[0]) * inv_32
B[1, 1] -= density * box[1] * term1 * wp.abs(lvel[1]) * inv_32
B[2, 2] -= density * box[2] * term2 * wp.abs(lvel[2]) * inv_32
B[3, 3] -= density * box[1] * box[2] * wp.abs(lvel[3])
B[4, 4] -= density * box[0] * box[2] * wp.abs(lvel[4])
B[5, 5] -= density * box[0] * box[1] * wp.abs(lvel[5])
if opt_integrator == IntegratorType.IMPLICITFAST:
B = 0.5 * (B + wp.transpose(B))
return B
@wp.func
def _get_jac_column_local(
# Model:
body_parentid: wp.array[int],
body_rootid: wp.array[int],
dof_bodyid: wp.array[int],
# Data in:
subtree_com_in: wp.array2d[wp.vec3],
cdof_in: wp.array2d[wp.spatial_vector],
# In:
point_global: wp.vec3,
bodyid: int,
dofid: int,
worldid: int,
b_imat: wp.mat33,
) -> wp.spatial_vector:
offset = point_global - subtree_com_in[worldid, body_rootid[bodyid]]
cdof_val = cdof_in[worldid, dofid]
cdof_ang = wp.spatial_top(cdof_val)
cdof_lin = wp.spatial_bottom(cdof_val)
jacp = cdof_lin + wp.cross(cdof_ang, offset)
jacr = cdof_ang
b_imat_T = wp.transpose(b_imat)
jacp_loc = b_imat_T @ jacp
jacr_loc = b_imat_T @ jacr
return wp.spatial_vector(jacr_loc, jacp_loc)
@wp.kernel
def _qderiv_box_fluid(
# Model:
opt_timestep: wp.array[float],
opt_wind: wp.array[wp.vec3],
opt_density: wp.array[float],
opt_viscosity: wp.array[float],
opt_integrator: int,
body_parentid: wp.array[int],
body_rootid: wp.array[int],
body_mass: wp.array2d[float],
body_inertia: wp.array2d[wp.vec3],
dof_bodyid: wp.array[int],
body_fluid_box_adr: wp.array[int],
body_isdofancestor: wp.array2d[int],
M_elemid: wp.array2d[int],
# Data in:
xipos_in: wp.array2d[wp.vec3],
ximat_in: wp.array2d[wp.mat33],
subtree_com_in: wp.array2d[wp.vec3],
cdof_in: wp.array2d[wp.spatial_vector],
cvel_in: wp.array2d[wp.spatial_vector],
# In:
Mi: wp.array[int],
Mj: wp.array[int],
# Out:
qDeriv_out: wp.array2d[float],
):
worldid, fluid_idx, elemid = wp.tid()
bodyid = body_fluid_box_adr[fluid_idx]
dofiid = Mi[elemid]
dofjid = Mj[elemid]
madr = M_elemid[dofiid, dofjid]
if madr < 0:
return
bodyid_i = dof_bodyid[dofiid]
if bodyid_i == 0:
return
if body_isdofancestor[bodyid, dofiid] == 0:
return
wind = opt_wind[worldid % opt_wind.shape[0]]
density = opt_density[worldid % opt_density.shape[0]]
viscosity = opt_viscosity[worldid % opt_viscosity.shape[0]]
timestep = opt_timestep[worldid % opt_timestep.shape[0]]
if density <= 0.0 and viscosity <= 0.0:
return
# Body velocity and kinematics
b_ipos = xipos_in[worldid, bodyid]
b_imat = ximat_in[worldid, bodyid]
subtree_root = subtree_com_in[worldid, body_rootid[bodyid]]
vel_subtree = cvel_in[worldid, bodyid]
v_subtree_ang = wp.vec3(vel_subtree[0], vel_subtree[1], vel_subtree[2])
v_subtree_lin = wp.vec3(vel_subtree[3], vel_subtree[4], vel_subtree[5])
lin_com = v_subtree_lin - wp.cross(b_ipos - subtree_root, v_subtree_ang)
b_imat_T = wp.transpose(b_imat)
v_local_ang = b_imat_T @ v_subtree_ang
v_local_lin = b_imat_T @ lin_com
wind_local = b_imat_T @ wind
lvel = wp.spatial_vector(v_local_ang, v_local_lin - wind_local)
B = _deriv_box_fluid(
opt_integrator,
body_mass,
body_inertia,
worldid,
bodyid,
lvel,
density,
viscosity,
)
# Jacobian transformation: J_i^T @ B @ J_j
J_i = _get_jac_column_local(
body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, b_ipos, bodyid, dofiid, worldid, b_imat
)
J_j = _get_jac_column_local(
body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, b_ipos, bodyid, dofjid, worldid, b_imat
)
contrib = wp.dot(J_i, B @ J_j) * timestep
if contrib != 0.0:
wp.atomic_add(qDeriv_out[worldid], madr, -contrib)
@event_scope
def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d[float]):
"""Analytical derivative of smooth forces w.r.t. velocities.
@@ -691,27 +1156,20 @@ def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d[float]):
],
outputs=[vel],
)
if m.is_sparse:
wp.launch(
_qderiv_actuator_passive_actuation_sparse,
dim=(d.nworld, m.nu),
inputs=[
m.M_elemid,
d.moment_rownnz,
d.moment_rowadr,
d.moment_colind,
d.actuator_moment,
vel,
],
outputs=[out],
)
else:
wp.launch(
_qderiv_actuator_passive_actuation_dense,
dim=(d.nworld, Mi.size),
inputs=[m.nu, d.moment_rownnz, d.moment_rowadr, d.moment_colind, d.actuator_moment, vel, Mi, Mj],
outputs=[out],
)
# out (qDeriv) is in M-structure.
wp.launch(
_qderiv_actuator_passive_actuation_sparse,
dim=(d.nworld, m.nu),
inputs=[
m.M_elemid,
d.moment_rownnz,
d.moment_rowadr,
d.moment_colind,
d.actuator_moment,
vel,
],
outputs=[out],
)
wp.launch(
_qderiv_actuator_passive,
dim=(d.nworld, Mi.size),
@@ -720,7 +1178,6 @@ def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d[float]):
m.opt.disableflags,
m.dof_damping,
m.dof_dampingpoly,
m.is_sparse,
m.M_elemid,
d.qvel,
d.M,
@@ -746,7 +1203,6 @@ def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d[float]):
m.ten_J_colind,
m.tendon_damping,
m.tendon_dampingpoly,
m.is_sparse,
m.M_elemid,
d.ten_J,
d.ten_velocity,
@@ -755,3 +1211,64 @@ def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d[float]):
],
outputs=[out],
)
if m.has_fluid:
if m.body_fluid_ellipsoid_adr.size > 0:
wp.launch(
_qderiv_ellipsoid_fluid,
dim=(d.nworld, m.body_fluid_ellipsoid_adr.size, Mi.size),
inputs=[
m.opt.timestep,
m.opt.wind,
m.opt.density,
m.opt.viscosity,
m.opt.integrator,
m.body_parentid,
m.body_rootid,
m.body_geomnum,
m.body_geomadr,
m.dof_bodyid,
m.geom_type,
m.geom_size,
m.geom_fluid,
m.body_fluid_ellipsoid_adr,
m.body_isdofancestor,
m.M_elemid,
d.xipos,
d.geom_xpos,
d.geom_xmat,
d.subtree_com,
d.cdof,
d.cvel,
Mi,
Mj,
],
outputs=[out],
)
if m.body_fluid_box_adr.size > 0:
wp.launch(
_qderiv_box_fluid,
dim=(d.nworld, m.body_fluid_box_adr.size, Mi.size),
inputs=[
m.opt.timestep,
m.opt.wind,
m.opt.density,
m.opt.viscosity,
m.opt.integrator,
m.body_parentid,
m.body_rootid,
m.body_mass,
m.body_inertia,
m.dof_bodyid,
m.body_fluid_box_adr,
m.body_isdofancestor,
m.M_elemid,
d.xipos,
d.ximat,
d.subtree_com,
d.cdof,
d.cvel,
Mi,
Mj,
],
outputs=[out],
)
+103 -130
View File
@@ -28,7 +28,6 @@ from mujoco.mjx.third_party.mujoco_warp._src import sensor
from mujoco.mjx.third_party.mujoco_warp._src import sleep
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 types
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.support import xfrc_accumulate
@@ -42,7 +41,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import GainType
from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType
from mujoco.mjx.third_party.mujoco_warp._src.types import JointType
from mujoco.mjx.third_party.mujoco_warp._src.types import Model
from mujoco.mjx.third_party.mujoco_warp._src.types import TileSet
from mujoco.mjx.third_party.mujoco_warp._src.types import OverflowType
from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType
from mujoco.mjx.third_party.mujoco_warp._src.types import vec10f
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel
@@ -219,46 +218,59 @@ def _next_activation(
act_out[worldid, j] = act
@wp.kernel
def _next_time(
# Model:
opt_timestep: wp.array[float],
is_sparse: bool,
# Data in:
nefc_in: wp.array[int],
time_in: wp.array[float],
efc_J_rownnz_in: wp.array2d[int],
efc_J_rowadr_in: wp.array2d[int],
nworld_in: int,
naconmax_in: int,
njmax_in: int,
njmax_nnz_in: int,
nacon_in: wp.array[int],
ncollision_in: wp.array[int],
# Data out:
time_out: wp.array[float],
):
worldid = wp.tid()
time_out[worldid] = time_in[worldid] + opt_timestep[worldid % opt_timestep.shape[0]]
nefc = nefc_in[worldid]
@cache_kernel
def _next_time_builder(warn_overflow: bool):
@wp.kernel(module="unique", enable_backward=False)
def _next_time(
# Model:
opt_timestep: wp.array[float],
is_sparse: bool,
# Data in:
nefc_in: wp.array[int],
time_in: wp.array[float],
efc_J_rownnz_in: wp.array2d[int],
efc_J_rowadr_in: wp.array2d[int],
nworld_in: int,
naconmax_in: int,
njmax_in: int,
njmax_nnz_in: int,
nacon_in: wp.array[int],
ncollision_in: wp.array[int],
# Data out:
time_out: wp.array[float],
overflow_out: wp.array[int],
):
worldid = wp.tid()
time_out[worldid] = time_in[worldid] + opt_timestep[worldid % opt_timestep.shape[0]]
nefc = nefc_in[worldid]
if nefc > njmax_in:
wp.printf("nefc overflow - please increase njmax to %u\n", nefc)
elif nefc > 0 and is_sparse:
efcid = wp.min(nefc, njmax_in) - 1
efc_nnz = efc_J_rowadr_in[worldid, efcid] + efc_J_rownnz_in[worldid, efcid]
if efc_nnz > njmax_nnz_in:
wp.printf("njmax_nnz overflow - please increase njmax_nnz to %u\n", efc_nnz)
if nefc > njmax_in:
if wp.static(warn_overflow):
wp.printf("nefc overflow - please increase njmax to %u\n", nefc)
overflow_out[worldid] = overflow_out[worldid] | OverflowType.NEFC
elif nefc > 0 and is_sparse:
efcid = wp.min(nefc, njmax_in) - 1
efc_nnz = efc_J_rowadr_in[worldid, efcid] + efc_J_rownnz_in[worldid, efcid]
if efc_nnz > njmax_nnz_in:
if wp.static(warn_overflow):
wp.printf("njmax_nnz overflow - please increase njmax_nnz to %u\n", efc_nnz)
overflow_out[worldid] = overflow_out[worldid] | OverflowType.NJMAX_NNZ
if worldid == 0:
ncollision = ncollision_in[0]
if ncollision > naconmax_in:
nconmax = int(wp.ceil(float(ncollision) / float(nworld_in)))
wp.printf("broadphase overflow - please increase nconmax to %u or naconmax to %u\n", nconmax, ncollision)
if worldid == 0 and wp.static(warn_overflow):
nconmax = int(wp.ceil(float(ncollision) / float(nworld_in)))
wp.printf("broadphase overflow - please increase nconmax to %u or naconmax to %u\n", nconmax, ncollision)
overflow_out[worldid] = overflow_out[worldid] | OverflowType.BROADPHASE
if nacon_in[0] > naconmax_in:
nconmax = int(wp.ceil(float(nacon_in[0]) / float(nworld_in)))
wp.printf("narrowphase overflow - please increase nconmax to %u or naconmax to %u\n", nconmax, nacon_in[0])
nacon = nacon_in[0]
if nacon > naconmax_in:
if worldid == 0 and wp.static(warn_overflow):
nconmax = int(wp.ceil(float(nacon) / float(nworld_in)))
wp.printf("narrowphase overflow - please increase nconmax to %u or naconmax to %u\n", nconmax, nacon)
overflow_out[worldid] = overflow_out[worldid] | OverflowType.NARROWPHASE
return _next_time
def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None):
@@ -309,7 +321,7 @@ def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None)
history.insert_ctrl_history(m, d)
wp.launch(
_next_time,
_next_time_builder(bool(m.opt.warn_overflow)),
dim=d.nworld,
inputs=[
m.opt.timestep,
@@ -325,12 +337,13 @@ def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None)
d.nacon,
d.ncollision,
],
outputs=[d.time],
outputs=[d.time, d.overflow],
)
wp.copy(d.qacc_warmstart, d.qacc)
if not (m.opt.disableflags & DisableBit.ISLAND) and (m.opt.enableflags & EnableBit.SLEEP):
sleep_enabled = bool(m.opt.enableflags & EnableBit.SLEEP) and not bool(m.opt.disableflags & DisableBit.ISLAND)
if sleep_enabled:
sleep.sleep(m, d)
fwd_velocity(m, d)
sleep.update_sleep(m, d)
@@ -354,7 +367,7 @@ def _compute_damping_deriv(
@wp.kernel
def _euler_damp_qfrc_sparse(
def _euler_damp_qfrc(
# Model:
opt_timestep: wp.array[float],
M_rownnz: wp.array[int],
@@ -362,46 +375,13 @@ def _euler_damp_qfrc_sparse(
# In:
damp_deriv: wp.array2d[float],
# Out:
M_integration_out: wp.array3d[float],
M_integration_out: wp.array2d[float],
):
worldid, tid = wp.tid()
timestep = opt_timestep[worldid % opt_timestep.shape[0]]
adr = M_rowadr[tid] + M_rownnz[tid] - 1
M_integration_out[worldid, 0, adr] += timestep * damp_deriv[worldid, tid]
@cache_kernel
def _tile_euler_dense(tile: TileSet):
@wp.kernel(module="unique", enable_backward=False)
def euler_dense(
# Model:
opt_timestep: wp.array[float],
# Data in:
M_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],
):
worldid, nodeid = wp.tid()
timestep = opt_timestep[worldid % opt_timestep.shape[0]]
TILE_SIZE = wp.static(tile.size)
dofid = adr_in[nodeid]
M_tile = wp.tile_load(M_in[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, 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)
Ma_tile = wp.tile_load(efc_Ma_in[worldid], shape=(TILE_SIZE,), offset=(dofid,))
L_tile = wp.tile_cholesky(qm_integration_tile, fill_mode="upper")
qacc_tile = wp.tile_cholesky_solve(L_tile, Ma_tile, fill_mode="upper")
wp.tile_store(qacc_out[worldid], qacc_tile, offset=(dofid))
return euler_dense
M_integration_out[worldid, adr] += timestep * damp_deriv[worldid, tid]
@event_scope
@@ -420,26 +400,18 @@ def euler(m: Model, d: Data):
outputs=[damp_deriv],
)
if m.is_sparse:
M = wp.clone(d.M)
qLD = wp.empty((d.nworld, 1, m.nC), dtype=float)
qLDiagInv = wp.empty((d.nworld, m.nv), dtype=float)
wp.launch(
_euler_damp_qfrc_sparse,
dim=(d.nworld, m.nv),
inputs=[m.opt.timestep, m.M_rownnz, m.M_rowadr, damp_deriv],
outputs=[M],
)
smooth.factor_solve_i(m, d, M, qLD, qLDiagInv, qacc, d.efc.Ma)
else:
for tile in m.M_tiles:
wp.launch_tiled(
_tile_euler_dense(tile),
dim=(d.nworld, tile.adr.size),
inputs=[m.opt.timestep, d.M, d.efc.Ma, damp_deriv, tile.adr],
outputs=[qacc],
block_dim=m.block_dim.euler_dense,
)
# Clone M, add the damping to the diagonal, and factor-solve. factor_solve_i factors each block
# per-block (packed dense and/or sparse LDL); the scratch qLD matches d.qLD.
M = wp.clone(d.M)
qLD = wp.empty_like(d.qLD)
qLDiagInv = wp.empty((d.nworld, m.nv), dtype=float)
wp.launch(
_euler_damp_qfrc,
dim=(d.nworld, m.nv),
inputs=[m.opt.timestep, m.M_rownnz, m.M_rowadr, damp_deriv],
outputs=[M],
)
smooth.factor_solve_i(m, d, M, qLD, qLDiagInv, qacc, d.efc.Ma)
_advance(m, d, qacc)
else:
_advance(m, d, d.qacc)
@@ -589,25 +561,18 @@ def rungekutta4(m: Model, d: Data):
def _map_m2d(
# Model:
mapM2D: wp.array[int],
is_sparse: bool,
# In:
qDi: wp.array[int],
qDj: wp.array[int],
qH_M: wp.array3d[float],
qH_M: wp.array2d[float],
# Data out:
qLU_out: wp.array3d[float],
qLU_out: wp.array2d[float],
):
# Scatter qH_M (M-structure) into the D-structure qLU via mapM2D.
worldid, elemid = wp.tid()
if is_sparse:
m_idx = mapM2D[elemid]
if m_idx >= 0:
qLU_out[worldid, 0, elemid] = qH_M[worldid, 0, m_idx]
else:
qLU_out[worldid, 0, elemid] = 0.0
m_idx = mapM2D[elemid]
if m_idx >= 0:
qLU_out[worldid, elemid] = qH_M[worldid, m_idx]
else:
i = qDi[elemid]
j = qDj[elemid]
qLU_out[worldid, 0, elemid] = qH_M[worldid, i, j]
qLU_out[worldid, elemid] = 0.0
@event_scope
@@ -623,7 +588,7 @@ def implicit(m: Model, d: Data):
wp.launch(
_map_m2d,
dim=(d.nworld, m.nD),
inputs=[m.mapM2D, m.is_sparse, m.qD_fullm_i, m.qD_fullm_j, qH_M],
inputs=[m.mapM2D, qH_M],
outputs=[d.qLU],
)
@@ -635,12 +600,9 @@ def implicit(m: Model, d: Data):
smooth.factor_solve_lu(m, d, d.qLU, qacc, d.efc.Ma)
_advance(m, d, qacc)
elif ~(m.opt.disableflags | ~(DisableBit.ACTUATION | DisableBit.SPRING | DisableBit.DAMPER)):
if m.is_sparse:
qDeriv = wp.empty((d.nworld, 1, m.nC), dtype=float)
qLD = wp.empty((d.nworld, 1, m.nC), dtype=float)
else:
qDeriv = wp.empty(d.M.shape, dtype=float)
qLD = wp.empty(d.M.shape, dtype=float)
# qDeriv is in M-structure; the scratch qLD matches d.qLD (per-block).
qDeriv = wp.empty((d.nworld, m.nC), dtype=float)
qLD = wp.empty_like(d.qLD)
qLDiagInv = wp.empty((d.nworld, m.nv), dtype=float)
derivative.deriv_smooth_vel(m, d, qDeriv)
qacc = wp.empty((d.nworld, m.nv), dtype=float)
@@ -665,7 +627,7 @@ def fwd_position(m: Model, d: Data, factorize: bool = True):
smooth.flex(m, d)
smooth.tendon(m, d)
sleep_enabled = not (m.opt.disableflags & DisableBit.ISLAND) and (m.opt.enableflags & EnableBit.SLEEP)
sleep_enabled = bool(m.opt.enableflags & EnableBit.SLEEP) and not bool(m.opt.disableflags & DisableBit.ISLAND)
if sleep_enabled and m.ntendon > 0:
sleep.wake_tendon(m, d)
@@ -679,12 +641,15 @@ def fwd_position(m: Model, d: Data, factorize: bool = True):
if sleep_enabled:
# pass 1
collision_driver.collision(m, d)
# check for newly awake
skip = wp.zeros(1, dtype=int)
sleep.wake_collision(m, d, skip)
# wake any sleeping tree touched by an awake one
sleep.wake_collision(m, d)
# snapshot the awake state pass 1 used, before update_sleep overwrites it. a body is "newly
# awakened" if it was asleep here but awake after update_sleep below.
awake_prev = wp.clone(d.body_awake)
sleep.update_sleep(m, d)
# pass 2: broadphase kernels early-return if skip[0] is 0
collision_driver.collision(m, d, skip)
# pass 2: passing awake_prev runs the incremental pass, emitting only pairs involving a
# newly-awakened body and appending them to the pass-1 buffer.
collision_driver.collision(m, d, awake_prev=awake_prev)
else:
collision_driver.collision(m, d)
@@ -695,7 +660,7 @@ def fwd_position(m: Model, d: Data, factorize: bool = True):
sleep.wake_equality(m, d)
sleep.update_sleep(m, d)
if m.ntree > 1 and not (m.opt.disableflags & types.DisableBit.ISLAND):
if sleep_enabled:
island.island(m, d)
smooth.transmission(m, d)
@@ -1142,7 +1107,6 @@ def _qfrc_actuator(
@wp.kernel
def _qfrc_actuator_gravcomp_limits(
# Model:
ngravcomp: int,
jnt_actfrclimited: wp.array[bool],
jnt_actgravcomp: wp.array[int],
jnt_actfrcrange: wp.array2d[wp.vec2],
@@ -1150,6 +1114,8 @@ def _qfrc_actuator_gravcomp_limits(
# Data in:
qfrc_gravcomp_in: wp.array2d[float],
qfrc_actuator_in: wp.array2d[float],
# In:
gravity_enabled: bool,
# Data out:
qfrc_actuator_out: wp.array2d[float],
):
@@ -1159,7 +1125,7 @@ def _qfrc_actuator_gravcomp_limits(
qfrc = qfrc_actuator_in[worldid, dofid]
# actuator-level gravity compensation, skip if added as passive force
if ngravcomp and jnt_actgravcomp[jntid]:
if gravity_enabled and jnt_actgravcomp[jntid]:
qfrc += qfrc_gravcomp_in[worldid, dofid]
# limits
@@ -1256,17 +1222,18 @@ def fwd_actuation(m: Model, d: Data):
],
outputs=[d.qfrc_actuator],
)
gravity_enabled = not (m.opt.disableflags & DisableBit.GRAVITY)
wp.launch(
_qfrc_actuator_gravcomp_limits,
dim=(d.nworld, m.nv),
inputs=[
m.ngravcomp,
m.jnt_actfrclimited,
m.jnt_actgravcomp,
m.jnt_actfrcrange,
m.dof_jntid,
d.qfrc_gravcomp,
d.qfrc_actuator,
gravity_enabled,
],
outputs=[d.qfrc_actuator],
)
@@ -1316,7 +1283,7 @@ def fwd_acceleration(m: Model, d: Data, factorize: bool = False):
d: The data object containing the current state and output arrays.
factorize: Flag to factorize inertia matrix.
"""
enable_sleep = bool(m.opt.enableflags & EnableBit.SLEEP)
enable_sleep = bool(m.opt.enableflags & EnableBit.SLEEP) and not bool(m.opt.disableflags & DisableBit.ISLAND)
wp.launch(
_qfrc_smooth(enable_sleep),
dim=(d.nworld, m.nv),
@@ -1333,7 +1300,12 @@ def fwd_acceleration(m: Model, d: Data, factorize: bool = False):
)
xfrc_accumulate(m, d, d.qfrc_smooth)
if factorize:
if enable_sleep:
# update the active-DOF set (needs contacts from fwd_position) and solve
# the smooth acceleration in compacted dense space.
island.update_active_dofs(m, d)
solver.smooth_solve_compact(m, d)
elif factorize:
smooth.factor_solve_i(m, d, d.M, d.qLD, d.qLDiagInv, d.qacc_smooth, d.qfrc_smooth)
else:
smooth.solve_m(m, d, d.qacc_smooth, d.qfrc_smooth)
@@ -1342,7 +1314,8 @@ def fwd_acceleration(m: Model, d: Data, factorize: bool = False):
@event_scope
def forward(m: Model, d: Data):
"""Forward dynamics."""
if not (m.opt.disableflags & DisableBit.ISLAND) and (m.opt.enableflags & EnableBit.SLEEP):
sleep_enabled = bool(m.opt.enableflags & EnableBit.SLEEP) and not bool(m.opt.disableflags & DisableBit.ISLAND)
if sleep_enabled:
sleep.wake(m, d)
sleep.update_sleep(m, d)
+1 -4
View File
@@ -108,10 +108,7 @@ def discrete_acc(m: Model, d: Data, qacc: wp.array2d[float]):
outputs=[qfrc],
)
elif m.opt.integrator == IntegratorType.IMPLICITFAST:
if m.is_sparse:
qDeriv = wp.empty((d.nworld, 1, m.nC), dtype=float)
else:
qDeriv = wp.empty((d.nworld, m.nv, m.nv), dtype=float)
qDeriv = wp.empty((d.nworld, m.nC), dtype=float)
derivative.deriv_smooth_vel(m, d, qDeriv)
mul_m(m, d, qfrc, d.qacc, M=qDeriv)
smooth.factor_solve_i(m, d, d.M, d.qLD, d.qLDiagInv, qacc, qfrc)
File diff suppressed because it is too large Load Diff
+88 -138
View File
@@ -18,8 +18,8 @@ import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src import types
from mujoco.mjx.third_party.mujoco_warp._src.types import ConstraintType
from mujoco.mjx.third_party.mujoco_warp._src.types import EqType
from mujoco.mjx.third_party.mujoco_warp._src.types import IslandSolverContext
from mujoco.mjx.third_party.mujoco_warp._src.types import ObjType
from mujoco.mjx.third_party.mujoco_warp._src.types import OverflowType
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
@@ -800,18 +800,17 @@ def _init_efc_arrays(
@event_scope
def compute_island_mapping(m: types.Model, d: types.Data, ctx: IslandSolverContext):
def compute_island_mapping(m: types.Model, d: types.Data):
"""Compute DOF/constraint island mappings after island discovery.
Populates island solver context arrays via ctx: nv, nefc, ne, nf,
iefcadr, nidof, map_dof2idof, map_idof2dof, dof_islandid, map_efc2iefc,
map_iefc2efc, efc_islandid. Also populates d.dof_island, d.efc.island,
d.island_idofadr, and d.island_dofadr.
Populates d.dof_island, d.efc.island, d.island_idofadr, d.island_dofadr,
d.island_nv, d.island_nefc, d.island_ne, d.island_nf, d.island_iefcadr,
d.nidof, d.map_dof2idof, d.map_idof2dof, d.dof_islandid, d.map_efc2iefc,
d.map_iefc2efc, d.efc_islandid.
Args:
m: Model.
d: Data.
ctx: IslandSolverContext.
"""
# Ensure dof_islandid / efc_islandid are allocated at the right shape
if d.dof_islandid.shape[1] != m.nv:
@@ -821,13 +820,6 @@ def compute_island_mapping(m: types.Model, d: types.Data, ctx: IslandSolverConte
if d.island_idofadr.shape[1] != m.ntree:
d.island_idofadr = wp.empty((d.nworld, m.ntree), dtype=int)
# Ensure island-local DOF arrays are allocated at the right shape
if d.iqacc.shape[1] != m.nv:
nw = d.nworld
d.iqacc = wp.empty((nw, m.nv), dtype=float)
d.iqacc_smooth = wp.empty((nw, m.nv), dtype=float)
d.iqfrc_smooth = wp.empty((nw, m.nv), dtype=float)
d.iqfrc_constraint = wp.empty((nw, m.nv), dtype=float)
wp.launch(
_init_island_arrays,
dim=(d.nworld, m.ntree),
@@ -838,7 +830,7 @@ def compute_island_mapping(m: types.Model, d: types.Data, ctx: IslandSolverConte
d.island_nefc,
d.island_ne,
d.island_nf,
d.island_efcadr,
d.island_iefcadr,
d.nidof,
],
)
@@ -905,7 +897,7 @@ def compute_island_mapping(m: types.Model, d: types.Data, ctx: IslandSolverConte
_island_scan_sizes,
dim=d.nworld,
inputs=[d.nisland],
outputs=[d.island_idofadr, d.island_nv, d.island_nefc, d.island_efcadr, d.nidof],
outputs=[d.island_idofadr, d.island_nv, d.island_nefc, d.island_iefcadr, d.nidof],
)
# 4. Map DOFs
@@ -930,7 +922,7 @@ def compute_island_mapping(m: types.Model, d: types.Data, ctx: IslandSolverConte
d.nefc,
d.njmax,
d.efc.island,
d.island_efcadr,
d.island_iefcadr,
d.island_ne,
d.island_nf,
d.efc.type,
@@ -941,131 +933,89 @@ def compute_island_mapping(m: types.Model, d: types.Data, ctx: IslandSolverConte
outputs=[d.island_nefc, d.map_efc2iefc, d.map_iefc2efc, d.efc_islandid],
)
# 6. Scan Sparse Rows (if sparse)
if m.is_sparse:
wp.launch(
_island_scan_sparse_rows,
dim=d.nworld,
inputs=[d.nisland, d.efc.J_rownnz, d.island_nefc, d.island_efcadr, d.map_iefc2efc],
outputs=[d.efc.iJ_rownnz, d.efc.iJ_rowadr],
)
# Active-DOF compaction (nvmax < nv).
#
# The active set is tracked per kinematic tree (the unit of coupling in the mass matrix)
# via the same tree_awake bookkeeping used by the sleep solver. Each step the active trees'
# DOFs are packed into a contiguous [0, ncdof) range so the dense factor/solve can run at
# size nvmax instead of nv. This is the active-set analog of compute_island_mapping.
@wp.kernel
def _reset_compact_maps(
# Model:
nv: int,
# Data in:
nvmax_pad_in: int,
# Data out:
dof_cdof_out: wp.array2d[int],
cdof_dof_out: wp.array2d[int],
):
worldid, idx = wp.tid()
if idx < nv:
dof_cdof_out[worldid, idx] = -1
# cdof_dof is nvmax_pad-wide: clear the whole row so the padded tail [ncdof, nvmax_pad)
# reads as -1 (the gather/solve run over nvmax_pad, not just the active ncdof).
if idx < nvmax_pad_in:
cdof_dof_out[worldid, idx] = -1
@wp.kernel
def _compact_dofs(
# Model:
ntree: int,
tree_dofadr: wp.array[int],
tree_dofnum: wp.array[int],
# Data in:
tree_awake_in: wp.array2d[int],
nvmax_in: int,
warn_overflow: bool,
# Data out:
ncdof_out: wp.array[int],
dof_cdof_out: wp.array2d[int],
cdof_dof_out: wp.array2d[int],
overflow_out: wp.array[int],
):
worldid = wp.tid()
count = int(0)
for t in range(ntree):
if tree_awake_in[worldid, t] == 1:
adr = tree_dofadr[t]
num = tree_dofnum[t]
for j in range(num):
dof = adr + j
if count < nvmax_in:
dof_cdof_out[worldid, dof] = count
cdof_dof_out[worldid, count] = dof
count += 1
if count > nvmax_in:
if warn_overflow:
wp.printf(
"nvmax overflow: world %d needs %d active DOFs but nvmax = %d (behavior undefined)\n",
worldid,
count,
nvmax_in,
)
overflow_out[worldid] = overflow_out[worldid] | OverflowType.NVMAX
ncdof_out[worldid] = nvmax_in
else:
ncdof_out[worldid] = count
@event_scope
def gather_island_inputs(m: types.Model, d: types.Data, ctx: IslandSolverContext):
"""Gather constraint and DOF arrays into island-local order.
Populates d.iefc (D, type, id, frictionloss, aref, J, J_colind) and
d.iqacc, d.iqacc_smooth, d.iqfrc_smooth.
Must be called after compute_island_mapping() and before per-island solving.
Args:
m: Model.
d: Data.
ctx: IslandSolverContext whose arrays are populated.
"""
# Gather constraint arrays and dense Jacobian (fused)
def update_active_dofs(m: types.Model, d: types.Data):
"""Rebuild the compaction maps (dof_cdof / cdof_dof) from tree_awake."""
wp.launch(
_gather_efc_and_jacobian,
dim=(d.nworld, d.njmax),
inputs=[
m.is_sparse,
d.nefc,
d.efc.D,
d.efc.type,
d.efc.id,
d.efc.frictionloss,
d.efc.aref,
d.efc.J,
d.efc.J_rownnz,
d.efc.J_rowadr,
d.efc.J_colind,
d.efc.iJ_rowadr,
d.njmax,
d.map_iefc2efc,
d.map_idof2dof,
d.map_dof2idof,
d.nidof,
],
outputs=[
d.efc.iD,
d.efc.itype,
d.efc.iid,
d.efc.ifrictionloss,
d.efc.iaref,
d.efc.iJ,
d.efc.iJ_colind,
],
_reset_compact_maps,
dim=(d.nworld, max(m.nv, d.nvmax_pad)),
inputs=[m.nv, d.nvmax_pad],
outputs=[d.dof_cdof, d.cdof_dof],
)
# Gather DOF arrays
wp.launch(
_gather_dof_arrays,
dim=(d.nworld, m.nv),
inputs=[
d.qacc,
d.qacc_smooth,
d.qfrc_smooth,
d.nidof,
d.map_idof2dof,
],
outputs=[
d.iqacc,
d.iqacc_smooth,
d.iqfrc_smooth,
],
)
@event_scope
def scatter_island_results(m: types.Model, d: types.Data, ctx: IslandSolverContext, scatter_Ma: bool):
"""Scatter island-local solver results back to global arrays.
Reads ctx qacc, qfrc_constraint, Ma, d.efc iforce, istate and
writes them back to d.qacc, d.qfrc_constraint, d.efc.Ma, d.efc.force,
d.efc.state. Unconstrained DOFs receive qacc_smooth and zero qfrc.
Args:
m: Model.
d: Data.
ctx: IslandSolverContext which contains results.
scatter_Ma: Whether to scatter Ma for Euler/implicit integrators.
"""
# Scatter DOF results (and optionally Ma for Euler/implicit integrators)
wp.launch(
_scatter_dof_arrays,
dim=(d.nworld, m.nv),
inputs=[
d.qacc_smooth,
d.qfrc_smooth,
d.dof_island,
d.iqacc,
d.iqfrc_constraint,
ctx.Ma,
d.map_dof2idof,
scatter_Ma,
],
outputs=[
d.qacc,
d.qfrc_constraint,
d.efc.Ma,
],
)
# Scatter constraint results (force and state from island-local arrays)
wp.launch(
_scatter_efc_arrays,
dim=(d.nworld, d.njmax),
inputs=[
d.nefc,
d.njmax,
d.map_iefc2efc,
d.efc.iforce,
d.efc.istate,
],
outputs=[
d.efc.force,
d.efc.state,
],
_compact_dofs,
dim=(d.nworld,),
inputs=[m.ntree, m.tree_dofadr, m.tree_dofnum, d.tree_awake, d.nvmax, m.opt.warn_overflow],
outputs=[d.ncdof, d.dof_cdof, d.cdof_dof, d.overflow],
)
+14 -19
View File
@@ -41,7 +41,7 @@ def _pow4(val: float) -> float:
@wp.func
def _geom_semiaxes(size: wp.vec3, geom_type: int) -> wp.vec3: # kernel_analyzer: ignore
def geom_semiaxes(size: wp.vec3, geom_type: int) -> wp.vec3: # kernel_analyzer: ignore
if geom_type == GeomType.SPHERE:
r = size[0]
return wp.vec3(r, r, r)
@@ -61,7 +61,7 @@ def _geom_semiaxes(size: wp.vec3, geom_type: int) -> wp.vec3: # kernel_analyzer
@wp.func
def _ellipsoid_max_moment(size: wp.vec3, dir: int) -> float:
def ellipsoid_max_moment(size: wp.vec3, dir: int) -> float:
d0 = size[dir]
d1 = size[(dir + 1) % 3]
d2 = size[(dir + 2) % 3]
@@ -368,7 +368,7 @@ def _fluid_force(
continue
size = geom_size[worldid % geom_size.shape[0], geomid]
semiaxes = _geom_semiaxes(size, geom_type[geomid])
semiaxes = geom_semiaxes(size, geom_type[geomid])
geom_rot = geom_xmat_in[worldid, geomid]
geom_rotT = wp.transpose(geom_rot)
geom_pos = geom_xpos_in[worldid, geomid]
@@ -430,12 +430,8 @@ def _fluid_force(
proj_denom = _pow4(s12) * _pow2(l_lin[0]) + _pow4(s20) * _pow2(l_lin[1]) + _pow4(s01) * _pow2(l_lin[2])
proj_num = _pow2(s12 * l_lin[0]) + _pow2(s20 * l_lin[1]) + _pow2(s01 * l_lin[2])
A_proj = 0.0
cos_alpha = 0.0
if proj_num > MJ_MINVAL and proj_denom > MJ_MINVAL:
A_proj = wp.pi * wp.sqrt(proj_denom / wp.max(MJ_MINVAL, proj_num))
if lin_speed > MJ_MINVAL:
cos_alpha = proj_num / wp.max(MJ_MINVAL, lin_speed * proj_denom)
A_proj = wp.pi * wp.sqrt(proj_denom / wp.max(MJ_MINVAL, proj_num))
cos_alpha = proj_num / wp.max(MJ_MINVAL, lin_speed * proj_denom)
norm = wp.vec3(
_pow2(s12) * l_lin[0],
@@ -453,9 +449,9 @@ def _fluid_force(
lin_visc_torq_coef = wp.pi * eq_sphere_D * eq_sphere_D * eq_sphere_D
I_max = wp.static(8.0 / 15.0 * wp.pi) * d_mid * _pow4(d_max)
II0 = _ellipsoid_max_moment(semiaxes, 0)
II1 = _ellipsoid_max_moment(semiaxes, 1)
II2 = _ellipsoid_max_moment(semiaxes, 2)
II0 = ellipsoid_max_moment(semiaxes, 0)
II1 = ellipsoid_max_moment(semiaxes, 1)
II2 = ellipsoid_max_moment(semiaxes, 2)
mom_visc = wp.vec3(
l_ang[0] * (ang_drag_coef * II0 + slender_drag_coef * (I_max - II0)),
@@ -573,7 +569,7 @@ def _qfrc_passive(
qfrc_gravcomp_in: wp.array2d[float],
qfrc_fluid_in: wp.array2d[float],
# In:
gravcomp: bool,
gravity_enabled: bool,
# Data out:
qfrc_passive_out: wp.array2d[float],
):
@@ -582,7 +578,7 @@ def _qfrc_passive(
qfrc_passive += qfrc_damper_in[worldid, dofid]
# add gravcomp unless added by actuators
if gravcomp and not jnt_actgravcomp[dof_jntid[dofid]]:
if gravity_enabled and not jnt_actgravcomp[dof_jntid[dofid]]:
qfrc_passive += qfrc_gravcomp_in[worldid, dofid]
# add fluid force
@@ -876,10 +872,9 @@ def passive(m: Model, d: Data):
outputs=[d.qfrc_spring],
)
gravcomp = m.ngravcomp and not (m.opt.disableflags & DisableBit.GRAVITY)
if gravcomp:
d.qfrc_gravcomp.zero_()
gravity_enabled = not (m.opt.disableflags & DisableBit.GRAVITY)
d.qfrc_gravcomp.zero_()
if gravity_enabled:
wp.launch(
_gravity_force,
dim=(d.nworld, m.nbody - 1, m.nv),
@@ -912,7 +907,7 @@ def passive(m: Model, d: Data):
d.qfrc_damper,
d.qfrc_gravcomp,
d.qfrc_fluid,
gravcomp,
gravity_enabled,
],
outputs=[
d.qfrc_passive,
+37 -46
View File
@@ -13,6 +13,7 @@
# limitations under the License.
# ==============================================================================
import dataclasses
from typing import Tuple
import warp as wp
@@ -555,25 +556,10 @@ def render(m: Model, d: Data, rc: RenderContext):
cast_ray_first_hit = _make_cast_ray(geom_ray_types, first_hit=True)
compute_lighting = _make_compute_lighting(cast_ray_first_hit)
# Extract static parameters to avoid capturing rc and m in the kernel closure.
static_use_precomputed_rays = rc.use_precomputed_rays
static_znear = rc.znear
static_enable_backface_culling = rc.enable_backface_culling
static_render_skybox = rc.render_skybox
static_skybox_tex_id = rc.skybox_tex_id
static_skybox_face_width = rc.skybox_face_width
static_use_textures = rc.use_textures
static_enable_specular = rc.enable_specular
static_enable_emission = rc.enable_emission
static_use_ambient_lighting = rc.use_ambient_lighting
static_headlight_active = rc.headlight_active
static_headlight_ambient = rc.headlight_ambient
static_headlight_diffuse = rc.headlight_diffuse
static_headlight_specular = rc.headlight_specular
static_enable_per_light_ambient = rc.enable_per_light_ambient
static_light_attenuation_is_default = rc.light_attenuation_is_default
static_has_spot_lights = rc.has_spot_lights
static_nlight = m.nlight
# Static parameters extracted for JAX FFI closure.
rc_static = {f.name: getattr(rc, f.name) for f in dataclasses.fields(rc) if f.type in (int, bool, float, wp.vec3)}
rc_static["enable_specular_or_emission"] = rc.enable_specular or rc.enable_emission
M_NLIGHT = m.nlight
@wp.kernel(module="unique", enable_backward=False)
def _render_megakernel(
@@ -641,6 +627,8 @@ def render(m: Model, d: Data, rc: RenderContext):
flex_rgba: wp.array[wp.vec4],
flex_geom_flexid: wp.array[int],
flex_geom_edgeid: wp.array[int],
skybox_tex_id: wp.array[int],
skybox_face_width: wp.array[int],
textures: wp.array[wp.Texture2D],
# Out:
rgb_out: wp.array2d[wp.uint32],
@@ -669,7 +657,7 @@ def render(m: Model, d: Data, rc: RenderContext):
# Map active camera index to MuJoCo camera ID
mujoco_cam_id = cam_id_map[camid]
if wp.static(static_use_precomputed_rays):
if wp.static(rc_static["use_precomputed_rays"]):
ray_dir_local_cam = ray[rayid]
else:
img_w = cam_res[camid][0]
@@ -685,7 +673,7 @@ def render(m: Model, d: Data, rc: RenderContext):
img_h,
px,
py,
wp.static(static_znear),
wp.static(rc_static["znear"]),
)
ray_origin_world = cam_xpos_in[worldid, mujoco_cam_id]
@@ -717,7 +705,7 @@ def render(m: Model, d: Data, rc: RenderContext):
ray_origin_world,
ray_dir_world,
float(MJ_MAXVAL),
wp.static(static_enable_backface_culling),
wp.static(rc_static["enable_backface_culling"]),
)
if render_seg[camid] and geom_id != -1:
@@ -728,10 +716,11 @@ def render(m: Model, d: Data, rc: RenderContext):
# Early Out
if geom_id == -1:
if wp.static(static_render_skybox) and render_rgb[camid]:
if wp.static(rc_static["render_skybox"]) and render_rgb[camid]:
skybox_id = skybox_tex_id[worldid % skybox_tex_id.shape[0]]
skybox_color = sample_skybox(
textures[wp.static(static_skybox_tex_id)],
wp.static(1.0 / float(static_skybox_face_width)),
textures[skybox_id],
1.0 / float(skybox_face_width[worldid % skybox_face_width.shape[0]]),
ray_dir_world,
)
rgb_out[worldid, rgb_adr[camid] + rayid_local] = pack_rgba_to_uint32(
@@ -765,7 +754,7 @@ def render(m: Model, d: Data, rc: RenderContext):
base_color = wp.vec3(color[0], color[1], color[2])
if wp.static(static_use_textures):
if wp.static(rc_static["use_textures"]):
if geom_id != -2:
mat_id = geom_matid[worldid % geom_matid.shape[0], geom_id]
if mat_id >= 0:
@@ -793,27 +782,27 @@ def render(m: Model, d: Data, rc: RenderContext):
mat_spec = DEFAULT_MAT_SPECULAR
mat_shin_exp = DEFAULT_MAT_SHININESS_EXPONENT
mat_emis = DEFAULT_MAT_EMISSION
if wp.static(static_enable_specular or static_enable_emission):
if wp.static(rc_static["enable_specular_or_emission"]):
if geom_id != -2:
mat_id_for_spec = geom_matid[worldid % geom_matid.shape[0], geom_id]
if mat_id_for_spec >= 0:
if wp.static(static_enable_specular):
if wp.static(rc_static["enable_specular"]):
mat_spec = mat_specular[worldid % mat_specular.shape[0], mat_id_for_spec]
mat_shin_exp = mat_shininess[worldid % mat_shininess.shape[0], mat_id_for_spec] * MAX_SHININESS
if wp.static(static_enable_emission):
if wp.static(rc_static["enable_emission"]):
mat_emis = mat_emission[worldid % mat_emission.shape[0], mat_id_for_spec]
result = wp.vec3(0.0)
if wp.static(static_enable_emission):
if wp.static(rc_static["enable_emission"]):
result = base_color * mat_emis
if wp.static(static_use_ambient_lighting):
if wp.static(static_headlight_active):
result = result + wp.cw_mul(base_color, wp.static(static_headlight_ambient))
elif wp.static(static_nlight == 0):
if wp.static(rc_static["use_ambient_lighting"]):
if wp.static(rc_static["headlight_active"]):
result = result + wp.cw_mul(base_color, wp.static(rc_static["headlight_ambient"]))
elif wp.static(M_NLIGHT == 0):
result = result + base_color * NO_LIGHT_AMBIENT_FALLBACK
if wp.static(static_enable_per_light_ambient):
for light_index in range(wp.static(static_nlight)):
if wp.static(rc_static["enable_per_light_ambient"]):
for light_index in range(wp.static(M_NLIGHT)):
if light_active[worldid % light_active.shape[0], light_index]:
result = result + wp.cw_mul(base_color, light_ambient[worldid % light_ambient.shape[0], light_index])
@@ -830,7 +819,7 @@ def render(m: Model, d: Data, rc: RenderContext):
light_diffuse_worldid = light_diffuse[worldid % light_diffuse.shape[0]]
light_specular_worldid = light_specular[worldid % light_specular.shape[0]]
# Apply Lighting for each light
for light_index in range(wp.static(static_nlight)):
for light_index in range(wp.static(M_NLIGHT)):
diff_rgb, spec_rgb = compute_lighting(
geom_type,
geom_dataid,
@@ -869,15 +858,15 @@ def render(m: Model, d: Data, rc: RenderContext):
view_dir,
mat_spec,
mat_shin_exp,
wp.static(static_enable_backface_culling),
wp.static(static_enable_specular),
wp.static(static_light_attenuation_is_default),
wp.static(static_has_spot_lights),
wp.static(rc_static["enable_backface_culling"]),
wp.static(rc_static["enable_specular"]),
wp.static(rc_static["light_attenuation_is_default"]),
wp.static(rc_static["has_spot_lights"]),
)
result = result + wp.cw_mul(base_color, diff_rgb) + spec_rgb
# Apply Headlight
if wp.static(static_headlight_active):
if wp.static(rc_static["headlight_active"]):
cam_pos = ray_origin_world
cam_fwd = -cam_mat_world[:, 2]
hl_diff, hl_spec = compute_lighting(
@@ -911,15 +900,15 @@ def render(m: Model, d: Data, rc: RenderContext):
wp.vec3(1.0, 0.0, 0.0),
0.0,
0.0,
wp.static(static_headlight_diffuse),
wp.static(static_headlight_specular),
wp.static(rc_static["headlight_diffuse"]),
wp.static(rc_static["headlight_specular"]),
normal,
hit_point,
view_dir,
mat_spec,
mat_shin_exp,
wp.static(static_enable_backface_culling),
wp.static(static_enable_specular),
wp.static(rc_static["enable_backface_culling"]),
wp.static(rc_static["enable_specular"]),
True,
False,
)
@@ -1000,6 +989,8 @@ def render(m: Model, d: Data, rc: RenderContext):
rc.flex_rgba,
rc.flex_geom_flexid,
rc.flex_geom_edgeid,
rc.skybox_tex_id,
rc.skybox_face_width,
rc.textures,
],
outputs=[
+164 -135
View File
@@ -27,7 +27,6 @@ from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import sdf
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXCONPAIR
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
from mujoco.mjx.third_party.mujoco_warp._src.types import TACTILE_DEPTH_SEMANTICS
from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType
from mujoco.mjx.third_party.mujoco_warp._src.types import ConstraintType
from mujoco.mjx.third_party.mujoco_warp._src.types import ContactType
@@ -37,6 +36,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
from mujoco.mjx.third_party.mujoco_warp._src.types import JointType
from mujoco.mjx.third_party.mujoco_warp._src.types import Model
from mujoco.mjx.third_party.mujoco_warp._src.types import ObjType
from mujoco.mjx.third_party.mujoco_warp._src.types import OverflowType
from mujoco.mjx.third_party.mujoco_warp._src.types import SensorType
from mujoco.mjx.third_party.mujoco_warp._src.types import Stage
from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType
@@ -152,46 +152,27 @@ def _cam_projection(
xpos = cam_xpos_in[worldid, refid]
xmat = cam_xmat_in[worldid, refid]
translation = wp.mat44f(1.0, 0.0, 0.0, -xpos[0], 0.0, 1.0, 0.0, -xpos[1], 0.0, 0.0, 1.0, -xpos[2], 0.0, 0.0, 0.0, 1.0)
rotation = wp.mat44f(
xmat[0, 0], xmat[1, 0], xmat[2, 0], 0.0,
xmat[0, 1], xmat[1, 1], xmat[2, 1], 0.0,
xmat[0, 2], xmat[1, 2], xmat[2, 2], 0.0,
0.0, 0.0, 0.0, 1.0,
) # fmt: skip
# Transform target position into camera-local frame: v = xmat^T @ (target - cam_pos)
v = wp.transpose(xmat) @ (target_xpos - xpos)
# focal transformation matrix (3 x 4)
# Compute focal lengths
if sensorsize[0] != 0.0 and sensorsize[1] != 0.0:
fx = intrinsic[0] / (sensorsize[0] + MJ_MINVAL) * float(res[0])
fy = intrinsic[1] / (sensorsize[1] + MJ_MINVAL) * float(res[1])
focal = wp.mat44f(-fx, 0.0, 0.0, 0.0, 0.0, fy, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0)
else:
f = 0.5 / wp.tan(fovy * wp.static(wp.pi / 360.0)) * float(res[1])
focal = wp.mat44f(-f, 0.0, 0.0, 0.0, 0.0, f, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0)
fx = f
fy = f
# image matrix (3 x 3)
image = wp.mat44f(
1.0, 0.0, 0.5 * float(res[0]), 0.0, 0.0, 1.0, 0.5 * float(res[1]), 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0
)
# Compute homogeneous coordinates
pixel_x = -fx * v[0]
pixel_y = fy * v[1]
# projection matrix (3 x 4): product of all 4 matrices
# TODO(team): compute proj directly
proj = image @ focal @ rotation @ translation
# projection matrix multiples homogeneous [x, y, z, 1] vectors
pos_hom = wp.vec4(target_xpos[0], target_xpos[1], target_xpos[2], 1.0)
# project world coordinates into pixel space, see:
# https://en.wikipedia.org/wiki/3D_projection#Mathematical_formula
pixel_coord_hom = proj @ pos_hom
# avoid dividing by tiny numbers
denom = pixel_coord_hom[2]
denom = v[2]
if wp.abs(denom) < MJ_MINVAL:
denom = wp.clamp(denom, -MJ_MINVAL, MJ_MINVAL)
# compute projection
return wp.vec2f(pixel_coord_hom[0], pixel_coord_hom[1]) / denom
return wp.vec2(pixel_x, pixel_y) / denom + 0.5 * wp.vec2(float(res[0]), float(res[1]))
@wp.kernel
@@ -281,6 +262,117 @@ def _limit_pos(
_write_scalar(sensor_type, sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, sensordata_out[worldid])
@wp.func
def _get_pos(
# Data in:
xpos_in: wp.array2d[wp.vec3],
xipos_in: wp.array2d[wp.vec3],
geom_xpos_in: wp.array2d[wp.vec3],
site_xpos_in: wp.array2d[wp.vec3],
cam_xpos_in: wp.array2d[wp.vec3],
# In:
worldid: int,
objtype: int,
objid: int,
) -> wp.vec3:
if objtype == ObjType.BODY:
return xipos_in[worldid, objid]
elif objtype == ObjType.XBODY:
return xpos_in[worldid, objid]
elif objtype == ObjType.GEOM:
return geom_xpos_in[worldid, objid]
elif objtype == ObjType.SITE:
return site_xpos_in[worldid, objid]
elif objtype == ObjType.CAMERA:
return cam_xpos_in[worldid, objid]
else:
return wp.vec3(0.0)
@wp.func
def _get_mat(
# Data in:
xmat_in: wp.array2d[wp.mat33],
ximat_in: wp.array2d[wp.mat33],
geom_xmat_in: wp.array2d[wp.mat33],
site_xmat_in: wp.array2d[wp.mat33],
cam_xmat_in: wp.array2d[wp.mat33],
# In:
worldid: int,
objtype: int,
objid: int,
) -> wp.mat33:
if objtype == ObjType.BODY:
return ximat_in[worldid, objid]
elif objtype == ObjType.XBODY:
return xmat_in[worldid, objid]
elif objtype == ObjType.GEOM:
return geom_xmat_in[worldid, objid]
elif objtype == ObjType.SITE:
return site_xmat_in[worldid, objid]
elif objtype == ObjType.CAMERA:
return cam_xmat_in[worldid, objid]
else:
return wp.identity(3, dtype=wp.float32)
@wp.func
def _get_body_id(
# Model:
geom_bodyid: wp.array[int],
site_bodyid: wp.array[int],
cam_bodyid: wp.array[int],
# In:
objtype: int,
objid: int,
) -> int:
if objtype == ObjType.BODY or objtype == ObjType.XBODY:
return objid
elif objtype == ObjType.GEOM:
return geom_bodyid[objid]
elif objtype == ObjType.SITE:
return site_bodyid[objid]
elif objtype == ObjType.CAMERA:
return cam_bodyid[objid]
else:
return 0
@wp.func
def _get_quat(
# Model:
body_iquat: wp.array2d[wp.quat],
geom_bodyid: wp.array[int],
geom_quat: wp.array2d[wp.quat],
site_bodyid: wp.array[int],
site_quat: wp.array2d[wp.quat],
cam_bodyid: wp.array[int],
cam_quat: wp.array2d[wp.quat],
# Data in:
xquat_in: wp.array2d[wp.quat],
# In:
worldid: int,
objtype: int,
objid: int,
) -> wp.quat:
if objtype == ObjType.BODY:
body_iquat_id = worldid % body_iquat.shape[0]
return math.mul_quat(xquat_in[worldid, objid], body_iquat[body_iquat_id, objid])
elif objtype == ObjType.XBODY:
return xquat_in[worldid, objid]
elif objtype == ObjType.GEOM:
geom_quat_id = worldid % geom_quat.shape[0]
return math.mul_quat(xquat_in[worldid, geom_bodyid[objid]], geom_quat[geom_quat_id, objid])
elif objtype == ObjType.SITE:
site_quat_id = worldid % site_quat.shape[0]
return math.mul_quat(xquat_in[worldid, site_bodyid[objid]], site_quat[site_quat_id, objid])
elif objtype == ObjType.CAMERA:
cam_quat_id = worldid % cam_quat.shape[0]
return math.mul_quat(xquat_in[worldid, cam_bodyid[objid]], cam_quat[cam_quat_id, objid])
else:
return wp.quat(1.0, 0.0, 0.0, 0.0)
@wp.func
def _frame_pos(
# Data in:
@@ -301,42 +393,12 @@ def _frame_pos(
refid: int,
reftype: int,
) -> wp.vec3:
if objtype == ObjType.BODY:
xpos = xipos_in[worldid, objid]
elif objtype == ObjType.XBODY:
xpos = xpos_in[worldid, objid]
elif objtype == ObjType.GEOM:
xpos = geom_xpos_in[worldid, objid]
elif objtype == ObjType.SITE:
xpos = site_xpos_in[worldid, objid]
elif objtype == ObjType.CAMERA:
xpos = cam_xpos_in[worldid, objid]
else: # UNKNOWN
xpos = wp.vec3(0.0)
xpos = _get_pos(xpos_in, xipos_in, geom_xpos_in, site_xpos_in, cam_xpos_in, worldid, objtype, objid)
if refid == -1:
return xpos
if reftype == ObjType.BODY:
xpos_ref = xipos_in[worldid, refid]
xmat_ref = ximat_in[worldid, refid]
elif objtype == ObjType.XBODY:
xpos_ref = xpos_in[worldid, refid]
xmat_ref = xmat_in[worldid, refid]
elif reftype == ObjType.GEOM:
xpos_ref = geom_xpos_in[worldid, refid]
xmat_ref = geom_xmat_in[worldid, refid]
elif reftype == ObjType.SITE:
xpos_ref = site_xpos_in[worldid, refid]
xmat_ref = site_xmat_in[worldid, refid]
elif reftype == ObjType.CAMERA:
xpos_ref = cam_xpos_in[worldid, refid]
xmat_ref = cam_xmat_in[worldid, refid]
else: # UNKNOWN
xpos_ref = wp.vec3(0.0)
xmat_ref = wp.identity(3, wp.float32)
xpos_ref = _get_pos(xpos_in, xipos_in, geom_xpos_in, site_xpos_in, cam_xpos_in, worldid, reftype, refid)
xmat_ref = _get_mat(xmat_in, ximat_in, geom_xmat_in, site_xmat_in, cam_xmat_in, worldid, reftype, refid)
return wp.transpose(xmat_ref) @ (xpos - xpos_ref)
@@ -356,40 +418,13 @@ def _frame_axis(
reftype: int,
frame_axis: int,
) -> wp.vec3:
if objtype == ObjType.BODY:
xmat = ximat_in[worldid, objid]
axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
elif objtype == ObjType.XBODY:
xmat = xmat_in[worldid, objid]
axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
elif objtype == ObjType.GEOM:
xmat = geom_xmat_in[worldid, objid]
axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
elif objtype == ObjType.SITE:
xmat = site_xmat_in[worldid, objid]
axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
elif objtype == ObjType.CAMERA:
xmat = cam_xmat_in[worldid, objid]
axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
else: # UNKNOWN
axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
xmat = _get_mat(xmat_in, ximat_in, geom_xmat_in, site_xmat_in, cam_xmat_in, worldid, objtype, objid)
axis = wp.vec3(xmat[0, frame_axis], xmat[1, frame_axis], xmat[2, frame_axis])
if refid == -1:
return axis
if reftype == ObjType.BODY:
xmat_ref = ximat_in[worldid, refid]
elif reftype == ObjType.XBODY:
xmat_ref = xmat_in[worldid, refid]
elif reftype == ObjType.GEOM:
xmat_ref = geom_xmat_in[worldid, refid]
elif reftype == ObjType.SITE:
xmat_ref = site_xmat_in[worldid, refid]
elif reftype == ObjType.CAMERA:
xmat_ref = cam_xmat_in[worldid, refid]
else: # UNKNOWN
xmat_ref = wp.identity(3, dtype=wp.float32)
xmat_ref = _get_mat(xmat_in, ximat_in, geom_xmat_in, site_xmat_in, cam_xmat_in, worldid, reftype, refid)
return wp.transpose(xmat_ref) @ axis
@@ -412,38 +447,36 @@ def _frame_quat(
refid: int,
reftype: int,
) -> wp.quat:
body_iquat_id = worldid % body_iquat.shape[0]
geom_quat_id = worldid % geom_quat.shape[0]
site_quat_id = worldid % site_quat.shape[0]
cam_quat_id = worldid % cam_quat.shape[0]
if objtype == ObjType.BODY:
quat = math.mul_quat(xquat_in[worldid, objid], body_iquat[body_iquat_id, objid])
elif objtype == ObjType.XBODY:
quat = xquat_in[worldid, objid]
elif objtype == ObjType.GEOM:
quat = math.mul_quat(xquat_in[worldid, geom_bodyid[objid]], geom_quat[geom_quat_id, objid])
elif objtype == ObjType.SITE:
quat = math.mul_quat(xquat_in[worldid, site_bodyid[objid]], site_quat[site_quat_id, objid])
elif objtype == ObjType.CAMERA:
quat = math.mul_quat(xquat_in[worldid, cam_bodyid[objid]], cam_quat[cam_quat_id, objid])
else: # UNKNOWN
quat = wp.quat(1.0, 0.0, 0.0, 0.0)
quat = _get_quat(
body_iquat,
geom_bodyid,
geom_quat,
site_bodyid,
site_quat,
cam_bodyid,
cam_quat,
xquat_in,
worldid,
objtype,
objid,
)
if refid == -1:
return quat
if reftype == ObjType.BODY:
refquat = math.mul_quat(xquat_in[worldid, refid], body_iquat[body_iquat_id, refid])
elif reftype == ObjType.XBODY:
refquat = xquat_in[worldid, refid]
elif reftype == ObjType.GEOM:
refquat = math.mul_quat(xquat_in[worldid, geom_bodyid[refid]], geom_quat[geom_quat_id, refid])
elif reftype == ObjType.SITE:
refquat = math.mul_quat(xquat_in[worldid, site_bodyid[refid]], site_quat[site_quat_id, refid])
elif reftype == ObjType.CAMERA:
refquat = math.mul_quat(xquat_in[worldid, cam_bodyid[refid]], cam_quat[cam_quat_id, refid])
else: # UNKNOWN
refquat = wp.quat(1.0, 0.0, 0.0, 0.0)
refquat = _get_quat(
body_iquat,
geom_bodyid,
geom_quat,
site_bodyid,
site_quat,
cam_bodyid,
cam_quat,
xquat_in,
worldid,
reftype,
refid,
)
return math.mul_quat(math.quat_inv(refquat), quat)
@@ -2263,23 +2296,14 @@ def _sensor_tactile(
vel_rel = vel_sensor - vel_other
forceT = wp.vec3(0.0, 0.0, 0.0)
if wp.static(TACTILE_DEPTH_SEMANTICS):
forceT[0] = -depth
else:
kMaxDepth = 0.05
pressure = depth / wp.max(kMaxDepth - depth, MJ_MINVAL)
force = wp.mul(normal, pressure)
forceT[0] = wp.dot(force, normal)
forceT[0] = -depth
if has_frame:
forceT[1] = wp.abs(wp.dot(vel_rel, tang1))
forceT[2] = wp.abs(wp.dot(vel_rel, tang2))
dim = sensor_dim[sensor_id] // 3
if wp.static(TACTILE_DEPTH_SEMANTICS):
wp.atomic_max(sensordata_out, worldid, sensor_adr[sensor_id] + 0 * dim + vertid, forceT[0])
else:
wp.atomic_add(sensordata_out, worldid, sensor_adr[sensor_id] + 0 * dim + vertid, forceT[0])
wp.atomic_max(sensordata_out, worldid, sensor_adr[sensor_id] + 0 * dim + vertid, forceT[0])
wp.atomic_add(sensordata_out, worldid, sensor_adr[sensor_id] + 1 * dim + vertid, forceT[1])
wp.atomic_add(sensordata_out, worldid, sensor_adr[sensor_id] + 2 * dim + vertid, forceT[2])
@@ -2308,6 +2332,7 @@ def _contact_match(
# Model:
opt_cone: int,
opt_contact_sensor_maxmatch: int,
opt_warn_overflow: bool,
body_parentid: wp.array[int],
geom_bodyid: wp.array[int],
site_type: wp.array[int],
@@ -2333,6 +2358,8 @@ def _contact_match(
efc_force_in: wp.array2d[float],
njmax_in: int,
nacon_in: wp.array[int],
# Data out:
overflow_out: wp.array[int],
# Out:
sensor_contact_nmatch_out: wp.array2d[int],
sensor_contact_matchid_out: wp.array3d[int],
@@ -2407,8 +2434,9 @@ def _contact_match(
contactmatchid = wp.atomic_add(sensor_contact_nmatch_out[worldid], contactsensorid, 1)
if contactmatchid >= opt_contact_sensor_maxmatch:
# TODO(team): alternative to wp.printf for reporting overflow?
wp.printf("contact match overflow: please increase Option.contact_sensor_maxmatch to %u\n", contactmatchid)
if opt_warn_overflow:
wp.printf("contact match overflow: please increase Option.contact_sensor_maxmatch to %u\n", contactmatchid)
wp.atomic_or(overflow_out, worldid, OverflowType.CONTACT_MATCH)
return
sensor_contact_matchid_out[worldid, contactsensorid, contactmatchid] = contactid
@@ -2585,6 +2613,7 @@ def sensor_acc(m: Model, d: Data):
inputs=[
m.opt.cone,
m.opt.contact_sensor_maxmatch,
m.opt.warn_overflow,
m.body_parentid,
m.geom_bodyid,
m.site_type,
@@ -2610,7 +2639,7 @@ def sensor_acc(m: Model, d: Data):
d.njmax,
d.nacon,
],
outputs=[sensor_contact_nmatch, sensor_contact_matchid, sensor_contact_criteria, sensor_contact_direction],
outputs=[d.overflow, sensor_contact_nmatch, sensor_contact_matchid, sensor_contact_criteria, sensor_contact_direction],
)
# sorting
+14 -50
View File
@@ -13,8 +13,6 @@
# limitations under the License.
# ==============================================================================
from typing import Optional
import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src import types
@@ -337,7 +335,6 @@ def _wake_kernel(
tree_awake_in: wp.array2d[int],
# Out:
tree_asleep_out: wp.array2d[int], # kernel_analyzer: ignore
nwoke_out: wp.array[int], # kernel_analyzer: ignore
):
worldid, treeid = wp.tid()
@@ -360,9 +357,7 @@ def _wake_kernel(
treeid,
0.0, # zero tolerance
):
woke = _wake_tree(ntree, worldid, treeid, K_AWAKE_VAL, tree_asleep_out)
if woke > 0:
wp.atomic_add(nwoke_out, worldid, woke)
_wake_tree(ntree, worldid, treeid, K_AWAKE_VAL, tree_asleep_out)
@wp.kernel
@@ -378,7 +373,6 @@ def _wake_collision_kernel(
nacon_in: wp.array[int],
# Out:
tree_asleep_out: wp.array2d[int], # kernel_analyzer: ignore
skip_out: wp.array[int], # kernel_analyzer: ignore
):
conid = wp.tid()
if conid >= nacon_in[0]:
@@ -413,9 +407,7 @@ def _wake_collision_kernel(
sleeping_tree = tree2 if awake1 == 1 else tree1
wakeval = tree_asleep_out[worldid, tree1] if awake1 == 1 else tree_asleep_out[worldid, tree2]
woke = _wake_tree(ntree, worldid, sleeping_tree, wakeval, tree_asleep_out)
if woke > 0:
wp.atomic_add(skip_out, 0, woke)
_wake_tree(ntree, worldid, sleeping_tree, wakeval, tree_asleep_out)
@wp.kernel
@@ -439,8 +431,6 @@ def _wake_tendon_kernel(
tree_awake_in: wp.array2d[int],
# Data out:
tree_asleep_out: wp.array2d[int],
# Out:
nwoke_out: wp.array[int],
):
worldid, tenid = wp.tid()
@@ -487,9 +477,7 @@ def _wake_tendon_kernel(
if t >= 0:
if tree_awake_in[worldid, t] == 0:
woke = _wake_tree(ntree, worldid, t, wakeval, tree_asleep_out)
if woke > 0:
wp.atomic_add(nwoke_out, worldid, woke)
_wake_tree(ntree, worldid, t, wakeval, tree_asleep_out)
@wp.func
@@ -560,8 +548,6 @@ def _wake_tendon_trees(
wakeval: int,
# Data out:
tree_asleep_out: wp.array2d[int],
# Out:
nwoke_out: wp.array[int],
):
"""Wakes up all sleeping trees associated with a tendon."""
if tenid < 0:
@@ -583,9 +569,7 @@ def _wake_tendon_trees(
if t >= 0:
if tree_awake_in[worldid, t] == 0:
woke = _wake_tree(ntree, worldid, t, wakeval, tree_asleep_out)
if woke > 0:
wp.atomic_add(nwoke_out, worldid, woke)
_wake_tree(ntree, worldid, t, wakeval, tree_asleep_out)
@wp.kernel
@@ -610,8 +594,6 @@ def _wake_equality_kernel(
tree_awake_in: wp.array2d[int],
# Data out:
tree_asleep_out: wp.array2d[int], # kernel_analyzer: ignore
# Out:
nwoke_out: wp.array[int], # kernel_analyzer: ignore
):
worldid, eqid = wp.tid()
@@ -651,15 +633,11 @@ def _wake_equality_kernel(
cycle1 = _sleep_cycle(tree_asleep_out, ntree, worldid, tree1)
cycle2 = _sleep_cycle(tree_asleep_out, ntree, worldid, tree2)
if cycle1 != cycle2:
w1 = _wake_tree(ntree, worldid, tree1, K_AWAKE_VAL, tree_asleep_out)
w2 = _wake_tree(ntree, worldid, tree2, K_AWAKE_VAL, tree_asleep_out)
if w1 + w2 > 0:
wp.atomic_add(nwoke_out, worldid, w1 + w2)
_wake_tree(ntree, worldid, tree1, K_AWAKE_VAL, tree_asleep_out)
_wake_tree(ntree, worldid, tree2, K_AWAKE_VAL, tree_asleep_out)
else:
sleeping_tree = tree1 if s1 == SleepState.ASLEEP else tree2
woke = _wake_tree(ntree, worldid, sleeping_tree, K_AWAKE_VAL, tree_asleep_out)
if woke > 0:
wp.atomic_add(nwoke_out, worldid, woke)
_wake_tree(ntree, worldid, sleeping_tree, K_AWAKE_VAL, tree_asleep_out)
elif eqtype == EqType.TENDON:
ten1 = id1
@@ -715,7 +693,6 @@ def _wake_equality_kernel(
ten1,
wakeval,
tree_asleep_out,
nwoke_out,
)
_wake_tendon_trees(
ntree,
@@ -732,7 +709,6 @@ def _wake_equality_kernel(
ten2,
wakeval,
tree_asleep_out,
nwoke_out,
)
# TODO(team): Implement waking for EqType.FLEX constraints.
@@ -741,7 +717,6 @@ def _wake_equality_kernel(
@event_scope
def wake(m: types.Model, d: types.Data):
"""Wakes sleeping trees due to user changes/perturbations."""
nwoke = wp.zeros((d.nworld,), dtype=int)
wp.launch(
_wake_kernel,
dim=(d.nworld, m.ntree),
@@ -757,16 +732,14 @@ def wake(m: types.Model, d: types.Data):
d.qfrc_applied,
d.xfrc_applied,
d.tree_awake,
d.tree_asleep,
],
outputs=[nwoke],
outputs=[d.tree_asleep],
)
@event_scope
def wake_collision(m: types.Model, d: types.Data, skip: Optional[wp.array] = None):
def wake_collision(m: types.Model, d: types.Data):
"""Wakes sleeping trees that touch awake trees."""
skip_out = skip if skip is not None else wp.zeros(1, dtype=int)
wp.launch(
_wake_collision_kernel,
dim=d.naconmax,
@@ -778,9 +751,8 @@ def wake_collision(m: types.Model, d: types.Data, skip: Optional[wp.array] = Non
d.contact.geom,
d.contact.worldid,
d.nacon,
d.tree_asleep,
],
outputs=[skip_out],
outputs=[d.tree_asleep],
)
@@ -790,7 +762,6 @@ def wake_tendon(m: types.Model, d: types.Data):
if m.ntendon == 0:
return
nwoke = wp.zeros((d.nworld,), dtype=int)
wp.launch(
_wake_tendon_kernel,
dim=(d.nworld, m.ntendon),
@@ -811,7 +782,7 @@ def wake_tendon(m: types.Model, d: types.Data):
d.ten_length,
d.tree_awake,
],
outputs=[d.tree_asleep, nwoke],
outputs=[d.tree_asleep],
)
@@ -821,7 +792,6 @@ def wake_equality(m: types.Model, d: types.Data):
if m.neq == 0:
return
nwoke = wp.zeros((d.nworld,), dtype=int)
wp.launch(
_wake_equality_kernel,
dim=(d.nworld, m.neq),
@@ -842,9 +812,8 @@ def wake_equality(m: types.Model, d: types.Data):
m.wrap_objid,
d.eq_active,
d.tree_awake,
d.tree_asleep,
],
outputs=[nwoke],
outputs=[d.tree_asleep],
)
@@ -928,7 +897,6 @@ def _build_cycles( # kernel_analyzer: ignore
tree_asleep_out: wp.array2d[int], # kernel_analyzer: ignore
qvel_out: wp.array2d[float], # kernel_analyzer: ignore
qacc_out: wp.array2d[float], # kernel_analyzer: ignore
nslept_out: wp.array[int], # kernel_analyzer: ignore
):
worldid = wp.tid()
@@ -937,7 +905,6 @@ def _build_cycles( # kernel_analyzer: ignore
if island_can_sleep_in[worldid, island_id] == 1:
first_tree = int(-1)
prev_tree = int(-1)
n = int(0)
for t in range(ntree):
if tree_island_in[worldid, t] == island_id:
if first_tree == -1:
@@ -945,7 +912,6 @@ def _build_cycles( # kernel_analyzer: ignore
if prev_tree != -1:
tree_asleep_out[worldid, prev_tree] = t
prev_tree = t
n += 1
# Zero velocities and accelerations
dofadr = tree_dofadr[t]
@@ -956,7 +922,6 @@ def _build_cycles( # kernel_analyzer: ignore
if first_tree != -1:
tree_asleep_out[worldid, prev_tree] = first_tree
wp.atomic_add(nslept_out, worldid, n)
# Sleep unconstrained trees
for t in range(ntree):
@@ -964,7 +929,6 @@ def _build_cycles( # kernel_analyzer: ignore
if island_id < 0 or island_id >= num_islands:
if tree_asleep_out[worldid, t] == -1:
tree_asleep_out[worldid, t] = t # self-cycle
wp.atomic_add(nslept_out, worldid, 1)
# Ensure sleeping tree dof velocity and acceleration remain exactly zero
if tree_asleep_out[worldid, t] >= 0:
@@ -1013,7 +977,6 @@ def sleep(m: types.Model, d: types.Data):
)
# 3. Build sleep cycles for sleeping islands and sleep unconstrained trees
nslept = wp.zeros((d.nworld,), dtype=int)
wp.launch(
_build_cycles,
dim=d.nworld,
@@ -1024,9 +987,10 @@ def sleep(m: types.Model, d: types.Data):
d.nisland,
d.tree_island,
island_can_sleep,
],
outputs=[
d.tree_asleep,
d.qvel,
d.qacc,
],
outputs=[nslept],
)
+236 -175
View File
@@ -823,7 +823,7 @@ def _crb_accumulate(
@wp.kernel
def _M_sparse(
def _M(
# Model:
dof_bodyid: wp.array[int],
dof_parentid: wp.array[int],
@@ -834,58 +834,25 @@ def _M_sparse(
cdof_in: wp.array2d[wp.spatial_vector],
crb_in: wp.array2d[vec10],
# Data out:
M_out: wp.array3d[float],
M_out: wp.array2d[float],
):
worldid, dofid = wp.tid()
bodyid = dof_bodyid[dofid]
madr_ij = M_rowadr[dofid] + M_rownnz[dofid] - 1
# init M(i,i) with armature inertia
M_out[worldid, 0, madr_ij] = dof_armature[worldid % dof_armature.shape[0], dofid]
M_out[worldid, madr_ij] = dof_armature[worldid % dof_armature.shape[0], dofid]
# precompute buf = crb_body_i * cdof_i
buf = math.inert_vec(crb_in[worldid, bodyid], cdof_in[worldid, dofid])
# sparse backward pass over ancestors
while dofid >= 0:
M_out[worldid, 0, madr_ij] += wp.dot(cdof_in[worldid, dofid], buf)
M_out[worldid, madr_ij] += wp.dot(cdof_in[worldid, dofid], buf)
madr_ij -= 1
dofid = dof_parentid[dofid]
@wp.kernel
def _M_dense(
# Model:
dof_bodyid: wp.array[int],
dof_parentid: wp.array[int],
dof_armature: wp.array2d[float],
# Data in:
cdof_in: wp.array2d[wp.spatial_vector],
crb_in: wp.array2d[vec10],
# Data out:
M_out: wp.array3d[float],
):
worldid, dofid = wp.tid()
bodyid = dof_bodyid[dofid]
# init M(i,i) with armature inertia.
M = dof_armature[worldid % dof_armature.shape[0], dofid]
# precompute buf = crb_body_i * cdof_i
buf = math.inert_vec(crb_in[worldid, bodyid], cdof_in[worldid, dofid])
M += wp.dot(cdof_in[worldid, dofid], buf)
M_out[worldid, dofid, dofid] = M
# sparse backward pass over ancestors
dofidi = dofid
dofid = dof_parentid[dofid]
while dofid >= 0:
Mij = wp.dot(cdof_in[worldid, dofid], buf)
M_out[worldid, dofidi, dofid] += Mij
M_out[worldid, dofid, dofidi] += Mij
dofid = dof_parentid[dofid]
@event_scope
def crb(m: Model, d: Data):
"""Computes composite rigid body inertias for each body and the joint-space inertia matrix.
@@ -900,17 +867,12 @@ def crb(m: Model, d: Data):
wp.launch(_crb_accumulate, dim=(d.nworld, body_tree.size), inputs=[m.body_parentid, d.crb, body_tree], outputs=[d.crb])
d.M.zero_()
if m.is_sparse:
wp.launch(
_M_sparse,
dim=(d.nworld, m.nv),
inputs=[m.dof_bodyid, m.dof_parentid, m.dof_armature, m.M_rownnz, m.M_rowadr, d.cdof, d.crb],
outputs=[d.M],
)
else:
wp.launch(
_M_dense, dim=(d.nworld, m.nv), inputs=[m.dof_bodyid, m.dof_parentid, m.dof_armature, d.cdof, d.crb], outputs=[d.M]
)
wp.launch(
_M,
dim=(d.nworld, m.nv),
inputs=[m.dof_bodyid, m.dof_parentid, m.dof_armature, m.M_rownnz, m.M_rowadr, d.cdof, d.crb],
outputs=[d.M],
)
@wp.kernel
@@ -923,11 +885,10 @@ def _tendon_armature(
tendon_armature: wp.array2d[float],
M_rownnz: wp.array[int],
M_rowadr: wp.array[int],
is_sparse: bool,
# Data in:
ten_J_in: wp.array2d[float],
# Data out:
M_out: wp.array3d[float],
M_out: wp.array2d[float],
):
worldid, tenid, dofid = wp.tid()
@@ -948,8 +909,8 @@ def _tendon_armature(
if ten_Ji == 0.0:
return
if is_sparse:
madr_ij = M_rowadr[dofid] + M_rownnz[dofid] - 1
# Walk the row's entries from the diagonal backward over ancestors.
madr_ij = M_rowadr[dofid] + M_rownnz[dofid] - 1
# sparse backward pass over ancestors
dofidi = dofid
@@ -971,13 +932,8 @@ def _tendon_armature(
Mij = armature * ten_Jj * ten_Ji
if is_sparse:
wp.atomic_add(M_out[worldid, 0], madr_ij, Mij)
madr_ij -= 1
else:
wp.atomic_add(M_out[worldid, dofidi], dofid, Mij)
if dofidi != dofid:
wp.atomic_add(M_out[worldid, dofid], dofidi, Mij)
wp.atomic_add(M_out[worldid], madr_ij, Mij)
madr_ij -= 1
dofid = dof_parentid[dofid]
@@ -996,7 +952,6 @@ def tendon_armature(m: Model, d: Data):
m.tendon_armature,
m.M_rownnz,
m.M_rowadr,
m.is_sparse,
d.ten_J,
],
outputs=[d.M],
@@ -1010,9 +965,9 @@ def _qLD_acc(
M_rowadr: wp.array[int],
# In:
qLD_updates_: wp.array[wp.vec3i],
L_in: wp.array3d[float],
L_in: wp.array2d[float],
# Out:
L_out: wp.array3d[float],
L_out: wp.array2d[float],
):
worldid, nodeid = wp.tid()
update = qLD_updates_[nodeid]
@@ -1020,12 +975,12 @@ def _qLD_acc(
Madr_i = M_rowadr[i] # Address of row being updated
diag_k = M_rowadr[k] + M_rownnz[k] - 1 # Address of diagonal element of k
# tmp = M(k,i) / M(k,k)
tmp = L_out[worldid, 0, Madr_ki] / L_out[worldid, 0, diag_k]
tmp = L_out[worldid, Madr_ki] / L_out[worldid, diag_k]
for j in range(M_rownnz[i]):
# M(i,j) -= M(k,j) * tmp
wp.atomic_sub(L_out[worldid, 0], Madr_i + j, L_in[worldid, 0, M_rowadr[k] + j] * tmp)
wp.atomic_sub(L_out[worldid], Madr_i + j, L_in[worldid, M_rowadr[k] + j] * tmp)
# M(k,i) = tmp
L_out[worldid, 0, Madr_ki] = tmp
L_out[worldid, Madr_ki] = tmp
@wp.kernel
@@ -1034,16 +989,35 @@ def _qLDiag_div(
M_rownnz: wp.array[int],
M_rowadr: wp.array[int],
# In:
L_in: wp.array3d[float],
L_in: wp.array2d[float],
# Out:
D_out: wp.array2d[float],
):
worldid, dofid = wp.tid()
diag_i = M_rowadr[dofid] + M_rownnz[dofid] - 1 # Address of diagonal element of i
D_out[worldid, dofid] = 1.0 / L_in[worldid, 0, diag_i]
D_out[worldid, dofid] = 1.0 / L_in[worldid, diag_i]
def _factor_i_sparse(m: Model, d: Data, M: wp.array3d[float], L: wp.array3d[float], D: wp.array2d[float]):
@wp.kernel
def _factor_simple(
# Model:
M_rownnz: wp.array[int],
M_rowadr: wp.array[int],
# Data in:
M_in: wp.array2d[float],
# In:
simple_dofs: wp.array[int],
# Out:
D_out: wp.array2d[float],
):
# A simple (decoupled) dof's whole factorization is D = 1/M(i,i): no L entries, no elimination.
worldid, s = wp.tid()
dofid = simple_dofs[s]
diag_i = M_rowadr[dofid] + M_rownnz[dofid] - 1
D_out[worldid, dofid] = 1.0 / M_in[worldid, diag_i]
def _factor_i_sparse(m: Model, d: Data, M: wp.array2d[float], L: wp.array2d[float], D: wp.array2d[float]):
"""Sparse L'*D*L factorization of inertia-like matrix M, assumed spd."""
wp.copy(L, M)
@@ -1055,36 +1029,46 @@ def _factor_i_sparse(m: Model, d: Data, M: wp.array3d[float], L: wp.array3d[floa
@cache_kernel
def _tile_cholesky_factorize(tile: TileSet):
"""Returns a kernel for dense Cholesky factorization of a tile."""
def _tile_cholesky_factorize_block(tile: TileSet):
# One diagonal block of `block_size` dofs per (world, block) tile group. tile_load_indexed gathers
# the block's dense slots from CSR via a precomputed per-slot index tile (block_elemid, laid out
# [block, slot]); structurally absent pairs carry an out-of-bounds index that reads as 0.
# Tile shapes must be compile-time constants, so the densify is inlined per kernel (sharing via a
# wp.func is not possible) -- keep it in sync with _tile_cholesky_factorize_solve_block.
block_size = tile.size
block_area = block_size * block_size
@wp.kernel(module="unique", enable_backward=False)
def cholesky_factorize(
def kernel(
# Model:
qLD_block_adr: wp.array[int],
# Data in:
M_in: wp.array3d[float],
M_in: wp.array2d[float],
# In:
adr: wp.array[int],
block_elemid: wp.array[int],
block_dof: wp.array[int],
# Out:
L_out: wp.array3d[float],
L_out: wp.array2d[float],
):
worldid, nodeid = wp.tid()
TILE_SIZE = wp.static(tile.size)
worldid, blk = wp.tid()
start = block_dof[blk]
dofid = adr[nodeid]
M_tile = wp.tile_load(M_in[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid))
L_tile = wp.tile_cholesky(M_tile, fill_mode="upper")
wp.tile_store(L_out[worldid], L_tile, offset=(dofid, dofid))
idx = wp.tile_load(block_elemid, shape=(block_area,), offset=(blk * block_area,), storage="shared")
block = wp.tile_load_indexed(M_in[worldid], idx, shape=(block_area,), storage="shared")
return cholesky_factorize
L = wp.tile_reshape(block, (block_size, block_size))
wp.tile_cholesky_inplace(L, fill_mode="upper")
wp.tile_store(L_out[worldid], wp.tile_reshape(L, (block_area,)), offset=(qLD_block_adr[start],))
return kernel
def _factor_i_dense(m: Model, d: Data, M: wp.array, L: wp.array):
"""Dense Cholesky factorization of inertia-like matrix M, assumed spd."""
def _factor_block_dense(m: Model, d: Data, M: wp.array2d[float], L: wp.array2d[float]):
for tile in m.M_tiles:
wp.launch_tiled(
_tile_cholesky_factorize(tile),
_tile_cholesky_factorize_block(tile),
dim=(d.nworld, tile.adr.size),
inputs=[M, tile.adr],
inputs=[m.qLD_block_adr, M, tile.elemid, tile.adr],
outputs=[L],
block_dim=m.block_dim.cholesky_factorize,
)
@@ -1092,11 +1076,23 @@ def _factor_i_dense(m: Model, d: Data, M: wp.array, L: wp.array):
@event_scope
def factor_m(m: Model, d: Data):
"""Factorization of inertia-like matrix M, assumed spd."""
if m.is_sparse:
_factor_i_sparse(m, d, d.M, d.qLD, d.qLDiagInv)
else:
_factor_i_dense(m, d, d.M, d.qLD)
"""Factorization of inertia-like matrix M, assumed spd.
The factor is a per-block decision: dense blocks factor as a packed tile-Cholesky (M_tiles),
sparse blocks via the LDL factor over the LDL region (offset qLD_block_total), and simple
(diagonal) blocks need only D = 1/diag. The passes write disjoint dofs and may all run at once.
"""
if m.qLD_has_dense:
_factor_block_dense(m, d, d.M, d.qLD)
if m.qLD_has_sparse:
_factor_i_sparse(m, d, d.M, d.qLD[:, m.qLD_block_total :], d.qLDiagInv)
if m.qLD_has_simple:
wp.launch(
_factor_simple,
dim=(d.nworld, m.qLD_simple_dofs.size),
inputs=[m.M_rownnz, m.M_rowadr, d.M, m.qLD_simple_dofs],
outputs=[d.qLDiagInv],
)
@wp.kernel
@@ -2738,7 +2734,9 @@ def _solve_LD_sparse_fused(nv: int, nlevels: int):
@wp.kernel(module="unique", enable_backward=False)
def kernel(
# In:
L: wp.array3d[float],
dof_dense: wp.array[int],
dof_simple: wp.array[int],
L: wp.array2d[float],
D: wp.array2d[float],
all_updates: wp.array[wp.vec3i],
level_offsets: wp.array[int],
@@ -2751,12 +2749,14 @@ def _solve_LD_sparse_fused(nv: int, nlevels: int):
NLEVELS = wp.static(nlevels)
BLOCK_DIM = wp.block_dim()
# Copy y to x_out
# Copy y to x_out for sparse-block dofs only; dense blocks use the packed pass and simple
# (diagonal) blocks use the dedicated 1/diag solve.
for dofid in range(tid, NV, BLOCK_DIM):
x_out[worldid, dofid] = y[worldid, dofid]
if dof_dense[dofid] == 0 and dof_simple[dofid] == 0:
x_out[worldid, dofid] = y[worldid, dofid]
_syncthreads()
# Forward substitution
# Forward substitution (all_updates only references sparse-block dofs)
for level in range(NLEVELS):
level_idx = NLEVELS - 1 - level
level_offset = level_offsets[level_idx]
@@ -2765,12 +2765,13 @@ def _solve_LD_sparse_fused(nv: int, nlevels: int):
for u in range(tid, level_size, BLOCK_DIM):
update = all_updates[level_offset + u]
i, k, Madr_ki = update[0], update[1], update[2]
wp.atomic_sub(x_out[worldid], i, L[worldid, 0, Madr_ki] * x_out[worldid, k])
wp.atomic_sub(x_out[worldid], i, L[worldid, Madr_ki] * x_out[worldid, k])
_syncthreads()
# Diagonal multiply
# Diagonal multiply (sparse-block dofs only)
for dofid in range(tid, NV, BLOCK_DIM):
x_out[worldid, dofid] *= D[worldid, dofid]
if dof_dense[dofid] == 0 and dof_simple[dofid] == 0:
x_out[worldid, dofid] *= D[worldid, dofid]
_syncthreads()
# Backward substitution
@@ -2782,7 +2783,7 @@ def _solve_LD_sparse_fused(nv: int, nlevels: int):
for u in range(tid, level_size, BLOCK_DIM):
update = all_updates[level_offset + u]
i, k, Madr_ki = update[0], update[1], update[2]
wp.atomic_sub(x_out[worldid], k, L[worldid, 0, Madr_ki] * x_out[worldid, i])
wp.atomic_sub(x_out[worldid], k, L[worldid, Madr_ki] * x_out[worldid, i])
_syncthreads()
return kernel
@@ -2791,7 +2792,7 @@ def _solve_LD_sparse_fused(nv: int, nlevels: int):
def _solve_LD_sparse(
m: Model,
d: Data,
L: wp.array3d[float],
L: wp.array2d[float],
D: wp.array2d[float],
x: wp.array2d[float],
y: wp.array2d[float],
@@ -2807,75 +2808,123 @@ def _solve_LD_sparse(
wp.launch(
_solve_LD_sparse_fused(m.nv, nlevels),
dim=(d.nworld, dim_block),
inputs=[L, D, m.qLD_all_updates, m.qLD_level_offsets, y],
inputs=[m.qLD_dof_dense, m.qLD_dof_simple, L, D, m.qLD_all_updates, m.qLD_level_offsets, y],
outputs=[x],
block_dim=dim_block,
)
@wp.kernel
def _solve_simple(
# In:
simple_dofs: wp.array[int],
D: wp.array2d[float],
y: wp.array2d[float],
# Out:
x_out: wp.array2d[float],
):
# A simple (decoupled) dof's solve is just x = (1/diag) * y.
worldid, s = wp.tid()
dofid = simple_dofs[s]
x_out[worldid, dofid] = D[worldid, dofid] * y[worldid, dofid]
@wp.kernel
def _factor_solve_simple(
# Model:
M_rownnz: wp.array[int],
M_rowadr: wp.array[int],
# Data in:
M_in: wp.array2d[float],
# In:
simple_dofs: wp.array[int],
y: wp.array2d[float],
# Out:
D_out: wp.array2d[float],
x_out: wp.array2d[float],
):
# Fused factor+solve for a simple dof: read M(i,i) once, emit D = 1/diag and x = D * y.
worldid, s = wp.tid()
dofid = simple_dofs[s]
diag_i = M_rowadr[dofid] + M_rownnz[dofid] - 1
d_inv = 1.0 / M_in[worldid, diag_i]
D_out[worldid, dofid] = d_inv
x_out[worldid, dofid] = d_inv * y[worldid, dofid]
@cache_kernel
def _tile_cholesky_solve(tile: TileSet):
"""Returns a kernel for dense Cholesky backsubstitution of a tile."""
def _tile_cholesky_solve_block(tile: TileSet):
# One diagonal block per (world, block) thread group; no densify, so a 2D grid suffices.
block_size = tile.size
block_area = block_size * block_size
@wp.kernel(module="unique", enable_backward=False)
def cholesky_solve(
def kernel(
# Model:
qLD_block_adr: wp.array[int],
# In:
L: wp.array3d[float],
block_dof: wp.array[int],
L_in: wp.array2d[float],
y: wp.array2d[float],
adr: wp.array[int],
# Out:
x: wp.array2d[float],
):
worldid, nodeid = wp.tid()
TILE_SIZE = wp.static(tile.size)
worldid, blk = wp.tid()
start = block_dof[blk]
dofid = adr[nodeid]
y_slice = wp.tile_load(y[worldid], shape=(TILE_SIZE,), offset=(dofid,))
L_tile = wp.tile_load(L[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid))
x_slice = wp.tile_cholesky_solve(L_tile, y_slice, fill_mode="upper")
wp.tile_store(x[worldid], x_slice, offset=(dofid,))
L = wp.tile_reshape(
wp.tile_load(L_in[worldid], shape=(block_area,), offset=(qLD_block_adr[start],)), (block_size, block_size)
)
rhs = wp.tile_load(y[worldid], shape=(block_size,), offset=(start,))
sol = wp.tile_cholesky_solve(L, rhs, fill_mode="upper")
wp.tile_store(x[worldid], sol, offset=(start,))
return cholesky_solve
return kernel
def _solve_LD_dense(m: Model, d: Data, L: wp.array3d[float], x: wp.array2d[float], y: wp.array2d[float]):
"""Computes dense backsubstitution: x = inv(U.T @ U) * y."""
def _solve_block_dense(m: Model, d: Data, L: wp.array2d[float], x: wp.array2d[float], y: wp.array2d[float]):
for tile in m.M_tiles:
# The triangular back-substitution is largely sequential, so large blocks prefer fewer threads
# for better occupancy while moderate blocks still want a couple warps (16/27->64, 60->32).
block_dim = m.block_dim.cholesky_solve if tile.size <= 40 else 32
wp.launch_tiled(
_tile_cholesky_solve(tile),
_tile_cholesky_solve_block(tile),
dim=(d.nworld, tile.adr.size),
inputs=[L, y, tile.adr],
inputs=[m.qLD_block_adr, tile.adr, L, y],
outputs=[x],
block_dim=m.block_dim.cholesky_solve,
block_dim=block_dim,
)
def solve_LD(
m: Model,
d: Data,
L: wp.array3d[float],
L: wp.array2d[float],
D: wp.array2d[float],
x: wp.array2d[float],
y: wp.array2d[float],
):
"""Computes backsubstitution for the inertia factorization.
Sparse models use MuJoCo's L'*D*L factors; dense models use an upper Cholesky factor U.
This function dispatches to either a sparse or dense solver depending on Model options.
The choice is per-block. Dense blocks back-substitute from the packed Cholesky region of L; sparse
blocks from the LDL region (offset qLD_block_total); simple (diagonal) blocks are a plain x = D*y.
The passes write disjoint dofs; the sparse pass skips dense and simple dofs so it does not clobber
their results.
Args:
m: The model containing factorization and sparsity information.
d: The data object containing workspace and factorization results.
L: Lower-triangular factor from the factorization (sparse or dense).
D: Diagonal factor from the factorization (only used for sparse).
L: The factor: packed dense region followed by the nC LDL region.
D: Diagonal factor (1/diag) for the sparse LDL and simple regions.
x: Output array for the solution.
y: Input right-hand side array.
"""
if m.is_sparse:
_solve_LD_sparse(m, d, L, D, x, y)
else:
_solve_LD_dense(m, d, L, x, y)
if m.qLD_has_dense:
_solve_block_dense(m, d, L, x, y)
if m.qLD_has_sparse:
_solve_LD_sparse(m, d, L[:, m.qLD_block_total :], D, x, y)
if m.qLD_has_simple:
wp.launch(_solve_simple, dim=(d.nworld, m.qLD_simple_dofs.size), inputs=[m.qLD_simple_dofs, D, y], outputs=[x])
@event_scope
@@ -2892,48 +2941,52 @@ def solve_m(m: Model, d: Data, x: wp.array2d[float], y: wp.array2d[float]):
@cache_kernel
def _tile_cholesky_factorize_solve(tile: TileSet):
"""Returns a kernel for dense Cholesky factorization and backsubstitution of a tile."""
def _tile_cholesky_factorize_solve_block(tile: TileSet):
# Fused factor+solve: densify the block, factor it, and back-substitute in one launch (avoids
# re-loading the factor). Grid/densify structure matches _tile_cholesky_factorize_block.
block_size = tile.size
block_area = block_size * block_size
@wp.kernel(module="unique", enable_backward=False)
def cholesky_factorize_solve(
def kernel(
# Model:
qLD_block_adr: wp.array[int],
# Data in:
M_in: wp.array3d[float],
M_in: wp.array2d[float],
# In:
block_elemid: wp.array[int],
block_dof: wp.array[int],
y: wp.array2d[float],
adr: wp.array[int],
# Out:
x: wp.array2d[float],
L: wp.array3d[float],
# Out:
L_out: wp.array2d[float],
):
worldid, nodeid = wp.tid()
TILE_SIZE = wp.static(tile.size)
worldid, blk = wp.tid()
start = block_dof[blk]
dofid = adr[nodeid]
M_tile = wp.tile_load(M_in[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid))
y_slice = wp.tile_load(y[worldid], shape=(TILE_SIZE,), offset=(dofid,))
# Densify the block (see _tile_cholesky_factorize_block for the gather rationale).
idx = wp.tile_load(block_elemid, shape=(block_area,), offset=(blk * block_area,), storage="shared")
block = wp.tile_load_indexed(M_in[worldid], idx, shape=(block_area,), storage="shared")
L_tile = wp.tile_cholesky(M_tile, fill_mode="upper")
wp.tile_store(L[worldid], L_tile, offset=(dofid, dofid))
x_slice = wp.tile_cholesky_solve(L_tile, y_slice, fill_mode="upper")
wp.tile_store(x[worldid], x_slice, offset=(dofid,))
L = wp.tile_reshape(block, (block_size, block_size))
wp.tile_cholesky_inplace(L, fill_mode="upper")
wp.tile_store(L_out[worldid], wp.tile_reshape(L, (block_area,)), offset=(qLD_block_adr[start],))
return cholesky_factorize_solve
rhs = wp.tile_load(y[worldid], shape=(block_size,), offset=(start,))
sol = wp.tile_cholesky_solve(L, rhs, fill_mode="upper")
wp.tile_store(x[worldid], sol, offset=(start,))
return kernel
def _factor_solve_i_dense(
m: Model,
d: Data,
M: wp.array3d[float],
x: wp.array2d[float],
y: wp.array2d[float],
L: wp.array3d[float],
def _factor_solve_block_dense(
m: Model, d: Data, M: wp.array2d[float], x: wp.array2d[float], y: wp.array2d[float], L: wp.array2d[float]
):
for tile in m.M_tiles:
wp.launch_tiled(
_tile_cholesky_factorize_solve(tile),
_tile_cholesky_factorize_solve_block(tile),
dim=(d.nworld, tile.adr.size),
inputs=[M, y, tile.adr],
inputs=[m.qLD_block_adr, M, tile.elemid, tile.adr, y],
outputs=[x, L],
block_dim=m.block_dim.cholesky_factorize_solve,
)
@@ -2942,25 +2995,33 @@ def _factor_solve_i_dense(
def factor_solve_i(m, d, M, L, D, x, y):
"""Factorizes and solves the inertia-like linear system.
Sparse models use MuJoCo's L'*D*L factors; dense models use an upper Cholesky factor U.
This function first factorizes the matrix M (sparse or dense depending on model options),
then solves the system for x given right-hand side y.
The choice is per-block (see factor_m): dense blocks factor+solve via the packed Cholesky, sparse
blocks via the LDL region, simple (diagonal) blocks via D = 1/diag. Factorizes M, solves for x.
Args:
m: The model containing factorization and sparsity information.
d: The data object containing workspace and factorization results.
M: The inertia-like matrix to factorize.
L: Output sparse factor or dense upper Cholesky factor.
D: Output diagonal factor from the factorization (only used for sparse).
M: The inertia-like matrix to factorize (CSR, length nC).
L: Output factor: packed dense region followed by the nC LDL region (sized like d.qLD).
D: Output diagonal factor (1/diag) for the sparse LDL and simple regions.
x: Output array for the solution.
y: Input right-hand side array.
"""
if m.is_sparse:
_factor_i_sparse(m, d, M, L, D)
_solve_LD_sparse(m, d, L, D, x, y)
else:
_factor_solve_i_dense(m, d, M, x, y, L)
# Per-block: dense blocks factor+solve via the packed Cholesky; sparse blocks via the LDL region
# (offset qLD_block_total); simple blocks via 1/diag. The passes write disjoint dofs.
if m.qLD_has_dense:
_factor_solve_block_dense(m, d, M, x, y, L)
if m.qLD_has_sparse:
L_ldl = L[:, m.qLD_block_total :]
_factor_i_sparse(m, d, M, L_ldl, D)
_solve_LD_sparse(m, d, L_ldl, D, x, y)
if m.qLD_has_simple:
wp.launch(
_factor_solve_simple,
dim=(d.nworld, m.qLD_simple_dofs.size),
inputs=[m.M_rownnz, m.M_rowadr, M, m.qLD_simple_dofs, y],
outputs=[D, x],
)
@cache_kernel
@@ -2978,7 +3039,7 @@ def _factor_solve_lu_sparse_fused(nv: int):
qfrc: wp.array2d[float],
# Data out:
qacc_out: wp.array2d[float],
qLU_out: wp.array3d[float],
qLU_out: wp.array2d[float],
):
worldid = wp.tid()
NV = wp.static(nv)
@@ -2997,7 +3058,7 @@ def _factor_solve_lu_sparse_fused(nv: int):
qacc_out[worldid, i] = float(rem_i - 1)
# cache diagonal element for row i
LUii = qLU_out[worldid, 0, ii]
LUii = qLU_out[worldid, ii]
# rows j above i (j < i), processed from i-1 down to 0
for c in range(i):
@@ -3015,8 +3076,8 @@ def _factor_solve_lu_sparse_fused(nv: int):
qacc_out[worldid, j] = float(rem_j)
# (j,i) = (j,i) / (i,i)
LUji = qLU_out[worldid, 0, ji] / LUii
qLU_out[worldid, 0, ji] = LUji
LUji = qLU_out[worldid, ji] / LUii
qLU_out[worldid, ji] = LUji
# (j,k) = (j,k) - (i,k) * (j,i) for k < i
icnt = rowadr_i
@@ -3026,7 +3087,7 @@ def _factor_solve_lu_sparse_fused(nv: int):
col_i = D_colind[icnt]
col_j = D_colind[jcnt]
if col_i == col_j:
qLU_out[worldid, 0, jcnt] = qLU_out[worldid, 0, jcnt] - qLU_out[worldid, 0, icnt] * LUji
qLU_out[worldid, jcnt] = qLU_out[worldid, jcnt] - qLU_out[worldid, icnt] * LUji
icnt = icnt + 1
jcnt = jcnt + 1
elif col_i > col_j:
@@ -3049,7 +3110,7 @@ def _factor_solve_lu_sparse_fused(nv: int):
for j in range(nnz_upper):
adr_j = rowadr_i + d1 + j
col = D_colind[adr_j]
acc = acc - qLU_out[worldid, 0, adr_j] * qacc_out[worldid, col]
acc = acc - qLU_out[worldid, adr_j] * qacc_out[worldid, col]
qacc_out[worldid, i] = acc
# Forward substitution: solve L*qacc = qacc
@@ -3061,15 +3122,15 @@ def _factor_solve_lu_sparse_fused(nv: int):
for j in range(diag_i):
adr_j = rowadr_i + j
col = D_colind[adr_j]
acc = acc - qLU_out[worldid, 0, adr_j] * qacc_out[worldid, col]
acc = acc - qLU_out[worldid, adr_j] * qacc_out[worldid, col]
qacc_out[worldid, i] = acc / qLU_out[worldid, 0, rowadr_i + diag_i]
qacc_out[worldid, i] = acc / qLU_out[worldid, rowadr_i + diag_i]
return kernel
@event_scope
def factor_solve_lu(m: Model, d: Data, qLU: wp.array3d[float], qacc: wp.array2d[float], qfrc: wp.array2d[float]):
def factor_solve_lu(m: Model, d: Data, qLU: wp.array2d[float], qacc: wp.array2d[float], qfrc: wp.array2d[float]):
r"""Factorize and solve non-symmetric implicit system: qacc = A \\ qfrc.
qLU is overwritten in-place with the LU factors, then used to solve for qacc.
File diff suppressed because it is too large Load Diff
+28 -347
View File
@@ -65,15 +65,15 @@ def next_act(
@cache_kernel
def mul_m_sparse(check_skip: bool):
def mul_m_kernel(check_skip: bool):
@wp.kernel(module="unique")
def _mul_m_sparse(
def _mul_m(
# Model:
M_mulm_rowadr: wp.array[int],
M_mulm_col: wp.array[int],
M_mulm_madr: wp.array[int],
# Data in:
M_in: wp.array3d[float],
M_in: wp.array2d[float],
# In:
vec: wp.array2d[float],
skip: wp.array[bool],
@@ -87,34 +87,33 @@ def mul_m_sparse(check_skip: bool):
if skip[worldid]:
return
# Gather all contributions (diagonal + off-diagonal)
# Gather all contributions (diagonal + off-diagonal).
acc = float(0.0)
start = M_mulm_rowadr[dofid]
end = M_mulm_rowadr[dofid + 1]
for k in range(start, end):
col = M_mulm_col[k]
madr = M_mulm_madr[k]
acc += M_in[worldid, 0, madr] * vec[worldid, col]
acc += M_in[worldid, madr] * vec[worldid, col]
res[worldid, dofid] = acc
return _mul_m_sparse
return _mul_m
@cache_kernel
def mul_m_dense(nv: int, check_skip: bool):
"""Simple SIMT dense matmul: one thread per output element."""
@wp.kernel(module="unique")
def _mul_m_dense(
# Data in:
M_in: wp.array3d[float],
M_in: wp.array3d[float], # kernel_analyzer: ignore
# In:
vec: wp.array2d[float],
skip: wp.array[bool],
# Out:
res: wp.array2d[float],
):
"""Dense matmul for the compact active-DOF inertia block (nworld, nv, nv)."""
worldid, i = wp.tid()
if wp.static(check_skip):
@@ -154,357 +153,23 @@ def mul_m(
if M is None:
M = d.M
if m.is_sparse:
wp.launch(
mul_m_sparse(check_skip),
dim=(d.nworld, m.nv),
inputs=[m.M_mulm_rowadr, m.M_mulm_col, m.M_mulm_madr, M, vec, skip],
outputs=[res],
)
else:
if M.ndim == 3:
# Dense compact active-DOF block (nworld, nv, nv) used by the compact solver.
wp.launch(
mul_m_dense(m.nv, check_skip),
dim=(d.nworld, m.nv),
inputs=[M, vec, skip],
outputs=[res],
)
@wp.kernel
def _mul_m_island_sparse(
# Model:
M_mulm_rowadr: wp.array[int],
M_mulm_col: wp.array[int],
M_mulm_madr: wp.array[int],
# Data in:
nidof_in: wp.array[int],
M_in: wp.array3d[float],
map_dof2idof_in: wp.array2d[int],
map_idof2dof_in: wp.array2d[int],
# In:
idof_islandid_in: wp.array2d[int],
vec: wp.array2d[float],
island_done_in: wp.array2d[bool],
check_skip: int,
# Out:
res: wp.array2d[float],
):
"""Sparse island mul_m for ALL islands in parallel."""
worldid, idofid = wp.tid()
nidof = nidof_in[worldid]
if idofid >= nidof:
return
islandid = idof_islandid_in[worldid, idofid]
if islandid < 0:
return
if check_skip:
if island_done_in[worldid, islandid]:
return
dof = map_idof2dof_in[worldid, idofid]
acc = float(0.0)
start = M_mulm_rowadr[dof]
end = M_mulm_rowadr[dof + 1]
for k in range(start, end):
col = M_mulm_col[k]
madr = M_mulm_madr[k]
col_idof = map_dof2idof_in[worldid, col]
# skip unconstrained DOFs
if col_idof < nidof:
acc += M_in[worldid, 0, madr] * vec[worldid, col_idof]
res[worldid, idofid] = acc
@wp.kernel
def _mul_m_island_dense(
# Model:
nv: int,
# Data in:
nidof_in: wp.array[int],
M_in: wp.array3d[float],
map_dof2idof_in: wp.array2d[int],
map_idof2dof_in: wp.array2d[int],
# In:
idof_islandid_in: wp.array2d[int],
vec: wp.array2d[float],
island_done_in: wp.array2d[bool],
check_skip: int,
# Out:
res: wp.array2d[float],
):
"""Dense island mul_m for ALL islands in parallel."""
worldid, idofid = wp.tid()
nidof = nidof_in[worldid]
if idofid >= nidof:
return
islandid = idof_islandid_in[worldid, idofid]
if islandid < 0:
return
if check_skip:
if island_done_in[worldid, islandid]:
return
dof = map_idof2dof_in[worldid, idofid]
acc = float(0.0)
for j in range(nv):
col_idof = map_dof2idof_in[worldid, j]
# skip unconstrained DOFs
if col_idof < nidof:
acc += M_in[worldid, dof, j] * vec[worldid, col_idof]
res[worldid, idofid] = acc
@event_scope
def mul_m_island(
m: Model,
d: Data,
res: wp.array2d[float],
vec: wp.array2d[float],
nidof: wp.array[int],
map_idof2dof: wp.array2d[int],
map_dof2idof: wp.array2d[int],
idof_islandid: wp.array2d[int],
island_done: Optional[wp.array] = None,
M: Optional[wp.array] = None,
):
"""Multiply island-local vectors by inertia matrix for all islands in parallel.
Args:
m: The model containing kinematic and dynamic information.
d: The data object containing the current state and output arrays.
res: Result: qM @ vec (island-local DOF order).
vec: Input vector (island-local DOF order).
nidof: Number of island DOFs per world.
map_idof2dof: Island-local DOF → global DOF map.
map_dof2idof: Global DOF → island-local DOF map.
idof_islandid: Island ID per island-local DOF.
island_done: Per-island done flags (nworld, ntree).
M: Optional mass matrix override.
"""
check_skip = int(island_done is not None)
island_done = island_done or wp.empty((0, 0), dtype=bool)
if M is None:
M = d.M
if m.is_sparse:
wp.launch(
_mul_m_island_sparse,
dim=(d.nworld, m.nv),
inputs=[
m.M_mulm_rowadr,
m.M_mulm_col,
m.M_mulm_madr,
nidof,
M,
map_dof2idof,
map_idof2dof,
idof_islandid,
vec,
island_done,
check_skip,
],
outputs=[res],
)
else:
wp.launch(
_mul_m_island_dense,
mul_m_kernel(check_skip),
dim=(d.nworld, m.nv),
inputs=[
m.nv,
nidof,
M,
map_dof2idof,
map_idof2dof,
idof_islandid,
vec,
island_done,
check_skip,
],
inputs=[m.M_mulm_rowadr, m.M_mulm_col, m.M_mulm_madr, M, vec, skip],
outputs=[res],
)
@cache_kernel
def _solve_LD_sparse_island(nv: int, nlevels: int):
"""Sparse backsubstitution with island-local index remapping.
Same algorithm as _solve_LD_sparse_fused, but reads/writes x/y in
island-local DOF order. The L/D factorization stays in global DOF order.
"""
@wp.func_native(snippet="WP_TILE_SYNC();")
def _syncthreads():
pass
@wp.kernel(module="unique", enable_backward=False)
def kernel(
# Data in:
nidof_in: wp.array[int],
map_dof2idof_in: wp.array2d[int],
map_idof2dof_in: wp.array2d[int],
# In:
L: wp.array3d[float],
D: wp.array2d[float],
all_updates: wp.array[wp.vec3i],
level_offsets: wp.array[int],
y: wp.array2d[float],
# Out:
x_out: wp.array2d[float],
):
worldid, tid = wp.tid()
NLEVELS = wp.static(nlevels)
BLOCK_DIM = wp.block_dim()
nid = nidof_in[worldid]
# Copy y to x_out (island-local, only iterate up to nid)
for idof in range(tid, nid, BLOCK_DIM):
x_out[worldid, idof] = y[worldid, idof]
_syncthreads()
# Forward substitution
for level in range(NLEVELS):
level_idx = NLEVELS - 1 - level
level_offset = level_offsets[level_idx]
level_size = level_offsets[level_idx + 1] - level_offset
for u in range(tid, level_size, BLOCK_DIM):
update = all_updates[level_offset + u]
i, k, Madr_ki = update[0], update[1], update[2]
idof_i = map_dof2idof_in[worldid, i]
if idof_i < nid:
idof_k = map_dof2idof_in[worldid, k]
wp.atomic_sub(x_out[worldid], idof_i, L[worldid, 0, Madr_ki] * x_out[worldid, idof_k])
_syncthreads()
# Diagonal multiply (only iterate up to nid)
for idof in range(tid, nid, BLOCK_DIM):
dofid = map_idof2dof_in[worldid, idof]
x_out[worldid, idof] *= D[worldid, dofid]
_syncthreads()
# Backward substitution
for level in range(NLEVELS):
level_idx = level
level_offset = level_offsets[level_idx]
level_size = level_offsets[level_idx + 1] - level_offset
for u in range(tid, level_size, BLOCK_DIM):
update = all_updates[level_offset + u]
i, k, Madr_ki = update[0], update[1], update[2]
idof_k = map_dof2idof_in[worldid, k]
if idof_k < nid:
idof_i = map_dof2idof_in[worldid, i]
wp.atomic_sub(x_out[worldid], idof_k, L[worldid, 0, Madr_ki] * x_out[worldid, idof_i])
_syncthreads()
return kernel
@cache_kernel
def _tile_cholesky_solve_island(tile):
"""Dense Cholesky backsubstitution with island-local index remapping.
L is loaded from global DOF offsets (factorization unchanged).
y/x are loaded/stored at island-local DOF offsets via map_dof2idof.
"""
@wp.kernel(module="unique", enable_backward=False)
def cholesky_solve(
# Data in:
nidof_in: wp.array[int],
map_dof2idof_in: wp.array2d[int],
# In:
L: wp.array3d[float],
y: wp.array2d[float],
adr: wp.array[int],
# Out:
x: wp.array2d[float],
):
worldid, nodeid = wp.tid()
TILE_SIZE = wp.static(tile.size)
dofid = adr[nodeid]
idofid = map_dof2idof_in[worldid, dofid]
# Skip unconstrained trees (uniform branch — all threads in block agree)
if idofid >= nidof_in[worldid]:
return
# L stays in global order
L_tile = wp.tile_load(L[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid))
# y and x use island-local offsets
y_slice = wp.tile_load(y[worldid], shape=(TILE_SIZE,), offset=(idofid,))
x_slice = wp.tile_cholesky_solve(L_tile, y_slice)
wp.tile_store(x[worldid], x_slice, offset=(idofid,))
return cholesky_solve
@event_scope
def solve_m_island(
m: Model,
d: Data,
res: wp.array2d[float],
vec: wp.array2d[float],
nidof: wp.array[int],
map_idof2dof: wp.array2d[int],
):
"""Compute res = M^{-1} @ vec for island-local DOFs.
Args:
m: Model.
d: Data.
res: Output in island-local DOF order.
vec: Input in island-local DOF order.
nidof: Number of island DOFs per world.
map_idof2dof: Island-local DOF -> global DOF map.
"""
if m.is_sparse:
nlevels = len(m.qLD_updates)
if wp.get_device().is_cuda:
dim_block = m.block_dim.solve_LD_sparse_fused
else:
dim_block = 1
wp.launch(
_solve_LD_sparse_island(m.nv, nlevels),
dim=(d.nworld, dim_block),
inputs=[
d.nidof,
d.map_dof2idof,
map_idof2dof,
d.qLD,
d.qLDiagInv,
m.qLD_all_updates,
m.qLD_level_offsets,
vec,
],
outputs=[res],
block_dim=dim_block,
)
else:
for tile in m.M_tiles:
wp.launch_tiled(
_tile_cholesky_solve_island(tile),
dim=(d.nworld, tile.adr.size),
inputs=[d.nidof, d.map_dof2idof, d.qLD, vec, tile.adr],
outputs=[res],
block_dim=m.block_dim.cholesky_solve,
)
@wp.kernel
def _apply_ft(
# Model:
@@ -938,6 +603,7 @@ def get_state(m: Model, d: Data, state: wp.array2d[float], sig: int, active: Opt
nbody: int,
neq: int,
nmocap: int,
nuserdata: int,
nhistory: int,
# Data in:
time_in: wp.array[float],
@@ -952,6 +618,7 @@ def get_state(m: Model, d: Data, state: wp.array2d[float], sig: int, active: Opt
eq_active_in: wp.array2d[bool],
mocap_pos_in: wp.array2d[wp.vec3],
mocap_quat_in: wp.array2d[wp.quat],
userdata_in: wp.array2d[float],
# In:
sig_in: int,
active_in: wp.array[bool],
@@ -1028,6 +695,10 @@ def get_state(m: Model, d: Data, state: wp.array2d[float], sig: int, active: Opt
state_out[worldid, adr + 2] = quat[2]
state_out[worldid, adr + 3] = quat[3]
adr += 4
elif element == State.USERDATA:
for j in range(nuserdata):
state_out[worldid, adr + j] = userdata_in[worldid, j]
adr += nuserdata
wp.launch(
_get_state,
@@ -1040,6 +711,7 @@ def get_state(m: Model, d: Data, state: wp.array2d[float], sig: int, active: Opt
m.nbody,
m.neq,
m.nmocap,
m.nuserdata,
m.nhistory,
d.time,
d.qpos,
@@ -1053,6 +725,7 @@ def get_state(m: Model, d: Data, state: wp.array2d[float], sig: int, active: Opt
d.eq_active,
d.mocap_pos,
d.mocap_quat,
d.userdata,
int(sig),
active or wp.ones(d.nworld, dtype=bool),
],
@@ -1085,6 +758,7 @@ def set_state(m: Model, d: Data, state: wp.array2d[float], sig: int, active: Opt
nbody: int,
neq: int,
nmocap: int,
nuserdata: int,
nhistory: int,
# In:
sig_in: int,
@@ -1103,6 +777,7 @@ def set_state(m: Model, d: Data, state: wp.array2d[float], sig: int, active: Opt
eq_active_out: wp.array2d[bool],
mocap_pos_out: wp.array2d[wp.vec3],
mocap_quat_out: wp.array2d[wp.quat],
userdata_out: wp.array2d[float],
):
worldid = wp.tid()
@@ -1180,6 +855,10 @@ def set_state(m: Model, d: Data, state: wp.array2d[float], sig: int, active: Opt
)
mocap_quat_out[worldid, j] = quat
adr += 4
elif element == State.USERDATA:
for j in range(nuserdata):
userdata_out[worldid, j] = state_in[worldid, adr + j]
adr += nuserdata
wp.launch(
_set_state,
@@ -1192,6 +871,7 @@ def set_state(m: Model, d: Data, state: wp.array2d[float], sig: int, active: Opt
m.nbody,
m.neq,
m.nmocap,
m.nuserdata,
m.nhistory,
int(sig),
active or wp.ones(d.nworld, dtype=bool),
@@ -1210,5 +890,6 @@ def set_state(m: Model, d: Data, state: wp.array2d[float], sig: int, active: Opt
d.eq_active,
d.mocap_pos,
d.mocap_quat,
d.userdata,
],
)
+181 -102
View File
@@ -20,8 +20,6 @@ import mujoco
import numpy as np
import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src.util_pkg import check_version
MJ_MINVAL = mujoco.mjMINVAL
MJ_MAXVAL = mujoco.mjMAXVAL
MJ_MINIMP = mujoco.mjMINIMP # minimum constraint impedance
@@ -29,8 +27,6 @@ MJ_MAXIMP = mujoco.mjMAXIMP # maximum constraint impedance
MJ_MAXCONPAIR = mujoco.mjMAXCONPAIR
MJ_MINMU = mujoco.mjMINMU # minimum friction
MJ_MINAWAKE = mujoco.mjMINAWAKE # minimum number of timesteps before sleeping
NEW_GAP_SEMANTICS = check_version("mujoco>=3.9.0.dev914519929")
TACTILE_DEPTH_SEMANTICS = check_version("mujoco>=3.9.0.dev921980899")
# maximum size (by number of edges) of an horizon in EPA algorithm
MJ_MAX_EPAHORIZON = 24
# maximum average number of trianglarfaces EPA can insert at each iteration
@@ -39,6 +35,9 @@ MJ_MAX_EPAFACES = 5
TILE_SIZE_JTDAJ_SPARSE = 16
TILE_SIZE_JTDAJ_DENSE = 16
# max M block size where dense tile-Cholesky beats sparse LDL (wins to ~64, degrades past ~80)
M_BLOCK_DENSE_MAX = 64
# maximum number of plugin attributes
_NPLUGINATTR = 128
@@ -52,14 +51,14 @@ class BlockDim:
Attributes:
segmented_sort: segmented sort block dimension (collision_driver)
euler_dense: Euler dense block dimension (forward)
convex_ccd: convex CCD kernel block dimension (collision_convex)
actuator_velocity: actuator velocity block dimension (forward)
ray: ray block dimension (ray)
contact_sort: contact sort block dimension (sensor)
energy_vel_kinetic: energy velocity kinetic block dimension (sensor)
cholesky_factorize: Cholesky factorize block dimension (smooth)
cholesky_factorize: block-dense Cholesky factorize block dimension (smooth)
cholesky_factorize_solve: block-dense Cholesky factorize+solve block dimension (smooth)
cholesky_solve: Cholesky solve block dimension (smooth)
cholesky_factorize_solve: Cholesky factorize and solve block dimension (smooth)
solve_LD_sparse_fused: solve LD sparse fused block dimension (smooth)
update_gradient_cholesky: update gradient Cholesky block dimension (solver)
update_gradient_cholesky_blocked: update gradient Cholesky blocked block dimension (solver)
@@ -77,23 +76,24 @@ class BlockDim:
# collision_driver
segmented_sort: int = 128
# collision_convex
convex_ccd: int = 256
# forward
euler_dense: int = 32
actuator_velocity: int = 32
# ray
ray: int = 64
# sensor
contact_sort: int = 64
energy_vel_kinetic: int = 32
# smooth
# smooth -- block tile-Cholesky widths
cholesky_factorize: int = 32
cholesky_solve: int = 32
cholesky_factorize_solve: int = 32
solve_LD_sparse_fused: int = 64
cholesky_solve: int = 64
solve_LD_sparse_fused: int = 128
# solver
update_gradient_cholesky: int = 64
update_gradient_cholesky_blocked: int = 32
update_gradient_JTDAJ_sparse: int = 64
update_gradient_JTDAJ_sparse: int = 128
update_gradient_JTDAJ_dense: int = 128
linesearch_iterative: int = 32
update_gradient_grad: int = 256
@@ -131,10 +131,34 @@ class BroadphaseFilter(enum.IntFlag):
OBB: collision between oriented bounding boxes
"""
PLANE = 1
SPHERE = 2
AABB = 4
OBB = 8
PLANE = 1 << 0
SPHERE = 1 << 1
AABB = 1 << 2
OBB = 1 << 3
class OverflowType(enum.IntFlag):
"""Bitmask for physics and collision overflows.
Attributes:
NEFC: nefc > njmax
NJMAX_NNZ: njmax_nnz overflow
BROADPHASE: broadphase overflow / flex broadphase overflow
NARROWPHASE: narrowphase overflow / flex narrowphase / contact overflow
CCD: CCD overflow / flex CCD overflow
HFIELD: height field collision overflow
CONTACT_MATCH: contact match sensor overflow
NVMAX: nvmax overflow (islands)
"""
NEFC = 1 << 0
NJMAX_NNZ = 1 << 1
BROADPHASE = 1 << 2
NARROWPHASE = 1 << 3
CCD = 1 << 4
HFIELD = 1 << 5
CONTACT_MATCH = 1 << 6
NVMAX = 1 << 7
class CamLightType(enum.IntEnum):
@@ -431,6 +455,8 @@ class GeomType(enum.IntEnum):
MESH = mujoco.mjtGeom.mjGEOM_MESH
SDF = mujoco.mjtGeom.mjGEOM_SDF
FLEX = mujoco.mjtGeom.mjGEOM_FLEX
# warp only
TRIANGLE = 999
# unsupported: NGEOMTYPES, ARROW*, LINE, SKIN, LABEL, NONE
@@ -708,7 +734,8 @@ class State(enum.IntEnum):
FULLPHYSICS = mujoco.mjtState.mjSTATE_FULLPHYSICS
USER = mujoco.mjtState.mjSTATE_USER
INTEGRATION = mujoco.mjtState.mjSTATE_INTEGRATION
# unsupported: USERDATA, PLUGIN
USERDATA = mujoco.mjtState.mjSTATE_USERDATA
# unsupported: PLUGIN
class vec5f(wp.types.vector(length=5, dtype=float)):
@@ -816,6 +843,7 @@ class Option:
zeros out the contacts at each step)
contact_sensor_maxmatch: max number of contacts considered by contact sensor matching criteria
contacts matched after this value is exceded will be ignored
warn_overflow: warn if overflow is encountered
"""
timestep: array("*", float)
@@ -845,6 +873,7 @@ class Option:
graph_conditional: bool
run_collision_detection: bool
contact_sensor_maxmatch: int
warn_overflow: bool
# TODO(team): remove in future version
@property
@@ -884,10 +913,12 @@ class TileSet:
Attributes:
adr: address of each tile in the set
size: size of all the tiles in this set
elemid: flat per-block gather indices into CSR M for tile_load_indexed (absent -> nC sentinel)
"""
adr: wp.array[int]
size: int
elemid: wp.array[int] = None
def __eq__(self, other) -> bool:
if self.__class__ is not other.__class__:
@@ -951,6 +982,7 @@ class Model:
nflexbending: number of bending parameters in all flexes
nflexelemedge: number of element edge ids in all flexes
nflexshelldata: number of shell fragment vertex ids in all flexes
nflexevpair: number of element-vertex pairs in all flexes
nJfe: number of non-zeros in sparse flexedge Jacobian
nmesh: number of meshes
nmeshvert: number of vertices for all meshes
@@ -973,7 +1005,7 @@ class Model:
nmocap: number of mocap bodies
nplugin: number of plugin instances
nJmom: number of non-zeros in actuator_moment
ngravcomp: number of bodies with nonzero gravcomp
nuserdata: number of custom user parameters
nsensordata: number of elements in sensor data vector
nhistory: number of history buffer entries
opt: physics options
@@ -1105,6 +1137,8 @@ class Model:
flex_friction: friction for (slide, spin, roll) (nflex, 3)
flex_margin: detect contact if dist<margin (nflex,)
flex_gap: include in solver if dist<margin-gap (nflex,)
flex_internal: internal collision enabled (nflex,)
flex_selfcollide: self-collision mode (nflex,)
flex_dim: 1: lines, 2: triangles, 3: tetrahedra (nflex,)
flex_vertadr: first vertex address (nflex,)
flex_vertnum: number of vertices (nflex,)
@@ -1118,12 +1152,15 @@ class Model:
flex_bendingadr: first bending data address (nflex,)
flex_shellnum: number of shells (nflex,)
flex_shelldataadr: first shell data address (nflex,)
flex_evpairadr: first element-vertex pair address (nflex,)
flex_evpairnum: number of element-vertex pairs (nflex,)
flex_vertbodyid: vertex body ids (nflexvert,)
flex_edge: edge vertex ids (2 per edge) (nflexedge, 2)
flex_edgeflap: adjacent vertex ids (dim=2 only) (nflexedge, 2)
flex_elem: element vertex ids (dim+1 per elem) (nflexelemdata,)
flex_elemedge: element edge ids (nflexelemedge,)
flex_shell: shell fragment vertex ids (dim per frag) (nflexshelldata,)
flex_evpair: element-vertex pair indices (nflexevpair, 2)
flex_vert: vertex local positions (nflexvert, 3)
flexedge_length0: edge lengths in qpos0 (nflexedge,)
flexedge_invweight0: inv. inertia for the edge (nflexedge,)
@@ -1146,6 +1183,7 @@ class Model:
mesh_normal: normals for all meshes (nmeshnormal, 3)
mesh_face: face indices for all meshes (nface, 3)
mesh_graph: convex graph data (nmeshgraph,)
mesh_pos: translation applied to asset vertices (nmesh, 3)
mesh_quat: rotation applied to asset vertices (nmesh, 4)
mesh_polynum: number of polygons per mesh (nmesh,)
mesh_polyadr: first polygon address per mesh (nmesh,)
@@ -1261,7 +1299,6 @@ class Model:
D_colind: column indices in D-structure (nD,)
mapM2D: index mapping from M to D (nD,)
mapD2M: index mapping from D to M (nC,)
flex_vertflexid: flex id for each flex vertex (nflexvert,)
warp only fields:
callback: custom physics callbacks
@@ -1277,15 +1314,28 @@ class Model:
nmaxpyramid: maximum number of pyramid directions
nmaxpolygon: maximum number of verts per polygon
nmaxmeshdeg: maximum number of polygons per vert
is_sparse: whether to use sparse representations
is_sparse: constraint Jacobian/Hessian layout (sparse vs dense). Does not affect M, whose
factorization is a per-block decision -- see qLD_* and m_block_layout
qLD_has_dense: any M block factors as a packed dense block
qLD_has_simple: any M block is simple (diagonal -> 1/diag, no factorization)
qLD_has_sparse: any M block factors via sparse LDL (oversized block / tendon armature)
qLD_block_total: packed length of the dense region per world (also the offset of the LDL region)
qLD_block_adr: packed offset of each dof's diagonal block; 0 if sparse (nv,)
qLD_dof_dense: per-dof flag, 1 if the dof's block is dense (packed) (nv,)
qLD_dof_simple: per-dof flag, 1 if the dof's block is simple (diagonal) (nv,)
qLD_simple_dofs: indices of the simple (diagonal) dofs (nsimple,)
has_fluid: True if wind, density, or viscosity are non-zero at put_model time
has_sdf_geom: whether the model contains SDF geoms
has_flex_selfcollide: whether any flex has self-collision enabled
max_flex_dim: maximum flex dimension in the model
block_dim: block dim options
body_tree: list of body ids by tree level
body_branches: flattened body ids for all branches
body_branch_start: start index in body_branches for each branch (nbranch + 1,)
mocap_bodyid: id of body for mocap (nmocap,)
body_fluid_ellipsoid: does body use ellipsoid fluid (nbody,)
body_fluid_ellipsoid_adr: body ids with ellipsoid fluid (nbody_fluid_ellipsoid,)
body_fluid_box_adr: body ids with box fluid (nbody_fluid_box,)
jnt_limited_slide_hinge_adr: limited/slide/hinge jntadr
jnt_limited_ball_adr: limited/ball jntadr
body_isdofancestor: precomputed mask of which DOFs affect each body
@@ -1353,6 +1403,7 @@ class Model:
M_fullm_i: sparse mass matrix addressing
M_fullm_j: sparse mass matrix addressing
M_elemid: (row, col) -> CSR madr addresses; -1 if not a chain ancestor
M_hinit_i: row index of each CSR M entry; for densifying M into the dense Newton H (nC,)
M_fullm_upper_i: upper-triangle row indices for solver h seeding
M_fullm_upper_j: upper-triangle column indices for solver h seeding
M_fullm_upper_elemid: source elemid into M_fullm_i/M_fullm_j
@@ -1361,6 +1412,14 @@ class Model:
M_mulm_rowadr: sparse matmul row pointers
M_mulm_col: sparse matmul column indices
M_mulm_madr: sparse matmul matrix addresses
flexelem_geom_pair_filtered: conaffinity-filtered element vs geom pairs (*, 2)
flexshell_geom_pair_filtered: conaffinity-filtered shell vs geom pairs (*, 2)
flexvert_geom_pair_filtered: conaffinity-filtered vertex vs geom pairs (*, 2)
flex_elemflexid: maps each element index directly to its flexid (nflexelem,)
flex_shellflexid: maps each shell index directly to its flexid (nflexshelldata,)
flex_evpairflexid: maps each element-vertex pair directly to its flexid (nflexevpair,)
flex_vertflexid: maps each vertex index directly to its flexid (nflexvert,)
flex_shelladr: maps each flex to its start shell index (nflex,)
"""
nq: int
@@ -1387,6 +1446,7 @@ class Model:
nflexbending: int
nflexelemedge: int
nflexshelldata: int
nflexevpair: int
nJfe: int
nmesh: int
nmeshvert: int
@@ -1409,7 +1469,7 @@ class Model:
nmocap: int
nplugin: int
nJmom: int
ngravcomp: int
nuserdata: int
nsensordata: int
nhistory: int
opt: Option
@@ -1541,6 +1601,8 @@ class Model:
flex_friction: array("nflex", wp.vec3)
flex_margin: array("nflex", float)
flex_gap: array("nflex", float)
flex_internal: array("nflex", int)
flex_selfcollide: array("nflex", int)
flex_dim: array("nflex", int)
flex_vertadr: array("nflex", int)
flex_vertnum: array("nflex", int)
@@ -1554,12 +1616,15 @@ class Model:
flex_bendingadr: array("nflex", int)
flex_shellnum: array("nflex", int)
flex_shelldataadr: array("nflex", int)
flex_evpairadr: array("nflex", int)
flex_evpairnum: array("nflex", int)
flex_vertbodyid: array("nflexvert", int)
flex_edge: array("nflexedge", wp.vec2i)
flex_edgeflap: array("nflexedge", wp.vec2i)
flex_elem: array("nflexelemdata", int)
flex_elemedge: array("nflexelemedge", int)
flex_shell: array("nflexshelldata", int)
flex_evpair: array("nflexevpair", wp.vec2i)
flex_vert: array("nflexvert", wp.vec3)
flexedge_length0: array("nflexedge", float)
flexedge_invweight0: array("nflexedge", float)
@@ -1582,6 +1647,7 @@ class Model:
mesh_normal: array("nmeshnormal", wp.vec3)
mesh_face: array("nmeshface", wp.vec3i)
mesh_graph: array("nmeshgraph", int)
mesh_pos: array("nmesh", wp.vec3)
mesh_quat: array("nmesh", wp.quat)
mesh_polynum: array("nmesh", int)
mesh_polyadr: array("nmesh", int)
@@ -1697,7 +1763,6 @@ class Model:
D_colind: array("nD", int)
mapM2D: array("nD", int)
mapD2M: array("nC", int)
flex_vertflexid: array("nflexvert", int)
# warp only fields:
callback: Callback
nbranch: int
@@ -1712,14 +1777,26 @@ class Model:
nmaxpolygon: int
nmaxmeshdeg: int
is_sparse: bool
qLD_has_dense: bool
qLD_has_simple: bool
qLD_has_sparse: bool
qLD_block_total: int
qLD_block_adr: wp.array[int]
qLD_dof_dense: wp.array[int]
qLD_dof_simple: wp.array[int]
qLD_simple_dofs: wp.array[int]
has_fluid: bool
has_sdf_geom: bool
has_flex_selfcollide: bool
max_flex_dim: int
block_dim: BlockDim
body_tree: tuple[wp.array[int], ...]
body_branches: wp.array[int]
body_branch_start: wp.array[int]
mocap_bodyid: array("nmocap", int)
body_fluid_ellipsoid: array("nbody", bool)
body_fluid_ellipsoid_adr: wp.array[int]
body_fluid_box_adr: wp.array[int]
jnt_limited_slide_hinge_adr: wp.array[int]
jnt_limited_ball_adr: wp.array[int]
body_isdofancestor: array("nbody", "nv_pad", int)
@@ -1780,6 +1857,7 @@ class Model:
M_fullm_i: wp.array[int]
M_fullm_j: wp.array[int]
M_elemid: wp.array2d[int] # (row, col) -> CSR madr address; -1 if col is not a chain ancestor of row
M_hinit_i: wp.array[int] # row index of each CSR M entry (for densifying M into the dense Newton H)
M_fullm_upper_i: wp.array[int]
M_fullm_upper_j: wp.array[int]
M_fullm_upper_elemid: wp.array[int]
@@ -1789,6 +1867,14 @@ class Model:
M_mulm_rowadr: wp.array[int] # start address for each row [nv+1]
M_mulm_col: wp.array[int] # column index to gather from
M_mulm_madr: wp.array[int] # matrix address to read
flexelem_geom_pair_filtered: wp.array[wp.vec2i]
flexshell_geom_pair_filtered: wp.array[wp.vec2i]
flexvert_geom_pair_filtered: wp.array[wp.vec2i]
flex_elemflexid: array("nflexelem", int)
flex_shellflexid: array("nflexshelldata", int)
flex_evpairflexid: array("nflexevpair", int)
flex_vertflexid: array("nflexvert", int)
flex_shelladr: array("nflex", int)
class ContactType(enum.IntFlag):
@@ -1799,8 +1885,8 @@ class ContactType(enum.IntFlag):
SENSOR: contact for collision sensor (GEOMDIST, GEOMNORMAL, GEOMFROMTO)
"""
CONSTRAINT = 1
SENSOR = 2
CONSTRAINT = 1 << 0
SENSOR = 1 << 1
@dataclasses.dataclass
@@ -1819,6 +1905,7 @@ class Contact:
dim: contact space dimensionality: 1, 3, 4 or 6 (naconmax,)
geom: geom ids; -1 for flex (naconmax, 2)
flex: flex ids; -1 for geom (naconmax, 2)
elem: element ids; -1 for geom or flex vertex (naconmax, 2)
vert: vertex ids for flex/mesh contact (naconmax, 2)
efc_address: address in efc; -1: not included (naconmax, nmaxpyramid)
worldid: world id (naconmax,)
@@ -1839,6 +1926,7 @@ class Contact:
dim: array("naconmax", int)
geom: array("naconmax", wp.vec2i)
flex: array("naconmax", wp.vec2i)
elem: array("naconmax", wp.vec2i)
vert: array("naconmax", wp.vec2i)
efc_address: array("naconmax", "nmaxpyramid", int)
worldid: array("naconmax", int)
@@ -1853,6 +1941,9 @@ class Constraint:
Attributes:
type: constraint type (ConstraintType) (nworld, njmax)
id: id of object of specific type (nworld, njmax)
jtdaj_adr: first efc row of each JTDAJ block (nworld, njmax)
jtdaj_nrow: efc rows per JTDAJ block (nworld, njmax)
jtdaj_nblock: number of JTDAJ blocks (nworld,)
J_rownnz: number of non-zeros in J row (nworld, 0) dense
(nworld, njmax) sparse
J_rowadr: row start address in colind array (nworld, 0) dense
@@ -1870,19 +1961,6 @@ class Constraint:
force: constraint force in constraint space (nworld, njmax)
state: constraint state (nworld, njmax_pad)
island: island ID per constraint (nworld, njmax)
itype: island constraint type (nworld, njmax)
iid: island constraint id (nworld, njmax)
iJ_rownnz: island J_rownnz (nworld, njmax)
iJ_rowadr: island J_rowadr (nworld, njmax)
iJ_colind: island J_colind (nworld, 0, 0) dense
(nworld, 1, njmax_nnz) sparse
iJ: island J (nworld, njmax, nv) dense
(nworld, 1, njmax_nnz) sparse
iD: island constraint mass (nworld, njmax_pad)
iaref: island aref (nworld, njmax)
ifrictionloss: island frictionloss (nworld, njmax)
iforce: island force (nworld, njmax)
istate: island state (nworld, njmax_pad)
warp only fields:
Ma: M*qacc (nworld, nv)
Jqvel: J*qvel (nworld, njmax)
@@ -1890,6 +1968,9 @@ class Constraint:
type: array("nworld", "njmax", int)
id: array("nworld", "njmax", int)
jtdaj_adr: array("nworld", "njmax", int)
jtdaj_nrow: array("nworld", "njmax", int)
jtdaj_nblock: array("nworld", int)
J_rownnz: array("nworld", "njmax", int)
J_rowadr: array("nworld", "njmax", int)
J_colind: wp.array3d[int]
@@ -1906,18 +1987,6 @@ class Constraint:
Ma: array("nworld", "nv", float)
Jqvel: array("nworld", "njmax", float)
itype: array("nworld", "njmax", int)
iid: array("nworld", "njmax", int)
iJ_rownnz: array("nworld", "njmax", int)
iJ_rowadr: array("nworld", "njmax", int)
iJ_colind: wp.array3d[int]
iJ: wp.array3d[float]
iD: array("nworld", "njmax_pad", float)
iaref: array("nworld", "njmax", float)
ifrictionloss: array("nworld", "njmax", float)
iforce: array("nworld", "njmax", float)
istate: array("nworld", "njmax_pad", int)
@dataclasses.dataclass
class Data:
@@ -1949,6 +2018,7 @@ class Data:
mocap_quat: orientation of mocap bodies (nworld, nmocap, 4)
qacc: acceleration (nworld, nv)
act_dot: time-derivative of actuator activation (nworld, na)
userdata: custom user data (nworld, nuserdata)
sensordata: sensor data array (nworld, nsensordata,)
tree_asleep: tree asleep counter; >=0: asleep cycle (nworld, ntree)
xpos: Cartesian position of body frame (nworld, nbody, 3)
@@ -1984,11 +2054,10 @@ class Data:
moment_colind: column indices in sparse actuator_moment (nworld, nJmom)
actuator_moment: actuator moments (nworld, nJmom)
crb: com-based composite inertia and mass (nworld, nbody, 10)
M: total inertia (nworld, nv_pad, nv_pad) if dense
(nworld, 1, nC) if sparse
qLD: upper Cholesky factorization (nworld, nv, nv) if dense
L'*D*L factorization of M (nworld, 1, nC) if sparse
qLDiagInv: 1/diag(D) (nworld, nv)
M: total inertia, CSR (nworld, nC)
qLD: per-block factor: packed dense region, then the nC (nworld, qLD_block_total + nC)
L'*D*L region at offset qLD_block_total (nC=0 if no sparse block)
qLDiagInv: 1/diag(D) for the sparse LDL region (nworld, nv)
tree_awake: is tree awake; 0: asleep; 1: awake (nworld, ntree)
body_awake: body sleep state (SleepState) (nworld, nbody)
body_awake_ind: indices of awake/static bodies (nworld, nbody)
@@ -2006,7 +2075,7 @@ class Data:
qfrc_passive: total passive force (nworld, nv)
subtree_linvel: linear velocity of subtree com (nworld, nbody, 3)
subtree_angmom: angular momentum about subtree com (nworld, nbody, 3)
qLU: sparse LU factorization of (M - dt*qDeriv) (nworld, 1, nD)
qLU: sparse LU factorization of (M - dt*qDeriv) (nworld, nD)
actuator_force: actuator force in actuation space (nworld, nu)
qfrc_actuator: actuator force (nworld, nv)
qfrc_smooth: net unconstrained force (nworld, nv)
@@ -2028,27 +2097,46 @@ class Data:
island_nefc: constraints per island (nworld, ntree)
island_ne: equality constraints per island (nworld, ntree)
island_nf: friction constraints per island (nworld, ntree)
island_efcadr: island start address in efc vector (nworld, ntree)
island_iefcadr: island start address in efc vector (nworld, ntree)
map_dof2idof: global DOF -> island-local DOF (nworld, nv)
map_idof2dof: island-local DOF -> global DOF (nworld, nv)
map_efc2iefc: global EFC -> island-local EFC (nworld, njmax)
map_iefc2efc: island-local EFC -> global EFC (nworld, njmax)
dof_islandid: island ID per island-DOF (nworld, nv)
efc_islandid: island ID per island-EFC (nworld, njmax)
iqacc: island-local qacc (nworld, nv)
iqacc_smooth: island-local qacc_smooth (nworld, nv)
iqfrc_smooth: island-local qfrc_smooth (nworld, nv)
iqfrc_constraint: island-local qfrc_constraint (nworld, nv)
ncdof: number of active (compacted) DOFs per world (nworld,)
dof_cdof: global DOF -> compacted DOF; -1 if inactive (nworld, nv)
cdof_dof: compacted DOF -> global DOF; -1 if unused (nworld, nvmax_pad)
ctol: compacted-solve main tolerance (nv/nvmax_pad scaled) (1,)
cls_tol: compacted-solve linesearch tolerance (1,)
cdof_tri_row: row index of compacted Hessian dof-pairs (nvmax_pad^2,)
cdof_tri_col: col index of compacted Hessian dof-pairs (nvmax_pad^2,)
cM: compacted dense inertia (nworld, nvmax_pad, nvmax_pad)
cqLD: compacted upper Cholesky factor (nworld, nvmax_pad, nvmax_pad)
crhs: compacted smooth-solve right-hand side (nworld, nvmax_pad, 1)
cx: compacted smooth-solve solution (nworld, nvmax_pad, 1)
cJ: compacted dense constraint Jacobian (nworld, njmax_pad, nvmax_pad)
cMa: compacted M @ qacc workspace (nworld, nvmax_pad)
cqfrc_smooth: compacted net unconstrained force (nworld, nvmax_pad)
cqacc_smooth: compacted unconstrained acceleration (nworld, nvmax_pad)
cqacc_warmstart: compacted warmstart acceleration (nworld, nvmax_pad)
cqacc: compacted acceleration (solve output) (nworld, nvmax_pad)
cqfrc_constraint: compacted constraint force (nworld, nvmax_pad)
warp only fields:
nworld: number of worlds
naconmax: maximum number of contacts (shared across all worlds)
naccdmax: maximum number of contacts for CCD (all worlds)
njmax: maximum number of constraints per world
nvmax: capacity for compacted active DOFs per world
nvmax_pad: nvmax rounded up to the nearest multiple of TILE_SIZE_JTDAJ_DENSE
njmax_pad: njmax rounded up to the nearest multiple of TILE_SIZE_JTDAJ
njmax_nnz: number of non-zeros in constraint Jacobian
nacon: number of detected contacts (across all worlds) (1,)
ncollision: collision count from broadphase (1,)
flex_aabb_min: dynamic flex object bounding box min (nworld, nflex, 3)
flex_aabb_max: dynamic flex object bounding box max (nworld, nflex, 3)
overflow: overflow bitmask (OverflowType) (nworld,)
"""
solver_niter: array("nworld", int)
@@ -2076,6 +2164,7 @@ class Data:
mocap_quat: array("nworld", "nmocap", wp.quat)
qacc: array("nworld", "nv", float)
act_dot: array("nworld", "na", float)
userdata: array("nworld", "nuserdata", float)
sensordata: array("nworld", "nsensordata", float)
tree_asleep: array("nworld", "ntree", int)
xpos: array("nworld", "nbody", wp.vec3)
@@ -2111,8 +2200,8 @@ class Data:
moment_colind: array("nworld", "nJmom", int)
actuator_moment: array("nworld", "nJmom", float)
crb: array("nworld", "nbody", vec10)
M: wp.array3d[float]
qLD: wp.array3d[float]
M: wp.array2d[float]
qLD: wp.array2d[float]
qLDiagInv: array("nworld", "nv", float)
tree_awake: array("nworld", "ntree", int)
body_awake: array("nworld", "nbody", int)
@@ -2131,7 +2220,7 @@ class Data:
qfrc_passive: array("nworld", "nv", float)
subtree_linvel: array("nworld", "nbody", wp.vec3)
subtree_angmom: array("nworld", "nbody", wp.vec3)
qLU: array("nworld", 1, "nD", float)
qLU: array("nworld", "nD", float)
actuator_force: array("nworld", "nu", float)
qfrc_actuator: array("nworld", "nv", float)
qfrc_smooth: array("nworld", "nv", float)
@@ -2151,27 +2240,46 @@ class Data:
island_nefc: array("nworld", "ntree", int)
island_ne: array("nworld", "ntree", int)
island_nf: array("nworld", "ntree", int)
island_efcadr: array("nworld", "ntree", int)
island_iefcadr: array("nworld", "ntree", int)
map_dof2idof: array("nworld", "nv", int)
map_idof2dof: array("nworld", "nv", int)
map_efc2iefc: array("nworld", "njmax", int)
map_iefc2efc: array("nworld", "njmax", int)
dof_islandid: array("nworld", "nv", int)
efc_islandid: array("nworld", "njmax", int)
iqacc: wp.array2d[float]
iqacc_smooth: wp.array2d[float]
iqfrc_smooth: wp.array2d[float]
iqfrc_constraint: wp.array2d[float]
ncdof: array("nworld", int)
dof_cdof: array("nworld", "nv", int)
cdof_dof: array("nworld", "nvmax_pad", int)
ctol: wp.array[float]
cls_tol: wp.array[float]
cdof_tri_row: wp.array[int]
cdof_tri_col: wp.array[int]
cM: wp.array3d[float]
cqLD: wp.array3d[float]
crhs: wp.array3d[float]
cx: wp.array3d[float]
cJ: wp.array3d[float]
cMa: wp.array2d[float]
cqfrc_smooth: wp.array2d[float]
cqacc_smooth: wp.array2d[float]
cqacc_warmstart: wp.array2d[float]
cqacc: wp.array2d[float]
cqfrc_constraint: wp.array2d[float]
# warp only fields:
nworld: int
naconmax: int
naccdmax: int
njmax: int
nvmax: int
nvmax_pad: int
njmax_pad: int
njmax_nnz: int
nacon: array(1, int)
ncollision: array(1, int)
flex_aabb_min: array("nworld", "nflex", wp.vec3)
flex_aabb_max: array("nworld", "nflex", wp.vec3)
overflow: array("nworld", int)
@dataclasses.dataclass
@@ -2185,35 +2293,6 @@ class InverseContext:
changed_efc_count: wp.array[int]
@dataclasses.dataclass
class IslandSolverContext:
"""Workspace arrays for island constraint solver."""
# Re-ordered workspace arrays (sized per-dof / per-constraint)
Jaref: wp.array2d[float]
jv: wp.array2d[float]
search: wp.array2d[float]
mv: wp.array2d[float]
grad: wp.array2d[float]
Mgrad: wp.array2d[float]
prev_grad: wp.array2d[float]
prev_Mgrad: wp.array2d[float]
h: wp.array3d[float]
# Per-island solver scalars (nworld, ntree)
cost: wp.array2d[float]
prev_cost: wp.array2d[float]
gauss: wp.array2d[float]
search_dot: wp.array2d[float]
grad_dot: wp.array2d[float]
done: wp.array2d[bool] # per-island convergence
solver_niter: wp.array2d[int] # iterations per island
beta: wp.array2d[float]
beta_den: wp.array2d[float]
alpha: wp.array2d[float]
Ma: wp.array2d[float] # island-local Ma (nworld, nv)
@dataclasses.dataclass
class SolverContext:
"""Workspace arrays for constraint solver."""
@@ -2294,9 +2373,9 @@ class RenderContext:
render_seg: per-camera segmentation render flags
znear: near plane distance
total_rays: total number of rays
render_skybox: whether to shade missed rays with the MuJoCo skybox texture
skybox_tex_id: index into textures of the skybox (MuJoCo tex_type == SKYBOX), -1 if none
skybox_face_width: pixel width of one skybox cube face (0 if no skybox)
render_skybox: whether to shade missed rays with a MuJoCo skybox texture
skybox_tex_id: per-world indices into textures of the skybox
skybox_face_width: per-world pixel widths of the skybox cube face
headlight_active: whether to inject MuJoCo's vis.headlight as a synthetic
directional light at the active camera. Read from `mjm.vis.headlight.active`
at context creation; users disable the headlight by configuring it on the
@@ -2342,8 +2421,8 @@ class RenderContext:
background_color: wp.uint32
use_precomputed_rays: bool
render_skybox: bool
skybox_tex_id: int
skybox_face_width: int
skybox_tex_id: array("*", int)
skybox_face_width: array("*", int)
headlight_active: bool
headlight_ambient: wp.vec3
headlight_diffuse: wp.vec3
+4 -4
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name="mujoco-warp"
version = "3.9.0.1"
version = "3.10.0.1"
# TODO(team): create a distribution list
authors = [
{name = "Newton Developers", email = "mujoco@deepmind.com"},
@@ -29,9 +29,9 @@ requires-python = ">=3.10"
dependencies = [
"absl-py",
"etils[epath]",
"mujoco>=3.8.0",
"mujoco>=3.9.0",
"numpy",
"warp-lang>=1.13",
"warp-lang>=1.14",
]
[[tool.uv.index]]
@@ -57,7 +57,7 @@ dev = [
"pygls>=1.0.0,<2.0.0",
"lsprotocol>=2023.0.1,<2024.0.0",
"mujoco>=3.8.0.dev0",
"warp-lang>=1.11.0.dev0",
"warp-lang>=1.14",
"mjviser>=0.0.10",
"pillow",
]
+2 -2
View File
@@ -180,9 +180,9 @@ def _main(argv: Sequence[str]) -> None:
print(f"MuJoCo Warp simulating with dt = {m.opt.timestep.numpy()[0]:.3f}...")
if _ENGINE.value == EngineOptions.WARP:
step_fn = _make_warp_step_fn(mjm, m, d, graph, ctrls)
step_fn = _make_warp_step_fn(mjm, m, d, graph, ctrls if cli.REPLAY.value else None)
else:
step_fn = _make_c_step_fn(ctrls)
step_fn = _make_c_step_fn(ctrls if cli.REPLAY.value else None)
mjd = mujoco.MjData(mjm)
mjw.get_data_into(mjd, mjm, d)
-4
View File
@@ -69,11 +69,7 @@ else:
def BlockDim(self, *args, **kwargs): # pylint: disable=invalid-name
pass
class _MjwpIoStub:
ENABLE_ISLANDS = True
WARP_INSTALLED: bool = True
warp: Any = _WpStub()
mujoco_warp: Any = _MjwpStub()
mjwp_types: Any = _MjwpTypesStub()
mjwp_io: Any = _MjwpIoStub()
+111 -25
View File
@@ -23,6 +23,7 @@ import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
import warp as wp
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
)
@@ -45,6 +46,7 @@ _cb = mjwp_types.Callback(
**{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init}
)
@ffi.format_args_for_warp
def _collision_shim(
# Model
@@ -57,20 +59,33 @@ def _collision_shim(
flex_elem: wp.array[int],
flex_elemadr: wp.array[int],
flex_elemdataadr: wp.array[int],
flex_elemflexid: wp.array[int],
flex_elemnum: wp.array[int],
flex_evpair: wp.array[wp.vec2i],
flex_evpairadr: wp.array[int],
flex_evpairflexid: wp.array[int],
flex_evpairnum: wp.array[int],
flex_friction: wp.array[wp.vec3],
flex_gap: wp.array[float],
flex_internal: wp.array[int],
flex_margin: wp.array[float],
flex_priority: wp.array[int],
flex_radius: wp.array[float],
flex_selfcollide: wp.array[int],
flex_shell: wp.array[int],
flex_shelladr: wp.array[int],
flex_shelldataadr: wp.array[int],
flex_shellnum: wp.array[int],
flex_shellflexid: 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_vertbodyid: wp.array[int],
flex_vertflexid: wp.array[int],
flex_vertnum: wp.array[int],
flexelem_geom_pair_filtered: wp.array[wp.vec2i],
flexshell_geom_pair_filtered: wp.array[wp.vec2i],
flexvert_geom_pair_filtered: wp.array[wp.vec2i],
geom_aabb: wp.array3d[wp.vec3],
geom_bodyid: wp.array[int],
geom_conaffinity: wp.array[int],
@@ -89,12 +104,14 @@ def _collision_shim(
geom_solmix: wp.array2d[float],
geom_solref: wp.array2d[wp.vec2],
geom_type: wp.array[int],
has_flex_selfcollide: bool,
has_sdf_geom: bool,
hfield_adr: wp.array[int],
hfield_data: wp.array[float],
hfield_ncol: wp.array[int],
hfield_nrow: wp.array[int],
hfield_size: wp.array[wp.vec4],
max_flex_dim: int,
mesh_face: wp.array[wp.vec3i],
mesh_faceadr: wp.array[int],
mesh_graph: wp.array[int],
@@ -109,16 +126,19 @@ def _collision_shim(
mesh_polyvert: wp.array[int],
mesh_polyvertadr: wp.array[int],
mesh_polyvertnum: wp.array[int],
mesh_pos: wp.array[wp.vec3],
mesh_vert: wp.array[wp.vec3],
mesh_vertadr: wp.array[int],
mesh_vertnum: wp.array[int],
nbody: int,
nflex: int,
nflexelem: int,
nflexshelldata: int,
nflexevpair: int,
nflexvert: int,
ngeom: int,
nmaxmeshdeg: int,
nmaxpolygon: int,
nmesh: int,
nmeshface: int,
nxn_geom_pair_filtered: wp.array[wp.vec2i],
nxn_pairid: wp.array[wp.vec2i],
@@ -141,20 +161,26 @@ def _collision_shim(
opt__ccd_tolerance: wp.array[float],
opt__disableflags: int,
opt__enableflags: int,
opt__graph_conditional: bool,
opt__sdf_initpoints: int,
opt__sdf_iterations: int,
opt__warn_overflow: bool,
# Data
naccdmax: int,
naconmax: int,
body_awake: wp.array2d[int],
flex_aabb_max: wp.array2d[wp.vec3],
flex_aabb_min: wp.array2d[wp.vec3],
flexvert_xpos: wp.array2d[wp.vec3],
geom_xmat: wp.array2d[wp.mat33],
geom_xpos: wp.array2d[wp.vec3],
nacon: wp.array[int],
ncollision: wp.array[int],
overflow: wp.array[int],
contact__dim: wp.array[int],
contact__dist: wp.array[float],
contact__efc_address: wp.array2d[int],
contact__elem: wp.array[wp.vec2i],
contact__flex: wp.array[wp.vec2i],
contact__frame: wp.array[wp.mat33],
contact__friction: wp.array[mjwp_types.vec5],
@@ -182,20 +208,33 @@ def _collision_shim(
_m.flex_elem = flex_elem
_m.flex_elemadr = flex_elemadr
_m.flex_elemdataadr = flex_elemdataadr
_m.flex_elemflexid = flex_elemflexid
_m.flex_elemnum = flex_elemnum
_m.flex_evpair = flex_evpair
_m.flex_evpairadr = flex_evpairadr
_m.flex_evpairflexid = flex_evpairflexid
_m.flex_evpairnum = flex_evpairnum
_m.flex_friction = flex_friction
_m.flex_gap = flex_gap
_m.flex_internal = flex_internal
_m.flex_margin = flex_margin
_m.flex_priority = flex_priority
_m.flex_radius = flex_radius
_m.flex_selfcollide = flex_selfcollide
_m.flex_shell = flex_shell
_m.flex_shelladr = flex_shelladr
_m.flex_shelldataadr = flex_shelldataadr
_m.flex_shellnum = flex_shellnum
_m.flex_shellflexid = flex_shellflexid
_m.flex_solimp = flex_solimp
_m.flex_solmix = flex_solmix
_m.flex_solref = flex_solref
_m.flex_vertadr = flex_vertadr
_m.flex_vertbodyid = flex_vertbodyid
_m.flex_vertflexid = flex_vertflexid
_m.flex_vertnum = flex_vertnum
_m.flexelem_geom_pair_filtered = flexelem_geom_pair_filtered
_m.flexshell_geom_pair_filtered = flexshell_geom_pair_filtered
_m.flexvert_geom_pair_filtered = flexvert_geom_pair_filtered
_m.geom_aabb = geom_aabb
_m.geom_bodyid = geom_bodyid
_m.geom_conaffinity = geom_conaffinity
@@ -214,12 +253,14 @@ def _collision_shim(
_m.geom_solmix = geom_solmix
_m.geom_solref = geom_solref
_m.geom_type = geom_type
_m.has_flex_selfcollide = has_flex_selfcollide
_m.has_sdf_geom = has_sdf_geom
_m.hfield_adr = hfield_adr
_m.hfield_data = hfield_data
_m.hfield_ncol = hfield_ncol
_m.hfield_nrow = hfield_nrow
_m.hfield_size = hfield_size
_m.max_flex_dim = max_flex_dim
_m.mesh_face = mesh_face
_m.mesh_faceadr = mesh_faceadr
_m.mesh_graph = mesh_graph
@@ -234,16 +275,19 @@ def _collision_shim(
_m.mesh_polyvert = mesh_polyvert
_m.mesh_polyvertadr = mesh_polyvertadr
_m.mesh_polyvertnum = mesh_polyvertnum
_m.mesh_pos = mesh_pos
_m.mesh_vert = mesh_vert
_m.mesh_vertadr = mesh_vertadr
_m.mesh_vertnum = mesh_vertnum
_m.nbody = nbody
_m.nflex = nflex
_m.nflexelem = nflexelem
_m.nflexshelldata = nflexshelldata
_m.nflexevpair = nflexevpair
_m.nflexvert = nflexvert
_m.ngeom = ngeom
_m.nmaxmeshdeg = nmaxmeshdeg
_m.nmaxpolygon = nmaxpolygon
_m.nmesh = nmesh
_m.nmeshface = nmeshface
_m.nxn_geom_pair_filtered = nxn_geom_pair_filtered
_m.nxn_pairid = nxn_pairid
@@ -257,8 +301,10 @@ def _collision_shim(
_m.opt.ccd_tolerance = opt__ccd_tolerance
_m.opt.disableflags = opt__disableflags
_m.opt.enableflags = opt__enableflags
_m.opt.graph_conditional = opt__graph_conditional
_m.opt.sdf_initpoints = opt__sdf_initpoints
_m.opt.sdf_iterations = opt__sdf_iterations
_m.opt.warn_overflow = opt__warn_overflow
_m.pair_dim = pair_dim
_m.pair_friction = pair_friction
_m.pair_gap = pair_gap
@@ -272,6 +318,7 @@ def _collision_shim(
_d.contact.dim = contact__dim
_d.contact.dist = contact__dist
_d.contact.efc_address = contact__efc_address
_d.contact.elem = contact__elem
_d.contact.flex = contact__flex
_d.contact.frame = contact__frame
_d.contact.friction = contact__friction
@@ -285,6 +332,8 @@ def _collision_shim(
_d.contact.type = contact__type
_d.contact.vert = contact__vert
_d.contact.worldid = contact__worldid
_d.flex_aabb_max = flex_aabb_max
_d.flex_aabb_min = flex_aabb_min
_d.flexvert_xpos = flexvert_xpos
_d.geom_xmat = geom_xmat
_d.geom_xpos = geom_xpos
@@ -292,17 +341,22 @@ def _collision_shim(
_d.nacon = nacon
_d.naconmax = naconmax
_d.ncollision = ncollision
_d.overflow = overflow
_d.nworld = nworld
mjwarp.collision(_m, _d)
def _collision_jax_impl(m: types.Model, d: types.Data):
output_dims = {
'flex_aabb_max': d._impl.flex_aabb_max.shape,
'flex_aabb_min': d._impl.flex_aabb_min.shape,
'nacon': d._impl.nacon.shape,
'ncollision': d._impl.ncollision.shape,
'overflow': d._impl.overflow.shape,
'contact__dim': d._impl.contact__dim.shape,
'contact__dist': d._impl.contact__dist.shape,
'contact__efc_address': d._impl.contact__efc_address.shape,
'contact__elem': d._impl.contact__elem.shape,
'contact__flex': d._impl.contact__flex.shape,
'contact__frame': d._impl.contact__frame.shape,
'contact__friction': d._impl.contact__friction.shape,
@@ -319,15 +373,19 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
}
jf = ffi.jax_callable_variadic_tuple(
_collision_shim,
num_outputs=18,
num_outputs=22,
output_dims=output_dims,
vmap_method=None,
in_out_argnames=set([
'flex_aabb_max',
'flex_aabb_min',
'nacon',
'ncollision',
'overflow',
'contact__dim',
'contact__dist',
'contact__efc_address',
'contact__elem',
'contact__flex',
'contact__frame',
'contact__friction',
@@ -376,20 +434,33 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
m._impl.flex_elem,
m._impl.flex_elemadr,
m._impl.flex_elemdataadr,
m._impl.flex_elemflexid,
m._impl.flex_elemnum,
m._impl.flex_evpair,
m._impl.flex_evpairadr,
m._impl.flex_evpairflexid,
m._impl.flex_evpairnum,
m._impl.flex_friction,
m._impl.flex_gap,
m._impl.flex_internal,
m._impl.flex_margin,
m._impl.flex_priority,
m._impl.flex_radius,
m._impl.flex_selfcollide,
m._impl.flex_shell,
m._impl.flex_shelladr,
m._impl.flex_shelldataadr,
m._impl.flex_shellnum,
m._impl.flex_shellflexid,
m._impl.flex_solimp,
m._impl.flex_solmix,
m._impl.flex_solref,
m.flex_vertadr,
m._impl.flex_vertbodyid,
m._impl.flex_vertflexid,
m.flex_vertnum,
m._impl.flexelem_geom_pair_filtered,
m._impl.flexshell_geom_pair_filtered,
m._impl.flexvert_geom_pair_filtered,
m.geom_aabb,
m.geom_bodyid,
m.geom_conaffinity,
@@ -408,12 +479,14 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
m.geom_solmix,
m.geom_solref,
m.geom_type,
m._impl.has_flex_selfcollide,
m._impl.has_sdf_geom,
m.hfield_adr,
m.hfield_data,
m.hfield_ncol,
m.hfield_nrow,
m.hfield_size,
m._impl.max_flex_dim,
m.mesh_face,
m.mesh_faceadr,
m.mesh_graph,
@@ -428,16 +501,19 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
m._impl.mesh_polyvert,
m._impl.mesh_polyvertadr,
m._impl.mesh_polyvertnum,
m.mesh_pos,
m.mesh_vert,
m.mesh_vertadr,
m.mesh_vertnum,
m.nbody,
m.nflex,
m._impl.nflexelem,
m._impl.nflexshelldata,
m._impl.nflexevpair,
m._impl.nflexvert,
m.ngeom,
m._impl.nmaxmeshdeg,
m._impl.nmaxpolygon,
m.nmesh,
m.nmeshface,
m._impl.nxn_geom_pair_filtered,
m._impl.nxn_pairid,
@@ -460,19 +536,25 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
m.opt._impl.ccd_tolerance,
m.opt.disableflags,
m.opt.enableflags,
m.opt._impl.graph_conditional,
m.opt._impl.sdf_initpoints,
m.opt._impl.sdf_iterations,
m.opt._impl.warn_overflow,
d._impl.naccdmax,
d._impl.naconmax,
d._impl.body_awake,
d._impl.flex_aabb_max,
d._impl.flex_aabb_min,
d._impl.flexvert_xpos,
d.geom_xmat,
d.geom_xpos,
d._impl.nacon,
d._impl.ncollision,
d._impl.overflow,
d._impl.contact__dim,
d._impl.contact__dist,
d._impl.contact__efc_address,
d._impl.contact__elem,
d._impl.contact__flex,
d._impl.contact__frame,
d._impl.contact__friction,
@@ -488,24 +570,28 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
d._impl.contact__worldid,
)
d = d.tree_replace({
'_impl.nacon': out[0],
'_impl.ncollision': out[1],
'_impl.contact__dim': out[2],
'_impl.contact__dist': out[3],
'_impl.contact__efc_address': out[4],
'_impl.contact__flex': out[5],
'_impl.contact__frame': out[6],
'_impl.contact__friction': out[7],
'_impl.contact__geom': out[8],
'_impl.contact__geomcollisionid': out[9],
'_impl.contact__includemargin': out[10],
'_impl.contact__pos': out[11],
'_impl.contact__solimp': out[12],
'_impl.contact__solref': out[13],
'_impl.contact__solreffriction': out[14],
'_impl.contact__type': out[15],
'_impl.contact__vert': out[16],
'_impl.contact__worldid': out[17],
'_impl.flex_aabb_max': out[0],
'_impl.flex_aabb_min': out[1],
'_impl.nacon': out[2],
'_impl.ncollision': out[3],
'_impl.overflow': out[4],
'_impl.contact__dim': out[5],
'_impl.contact__dist': out[6],
'_impl.contact__efc_address': out[7],
'_impl.contact__elem': out[8],
'_impl.contact__flex': out[9],
'_impl.contact__frame': out[10],
'_impl.contact__friction': out[11],
'_impl.contact__geom': out[12],
'_impl.contact__geomcollisionid': out[13],
'_impl.contact__includemargin': out[14],
'_impl.contact__pos': out[15],
'_impl.contact__solimp': out[16],
'_impl.contact__solref': out[17],
'_impl.contact__solreffriction': out[18],
'_impl.contact__type': out[19],
'_impl.contact__vert': out[20],
'_impl.contact__worldid': out[21],
})
return d
+4 -14
View File
@@ -168,20 +168,6 @@ def _format_arg(arg: Any, name: str, annotation: Any, verbose: bool):
f'Arg ndim {arg.ndim} does not match expected ndim {expected_ndim}.'
)
# Add stride 0 to first axis in case the underlying argument should be
# batched.
# NB: the outer marshalling does an "expand_dims" on Model fields.
is_batch_field = mjx_warp_types._BATCH_DIM['Model'].get(name, False) # pylint: disable=protected-access
if arg.shape[0] == 1 and is_batch_field:
old_strides = arg.strides
arg.strides = (0,) + arg.strides[1:]
if verbose:
print(
f'Leading batch dim of 1, adding stride: {name} {old_strides} =>'
f' {arg.strides}'
)
return arg
if verbose:
print(f'Did nothing: {name}: {arg.shape}')
return arg
@@ -347,6 +333,8 @@ def _check_leading_dim(
not has_batch_dim
and attr.startswith('contact__')
and leaf.shape[0] != expected_naconmax
and leaf.shape[0]
!= 0 # Allow empty arrays (shape[0] == 0) for unused contact fields (e.g., flex).
):
raise ValueError(
f'Leaf node leading dim ({leaf.shape[0]}) does not match naconmax'
@@ -356,6 +344,8 @@ def _check_leading_dim(
not has_batch_dim
and attr.startswith('efc__')
and leaf.shape[0] != expected_njmax
and leaf.shape[0]
!= 0 # Allow empty arrays (shape[0] == 0) for unused constraint fields (e.g., islands).
):
raise ValueError(
f'Leaf node leading dim ({leaf.shape[0]}) does not match njmax'
File diff suppressed because it is too large Load Diff
+9 -2
View File
@@ -176,8 +176,15 @@ class ForwardTest(parameterized.TestCase):
qm = np.zeros((m.nv, m.nv))
mujoco.mju_sym2dense(qm, d.M, m.M_rownnz, m.M_rowadr, m.M_colind)
# mjwarp adds padding to M
tu.assert_eq(qm, dx._impl.M[: m.nv, : m.nv], 'M')
warp_M = np.zeros((m.nv, m.nv))
mujoco.mju_sym2dense(
warp_M,
np.asarray(dx._impl.M),
np.asarray(mx.M_rownnz),
np.asarray(mx.M_rowadr),
np.asarray(mx.M_colind),
)
tu.assert_eq(qm, warp_M, 'M')
# qLD is fused in a cholesky factorize and solve, and not written to.
tu.assert_contact_eq(d, dx, worldid=i)
-1
View File
@@ -25,7 +25,6 @@ import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
import warp as wp
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
)
+1
View File
@@ -23,6 +23,7 @@ import mujoco.mjx.third_party.mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
import warp as wp
_m = mjwarp.Model(
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
)
+226 -83
View File
@@ -53,9 +53,12 @@ class TileSet:
Attributes:
adr: address of each tile in the set
size: size of all the tiles in this set
elemid: flat per-block gather indices into CSR M for tile_load_indexed
(absent -> nC sentinel)
"""
adr: np.ndarray
size: int
elemid: typing.Optional[np.ndarray] = None
def __eq__(self, other) -> bool:
if self.__class__ is not other.__class__:
@@ -87,15 +90,15 @@ class BlockDim:
Attributes:
segmented_sort: segmented sort block dimension (collision_driver)
euler_dense: Euler dense block dimension (forward)
convex_ccd: convex CCD kernel block dimension (collision_convex)
actuator_velocity: actuator velocity block dimension (forward)
ray: ray block dimension (ray)
contact_sort: contact sort block dimension (sensor)
energy_vel_kinetic: energy velocity kinetic block dimension (sensor)
cholesky_factorize: Cholesky factorize block dimension (smooth)
cholesky_factorize: block-dense Cholesky factorize block dimension (smooth)
cholesky_factorize_solve: block-dense Cholesky factorize+solve block
dimension (smooth)
cholesky_solve: Cholesky solve block dimension (smooth)
cholesky_factorize_solve: Cholesky factorize and solve block dimension
(smooth)
solve_LD_sparse_fused: solve LD sparse fused block dimension (smooth)
update_gradient_cholesky: update gradient Cholesky block dimension (solver)
update_gradient_cholesky_blocked: update gradient Cholesky blocked block
@@ -113,28 +116,28 @@ class BlockDim:
qderiv_actuator_dense: qderiv actuator dense block dimension (derivative)
render: render block dimension (render)
"""
actuator_velocity: int
cholesky_factorize: int
cholesky_factorize_solve: int
cholesky_solve: int
contact_jac_tiled: int
contact_sort: int
energy_vel_kinetic: int
euler_dense: int
linesearch_iterative: int
qderiv_actuator_dense: int
ray: int
render: int
segmented_sort: int
solve_LD_sparse_fused: int
solve_beta_accumulate: int
solve_init_search_cg: int
solve_search_update_cg: int
update_gradient_JTDAJ_dense: int
update_gradient_JTDAJ_sparse: int
update_gradient_cholesky: int
update_gradient_cholesky_blocked: int
update_gradient_grad: int
segmented_sort: int = 128
convex_ccd: int = 256
actuator_velocity: int = 32
ray: int = 64
contact_sort: int = 64
energy_vel_kinetic: int = 32
cholesky_factorize: int = 32
cholesky_factorize_solve: int = 32
cholesky_solve: int = 64
solve_LD_sparse_fused: int = 128
update_gradient_cholesky: int = 64
update_gradient_cholesky_blocked: int = 32
update_gradient_JTDAJ_sparse: int = 128
update_gradient_JTDAJ_dense: int = 128
linesearch_iterative: int = 32
update_gradient_grad: int = 256
solve_beta_accumulate: int = 256
solve_search_update_cg: int = 256
solve_init_search_cg: int = 256
contact_jac_tiled: int = 32
qderiv_actuator_dense: int = 32
render: int = 64
def tree_flatten(self):
children = list((getattr(self, k) for k in self.__dataclass_fields__))
@@ -164,6 +167,7 @@ class OptionWarp(PyTreeNode):
sdf_initpoints: int
sdf_iterations: int
sleep_tolerance: jax.Array
warn_overflow: bool
class ModelWarp(PyTreeNode):
"""Derived fields from Model."""
@@ -177,6 +181,7 @@ class ModelWarp(PyTreeNode):
M_fullm_upper_elemid: np.ndarray
M_fullm_upper_i: np.ndarray
M_fullm_upper_j: np.ndarray
M_hinit_i: np.ndarray
M_mulm_col: np.ndarray
M_mulm_madr: np.ndarray
M_mulm_rowadr: np.ndarray
@@ -188,7 +193,9 @@ class ModelWarp(PyTreeNode):
block_dim: BlockDim
body_branch_start: np.ndarray
body_branches: np.ndarray
body_fluid_box_adr: np.ndarray
body_fluid_ellipsoid: np.ndarray
body_fluid_ellipsoid_adr: np.ndarray
body_isdofancestor: np.ndarray
body_tree: Tuple[np.ndarray, ...]
callback: Callback
@@ -219,14 +226,23 @@ class ModelWarp(PyTreeNode):
flex_elemdataadr: np.ndarray
flex_elemedge: np.ndarray
flex_elemedgeadr: np.ndarray
flex_elemflexid: np.ndarray
flex_elemnum: np.ndarray
flex_evpair: np.ndarray
flex_evpairadr: np.ndarray
flex_evpairflexid: np.ndarray
flex_evpairnum: np.ndarray
flex_friction: np.ndarray
flex_gap: np.ndarray
flex_internal: np.ndarray
flex_margin: np.ndarray
flex_priority: np.ndarray
flex_radius: np.ndarray
flex_selfcollide: np.ndarray
flex_shell: np.ndarray
flex_shelladr: np.ndarray
flex_shelldataadr: np.ndarray
flex_shellflexid: np.ndarray
flex_shellnum: np.ndarray
flex_solimp: np.ndarray
flex_solmix: np.ndarray
@@ -241,8 +257,12 @@ class ModelWarp(PyTreeNode):
flexedge_J_rownnz: np.ndarray
flexedge_invweight0: np.ndarray
flexedge_length0: np.ndarray
flexelem_geom_pair_filtered: np.ndarray
flexshell_geom_pair_filtered: np.ndarray
flexvert_geom_pair_filtered: np.ndarray
geom_pair_type_count: Tuple[int, ...]
geom_plugin_index: np.ndarray
has_flex_selfcollide: bool
has_fluid: bool
has_sdf_geom: bool
is_sparse: bool
@@ -255,6 +275,7 @@ class ModelWarp(PyTreeNode):
mapM2D: np.ndarray
mapM2M: np.ndarray
mat_texrepeat: jax.Array
max_flex_dim: int
max_ten_J_rownnz: int
mesh_polyadr: np.ndarray
mesh_polymap: np.ndarray
@@ -274,6 +295,7 @@ class ModelWarp(PyTreeNode):
nflexelem: int
nflexelemdata: int
nflexelemedge: int
nflexevpair: int
nflexshelldata: int
nflexstiffness: int
nflexvert: int
@@ -301,7 +323,15 @@ class ModelWarp(PyTreeNode):
qD_fullm_i: np.ndarray
qD_fullm_j: np.ndarray
qLD_all_updates: np.ndarray
qLD_block_adr: np.ndarray
qLD_block_total: int
qLD_dof_dense: np.ndarray
qLD_dof_simple: np.ndarray
qLD_has_dense: bool
qLD_has_simple: bool
qLD_has_sparse: bool
qLD_level_offsets: np.ndarray
qLD_simple_dofs: np.ndarray
qLD_updates: Tuple[np.ndarray, ...]
rangefinder_sensor_adr: np.ndarray
sensor_acc_adr: np.ndarray
@@ -353,13 +383,21 @@ class DataWarp(PyTreeNode):
actuator_velocity: jax.Array
body_awake: jax.Array
body_awake_ind: jax.Array
cJ: jax.Array
cM: jax.Array
cMa: jax.Array
cacc: jax.Array
cdof_dof: jax.Array
cdof_tri_col: jax.Array
cdof_tri_row: jax.Array
cfrc_ext: jax.Array
cfrc_int: jax.Array
cinert: jax.Array
cls_tol: jax.Array
contact__dim: jax.Array
contact__dist: jax.Array
contact__efc_address: jax.Array
contact__elem: jax.Array
contact__flex: jax.Array
contact__frame: jax.Array
contact__friction: jax.Array
@@ -373,8 +411,18 @@ class DataWarp(PyTreeNode):
contact__type: jax.Array
contact__vert: jax.Array
contact__worldid: jax.Array
cqLD: jax.Array
cqacc: jax.Array
cqacc_smooth: jax.Array
cqacc_warmstart: jax.Array
cqfrc_constraint: jax.Array
cqfrc_smooth: jax.Array
crb: jax.Array
crhs: jax.Array
ctol: jax.Array
cx: jax.Array
dof_awake_ind: jax.Array
dof_cdof: jax.Array
dof_island: jax.Array
dof_islandid: jax.Array
efc__D: jax.Array
@@ -387,19 +435,11 @@ class DataWarp(PyTreeNode):
efc__aref: jax.Array
efc__force: jax.Array
efc__frictionloss: jax.Array
efc__iD: jax.Array
efc__iJ: jax.Array
efc__iJ_colind: jax.Array
efc__iJ_rowadr: jax.Array
efc__iJ_rownnz: jax.Array
efc__iaref: jax.Array
efc__id: jax.Array
efc__iforce: jax.Array
efc__ifrictionloss: jax.Array
efc__iid: jax.Array
efc__island: jax.Array
efc__istate: jax.Array
efc__itype: jax.Array
efc__jtdaj_adr: jax.Array
efc__jtdaj_nblock: jax.Array
efc__jtdaj_nrow: jax.Array
efc__margin: jax.Array
efc__pos: jax.Array
efc__state: jax.Array
@@ -407,17 +447,15 @@ class DataWarp(PyTreeNode):
efc__vel: jax.Array
efc_islandid: jax.Array
energy: jax.Array
flex_aabb_max: jax.Array
flex_aabb_min: jax.Array
flexedge_J: jax.Array
flexedge_length: jax.Array
flexedge_velocity: jax.Array
flexvert_xpos: jax.Array
iqacc: jax.Array
iqacc_smooth: jax.Array
iqfrc_constraint: jax.Array
iqfrc_smooth: jax.Array
island_dofadr: jax.Array
island_efcadr: jax.Array
island_idofadr: jax.Array
island_iefcadr: jax.Array
island_ne: jax.Array
island_nefc: jax.Array
island_nf: jax.Array
@@ -435,6 +473,7 @@ class DataWarp(PyTreeNode):
nacon: jax.Array
naconmax: int
nbody_awake: jax.Array
ncdof: jax.Array
ncollision: jax.Array
ne: jax.Array
nefc: jax.Array
@@ -447,7 +486,10 @@ class DataWarp(PyTreeNode):
nl: jax.Array
ntree_awake: jax.Array
nv_awake: jax.Array
nvmax: int
nvmax_pad: int
nworld: int
overflow: jax.Array
qLD: jax.Array
qLDiagInv: jax.Array
qLU: jax.Array
@@ -467,9 +509,16 @@ class DataWarp(PyTreeNode):
wrap_xpos: jax.Array
shape = property(lambda self: self.cacc.shape)
DATA_NON_VMAP = {
'cJ',
'cM',
'cMa',
'cdof_tri_col',
'cdof_tri_row',
'cls_tol',
'contact__dim',
'contact__dist',
'contact__efc_address',
'contact__elem',
'contact__flex',
'contact__frame',
'contact__friction',
@@ -483,14 +532,24 @@ DATA_NON_VMAP = {
'contact__type',
'contact__vert',
'contact__worldid',
'cqLD',
'cqacc',
'cqacc_smooth',
'cqacc_warmstart',
'cqfrc_constraint',
'cqfrc_smooth',
'crhs',
'ctol',
'cx',
'naccdmax',
'nacon',
'naconmax',
'ncollision',
'nidof',
'njmax',
'njmax_nnz',
'njmax_pad',
'nvmax',
'nvmax_pad',
'nworld',
}
@@ -520,7 +579,7 @@ batching.register_vmappable(DataWarp, int, int, _to_elt, _from_elt, None)
_NDIM = {
'Data': {
'M': 3,
'M': 2,
'act': 2,
'act_dot': 2,
'actuator_force': 2,
@@ -529,17 +588,25 @@ _NDIM = {
'actuator_velocity': 2,
'body_awake': 2,
'body_awake_ind': 2,
'cJ': 3,
'cM': 3,
'cMa': 2,
'cacc': 3,
'cam_xmat': 4,
'cam_xpos': 3,
'cdof': 3,
'cdof_dof': 2,
'cdof_dot': 3,
'cdof_tri_col': 1,
'cdof_tri_row': 1,
'cfrc_ext': 3,
'cfrc_int': 3,
'cinert': 3,
'cls_tol': 1,
'contact__dim': 1,
'contact__dist': 1,
'contact__efc_address': 2,
'contact__elem': 2,
'contact__flex': 2,
'contact__frame': 3,
'contact__friction': 2,
@@ -553,10 +620,20 @@ _NDIM = {
'contact__type': 1,
'contact__vert': 2,
'contact__worldid': 1,
'cqLD': 3,
'cqacc': 2,
'cqacc_smooth': 2,
'cqacc_warmstart': 2,
'cqfrc_constraint': 2,
'cqfrc_smooth': 2,
'crb': 3,
'crhs': 3,
'ctol': 1,
'ctrl': 2,
'cvel': 3,
'cx': 3,
'dof_awake_ind': 2,
'dof_cdof': 2,
'dof_island': 2,
'dof_islandid': 2,
'efc__D': 2,
@@ -569,19 +646,11 @@ _NDIM = {
'efc__aref': 2,
'efc__force': 2,
'efc__frictionloss': 2,
'efc__iD': 2,
'efc__iJ': 3,
'efc__iJ_colind': 3,
'efc__iJ_rowadr': 2,
'efc__iJ_rownnz': 2,
'efc__iaref': 2,
'efc__id': 2,
'efc__iforce': 2,
'efc__ifrictionloss': 2,
'efc__iid': 2,
'efc__island': 2,
'efc__istate': 2,
'efc__itype': 2,
'efc__jtdaj_adr': 2,
'efc__jtdaj_nblock': 1,
'efc__jtdaj_nrow': 2,
'efc__margin': 2,
'efc__pos': 2,
'efc__state': 2,
@@ -590,6 +659,8 @@ _NDIM = {
'efc_islandid': 2,
'energy': 2,
'eq_active': 2,
'flex_aabb_max': 3,
'flex_aabb_min': 3,
'flexedge_J': 2,
'flexedge_length': 2,
'flexedge_velocity': 2,
@@ -597,13 +668,9 @@ _NDIM = {
'geom_xmat': 4,
'geom_xpos': 3,
'history': 2,
'iqacc': 2,
'iqacc_smooth': 2,
'iqfrc_constraint': 2,
'iqfrc_smooth': 2,
'island_dofadr': 2,
'island_efcadr': 2,
'island_idofadr': 2,
'island_iefcadr': 2,
'island_ne': 2,
'island_nefc': 2,
'island_nf': 2,
@@ -623,6 +690,7 @@ _NDIM = {
'nacon': 1,
'naconmax': 0,
'nbody_awake': 1,
'ncdof': 1,
'ncollision': 1,
'ne': 1,
'nefc': 1,
@@ -635,10 +703,13 @@ _NDIM = {
'nl': 1,
'ntree_awake': 1,
'nv_awake': 1,
'nvmax': 0,
'nvmax_pad': 0,
'nworld': 0,
'qLD': 3,
'overflow': 1,
'qLD': 2,
'qLDiagInv': 2,
'qLU': 3,
'qLU': 2,
'qacc': 2,
'qacc_smooth': 2,
'qacc_warmstart': 2,
@@ -671,6 +742,7 @@ _NDIM = {
'tree_asleep': 2,
'tree_awake': 2,
'tree_island': 2,
'userdata': 2,
'wrap_obj': 3,
'wrap_xpos': 3,
'xanchor': 3,
@@ -694,6 +766,7 @@ _NDIM = {
'M_fullm_upper_elemid': 1,
'M_fullm_upper_i': 1,
'M_fullm_upper_j': 1,
'M_hinit_i': 1,
'M_mulm_col': 1,
'M_mulm_madr': 1,
'M_mulm_rowadr': 1,
@@ -731,8 +804,8 @@ _NDIM = {
'block_dim__cholesky_solve': 0,
'block_dim__contact_jac_tiled': 0,
'block_dim__contact_sort': 0,
'block_dim__convex_ccd': 0,
'block_dim__energy_vel_kinetic': 0,
'block_dim__euler_dense': 0,
'block_dim__linesearch_iterative': 0,
'block_dim__qderiv_actuator_dense': 0,
'block_dim__ray': 0,
@@ -753,7 +826,9 @@ _NDIM = {
'body_contype': 1,
'body_dofadr': 1,
'body_dofnum': 1,
'body_fluid_box_adr': 1,
'body_fluid_ellipsoid': 1,
'body_fluid_ellipsoid_adr': 1,
'body_geomadr': 1,
'body_geomnum': 1,
'body_gravcomp': 2,
@@ -834,14 +909,23 @@ _NDIM = {
'flex_elemdataadr': 1,
'flex_elemedge': 1,
'flex_elemedgeadr': 1,
'flex_elemflexid': 1,
'flex_elemnum': 1,
'flex_evpair': 2,
'flex_evpairadr': 1,
'flex_evpairflexid': 1,
'flex_evpairnum': 1,
'flex_friction': 2,
'flex_gap': 1,
'flex_internal': 1,
'flex_margin': 1,
'flex_priority': 1,
'flex_radius': 1,
'flex_selfcollide': 1,
'flex_shell': 1,
'flex_shelladr': 1,
'flex_shelldataadr': 1,
'flex_shellflexid': 1,
'flex_shellnum': 1,
'flex_solimp': 2,
'flex_solmix': 1,
@@ -858,6 +942,9 @@ _NDIM = {
'flexedge_J_rownnz': 1,
'flexedge_invweight0': 1,
'flexedge_length0': 1,
'flexelem_geom_pair_filtered': 2,
'flexshell_geom_pair_filtered': 2,
'flexvert_geom_pair_filtered': 2,
'geom_aabb': 4,
'geom_bodyid': 1,
'geom_conaffinity': 1,
@@ -882,6 +969,7 @@ _NDIM = {
'geom_solmix': 2,
'geom_solref': 3,
'geom_type': 1,
'has_flex_selfcollide': 0,
'has_fluid': 0,
'has_sdf_geom': 0,
'hfield_adr': 1,
@@ -935,6 +1023,7 @@ _NDIM = {
'mat_specular': 2,
'mat_texid': 3,
'mat_texrepeat': 3,
'max_flex_dim': 0,
'max_ten_J_rownnz': 0,
'mesh_face': 2,
'mesh_faceadr': 1,
@@ -953,6 +1042,7 @@ _NDIM = {
'mesh_polyvert': 1,
'mesh_polyvertadr': 1,
'mesh_polyvertnum': 1,
'mesh_pos': 2,
'mesh_quat': 2,
'mesh_vert': 2,
'mesh_vertadr': 1,
@@ -977,11 +1067,11 @@ _NDIM = {
'nflexelem': 0,
'nflexelemdata': 0,
'nflexelemedge': 0,
'nflexevpair': 0,
'nflexshelldata': 0,
'nflexstiffness': 0,
'nflexvert': 0,
'ngeom': 0,
'ngravcomp': 0,
'nhfield': 0,
'nhfielddata': 0,
'nhistory': 0,
@@ -1015,6 +1105,7 @@ _NDIM = {
'ntendon': 0,
'ntree': 0,
'nu': 0,
'nuserdata': 0,
'nv': 0,
'nv_pad': 0,
'nwrap': 0,
@@ -1048,6 +1139,7 @@ _NDIM = {
'opt__timestep': 1,
'opt__tolerance': 1,
'opt__viscosity': 1,
'opt__warn_overflow': 0,
'opt__wind': 2,
'pair_dim': 1,
'pair_friction': 3,
@@ -1063,7 +1155,15 @@ _NDIM = {
'qD_fullm_i': 1,
'qD_fullm_j': 1,
'qLD_all_updates': 2,
'qLD_block_adr': 1,
'qLD_block_total': 0,
'qLD_dof_dense': 1,
'qLD_dof_simple': 1,
'qLD_has_dense': 0,
'qLD_has_simple': 0,
'qLD_has_sparse': 0,
'qLD_level_offsets': 1,
'qLD_simple_dofs': 1,
'qLD_updates': -1,
'qpos0': 2,
'qpos_spring': 2,
@@ -1173,6 +1273,7 @@ _NDIM = {
'timestep': 1,
'tolerance': 1,
'viscosity': 1,
'warn_overflow': 0,
'wind': 2,
},
'Statistic': {'meaninertia': 1},
@@ -1188,17 +1289,25 @@ _BATCH_DIM = {
'actuator_velocity': True,
'body_awake': True,
'body_awake_ind': True,
'cJ': False,
'cM': False,
'cMa': False,
'cacc': True,
'cam_xmat': True,
'cam_xpos': True,
'cdof': True,
'cdof_dof': True,
'cdof_dot': True,
'cdof_tri_col': False,
'cdof_tri_row': False,
'cfrc_ext': True,
'cfrc_int': True,
'cinert': True,
'cls_tol': False,
'contact__dim': False,
'contact__dist': False,
'contact__efc_address': False,
'contact__elem': False,
'contact__flex': False,
'contact__frame': False,
'contact__friction': False,
@@ -1212,10 +1321,20 @@ _BATCH_DIM = {
'contact__type': False,
'contact__vert': False,
'contact__worldid': False,
'cqLD': False,
'cqacc': False,
'cqacc_smooth': False,
'cqacc_warmstart': False,
'cqfrc_constraint': False,
'cqfrc_smooth': False,
'crb': True,
'crhs': False,
'ctol': False,
'ctrl': True,
'cvel': True,
'cx': False,
'dof_awake_ind': True,
'dof_cdof': True,
'dof_island': True,
'dof_islandid': True,
'efc__D': True,
@@ -1228,19 +1347,11 @@ _BATCH_DIM = {
'efc__aref': True,
'efc__force': True,
'efc__frictionloss': True,
'efc__iD': True,
'efc__iJ': True,
'efc__iJ_colind': True,
'efc__iJ_rowadr': True,
'efc__iJ_rownnz': True,
'efc__iaref': True,
'efc__id': True,
'efc__iforce': True,
'efc__ifrictionloss': True,
'efc__iid': True,
'efc__island': True,
'efc__istate': True,
'efc__itype': True,
'efc__jtdaj_adr': True,
'efc__jtdaj_nblock': True,
'efc__jtdaj_nrow': True,
'efc__margin': True,
'efc__pos': True,
'efc__state': True,
@@ -1249,6 +1360,8 @@ _BATCH_DIM = {
'efc_islandid': True,
'energy': True,
'eq_active': True,
'flex_aabb_max': True,
'flex_aabb_min': True,
'flexedge_J': True,
'flexedge_length': True,
'flexedge_velocity': True,
@@ -1256,13 +1369,9 @@ _BATCH_DIM = {
'geom_xmat': True,
'geom_xpos': True,
'history': True,
'iqacc': True,
'iqacc_smooth': True,
'iqfrc_constraint': True,
'iqfrc_smooth': True,
'island_dofadr': True,
'island_efcadr': True,
'island_idofadr': True,
'island_iefcadr': True,
'island_ne': True,
'island_nefc': True,
'island_nf': True,
@@ -1282,11 +1391,12 @@ _BATCH_DIM = {
'nacon': False,
'naconmax': False,
'nbody_awake': True,
'ncdof': True,
'ncollision': False,
'ne': True,
'nefc': True,
'nf': True,
'nidof': False,
'nidof': True,
'nisland': True,
'njmax': False,
'njmax_nnz': False,
@@ -1294,7 +1404,10 @@ _BATCH_DIM = {
'nl': True,
'ntree_awake': True,
'nv_awake': True,
'nvmax': False,
'nvmax_pad': False,
'nworld': False,
'overflow': True,
'qLD': True,
'qLDiagInv': True,
'qLU': True,
@@ -1330,6 +1443,7 @@ _BATCH_DIM = {
'tree_asleep': True,
'tree_awake': True,
'tree_island': True,
'userdata': True,
'wrap_obj': True,
'wrap_xpos': True,
'xanchor': True,
@@ -1353,6 +1467,7 @@ _BATCH_DIM = {
'M_fullm_upper_elemid': False,
'M_fullm_upper_i': False,
'M_fullm_upper_j': False,
'M_hinit_i': False,
'M_mulm_col': False,
'M_mulm_madr': False,
'M_mulm_rowadr': False,
@@ -1390,8 +1505,8 @@ _BATCH_DIM = {
'block_dim__cholesky_solve': False,
'block_dim__contact_jac_tiled': False,
'block_dim__contact_sort': False,
'block_dim__convex_ccd': False,
'block_dim__energy_vel_kinetic': False,
'block_dim__euler_dense': False,
'block_dim__linesearch_iterative': False,
'block_dim__qderiv_actuator_dense': False,
'block_dim__ray': False,
@@ -1412,7 +1527,9 @@ _BATCH_DIM = {
'body_contype': False,
'body_dofadr': False,
'body_dofnum': False,
'body_fluid_box_adr': False,
'body_fluid_ellipsoid': False,
'body_fluid_ellipsoid_adr': False,
'body_geomadr': False,
'body_geomnum': False,
'body_gravcomp': True,
@@ -1493,14 +1610,23 @@ _BATCH_DIM = {
'flex_elemdataadr': False,
'flex_elemedge': False,
'flex_elemedgeadr': False,
'flex_elemflexid': False,
'flex_elemnum': False,
'flex_evpair': False,
'flex_evpairadr': False,
'flex_evpairflexid': False,
'flex_evpairnum': False,
'flex_friction': False,
'flex_gap': False,
'flex_internal': False,
'flex_margin': False,
'flex_priority': False,
'flex_radius': False,
'flex_selfcollide': False,
'flex_shell': False,
'flex_shelladr': False,
'flex_shelldataadr': False,
'flex_shellflexid': False,
'flex_shellnum': False,
'flex_solimp': False,
'flex_solmix': False,
@@ -1517,6 +1643,9 @@ _BATCH_DIM = {
'flexedge_J_rownnz': False,
'flexedge_invweight0': False,
'flexedge_length0': False,
'flexelem_geom_pair_filtered': False,
'flexshell_geom_pair_filtered': False,
'flexvert_geom_pair_filtered': False,
'geom_aabb': True,
'geom_bodyid': False,
'geom_conaffinity': False,
@@ -1541,6 +1670,7 @@ _BATCH_DIM = {
'geom_solmix': True,
'geom_solref': True,
'geom_type': False,
'has_flex_selfcollide': False,
'has_fluid': False,
'has_sdf_geom': False,
'hfield_adr': False,
@@ -1594,6 +1724,7 @@ _BATCH_DIM = {
'mat_specular': True,
'mat_texid': True,
'mat_texrepeat': True,
'max_flex_dim': False,
'max_ten_J_rownnz': False,
'mesh_face': False,
'mesh_faceadr': False,
@@ -1612,6 +1743,7 @@ _BATCH_DIM = {
'mesh_polyvert': False,
'mesh_polyvertadr': False,
'mesh_polyvertnum': False,
'mesh_pos': False,
'mesh_quat': False,
'mesh_vert': False,
'mesh_vertadr': False,
@@ -1636,11 +1768,11 @@ _BATCH_DIM = {
'nflexelem': False,
'nflexelemdata': False,
'nflexelemedge': False,
'nflexevpair': False,
'nflexshelldata': False,
'nflexstiffness': False,
'nflexvert': False,
'ngeom': False,
'ngravcomp': False,
'nhfield': False,
'nhfielddata': False,
'nhistory': False,
@@ -1674,6 +1806,7 @@ _BATCH_DIM = {
'ntendon': False,
'ntree': False,
'nu': False,
'nuserdata': False,
'nv': False,
'nv_pad': False,
'nwrap': False,
@@ -1707,6 +1840,7 @@ _BATCH_DIM = {
'opt__timestep': True,
'opt__tolerance': True,
'opt__viscosity': True,
'opt__warn_overflow': False,
'opt__wind': True,
'pair_dim': False,
'pair_friction': True,
@@ -1722,7 +1856,15 @@ _BATCH_DIM = {
'qD_fullm_i': False,
'qD_fullm_j': False,
'qLD_all_updates': False,
'qLD_block_adr': False,
'qLD_block_total': False,
'qLD_dof_dense': False,
'qLD_dof_simple': False,
'qLD_has_dense': False,
'qLD_has_simple': False,
'qLD_has_sparse': False,
'qLD_level_offsets': False,
'qLD_simple_dofs': False,
'qLD_updates': False,
'qpos0': True,
'qpos_spring': True,
@@ -1832,6 +1974,7 @@ _BATCH_DIM = {
'timestep': True,
'tolerance': True,
'viscosity': True,
'warn_overflow': False,
'wind': True,
},
'Statistic': {'meaninertia': True},