diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index fc1a6595..13df4f7f 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -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}' diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index ff014402..b7ca1257 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -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.""" diff --git a/mjx/mujoco/mjx/codegen/generate_warp_shim.py b/mjx/mujoco/mjx/codegen/generate_warp_shim.py index 2596c725..a513b88c 100644 --- a/mjx/mujoco/mjx/codegen/generate_warp_shim.py +++ b/mjx/mujoco/mjx/codegen/generate_warp_shim.py @@ -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: diff --git a/mjx/mujoco/mjx/codegen/generate_warp_types.py b/mjx/mujoco/mjx/codegen/generate_warp_types.py index 6c9e3c05..48e090e7 100644 --- a/mjx/mujoco/mjx/codegen/generate_warp_types.py +++ b/mjx/mujoco/mjx/codegen/generate_warp_types.py @@ -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') diff --git a/mjx/mujoco/mjx/codegen/trace.py b/mjx/mujoco/mjx/codegen/trace.py index 5ba8359d..b62896a3 100644 --- a/mjx/mujoco/mjx/codegen/trace.py +++ b/mjx/mujoco/mjx/codegen/trace.py @@ -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) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py index ec62bc20..ebfc44a6 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py @@ -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 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/cli.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/cli.py index 425b54bb..b3fe92d7 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/cli.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/cli.py @@ -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(): diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py index 9bf7ce5e..19d97952 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py @@ -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], ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_core.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_core.py index 35bb0f27..dbf0fcec 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_core.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_core.py @@ -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 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py index e3199833..9b53f938 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py @@ -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 `` 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) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_flex.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_flex.py index af166dd8..0e2e3117 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_flex.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_flex.py @@ -13,22 +13,452 @@ # limitations under the License. """Flex collision detection (geom vs flex triangles).""" +from typing import Tuple + import warp as wp from mujoco.mjx.third_party.mujoco_warp._src import collision_primitive_core +from mujoco.mjx.third_party.mujoco_warp._src.collision_core import Geom +from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import ccd from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame +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_MAXVAL from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINMU from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import ContactType from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import OverflowType from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope wp.set_module_options({"enable_backward": False}) +@wp.func +def _flex_element_aabb_filter( + # In: + box1_min: wp.vec3, + box1_max: wp.vec3, + box2_min: wp.vec3, + box2_max: wp.vec3, +): + """Return True if the two AABBs do NOT intersect (discard pair).""" + if box1_max[0] < box2_min[0] or box1_min[0] > box2_max[0]: + return True + if box1_max[1] < box2_min[1] or box1_min[1] > box2_max[1]: + return True + if box1_max[2] < box2_min[2] or box1_min[2] > box2_max[2]: + return True + return False + + +@wp.kernel +def _flex_broadphase_bounds( + # Model: + flex_margin: wp.array[float], + flex_gap: wp.array[float], + flex_vertadr: wp.array[int], + flex_vertnum: wp.array[int], + flex_radius: wp.array[float], + # Data in: + flexvert_xpos_in: wp.array2d[wp.vec3], + # Data out: + flex_aabb_min_out: wp.array2d[wp.vec3], + flex_aabb_max_out: wp.array2d[wp.vec3], +): + worldid, flexid = wp.tid() + + start = flex_vertadr[flexid] + num = flex_vertnum[flexid] + if num == 0: + return + + min_bound = wp.vec3(MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL) + max_bound = wp.vec3(-MJ_MAXVAL, -MJ_MAXVAL, -MJ_MAXVAL) + + for i in range(num): + pos = flexvert_xpos_in[worldid, start + i] + min_bound = wp.min(min_bound, pos) + max_bound = wp.max(max_bound, pos) + + margin = flex_margin[flexid] + flex_gap[flexid] + bound = flex_radius[flexid] + margin + inflate = wp.vec3(bound, bound, bound) + + flex_aabb_min_out[worldid, flexid] = min_bound - inflate + flex_aabb_max_out[worldid, flexid] = max_bound + inflate + + +@wp.func +def _flex_triangle_geom_broadphase( + # Model: + ngeom: int, + opt_warn_overflow: bool, + geom_type: wp.array[int], + geom_aabb: wp.array3d[wp.vec3], + geom_margin: wp.array2d[float], + # Data in: + geom_xpos_in: wp.array2d[wp.vec3], + geom_xmat_in: wp.array2d[wp.mat33], + naconmax_in: int, + # In: + worldid: int, + element_or_shell_id: int, + flexid: int, + geomid: int, + t1: wp.vec3, + t2: wp.vec3, + t3: wp.vec3, + tri_radius: float, + tri_margin: float, + flex_aabb_min_val: wp.vec3, + flex_aabb_max_val: wp.vec3, + # Data out: + ncollision_out: wp.array[int], + # Data out: + overflow_out: wp.array[int], + # Out: + collision_pair_out: wp.array[wp.vec2i], + collision_worldid_out: wp.array[int], +): + gtype = geom_type[geomid] + if ( + gtype != int(GeomType.SPHERE) + and gtype != int(GeomType.CAPSULE) + and gtype != int(GeomType.BOX) + and gtype != int(GeomType.CYLINDER) + and gtype != int(GeomType.MESH) + ): + return + + geom_margin_val = geom_margin[worldid % geom_margin.shape[0], geomid] + margin = geom_margin_val + tri_margin + + aabb_id = worldid % geom_aabb.shape[0] + geom_center_local = geom_aabb[aabb_id, geomid, 0] + geom_half_size_local = geom_aabb[aabb_id, geomid, 1] + + geom_pos = geom_xpos_in[worldid, geomid] + geom_rot = geom_xmat_in[worldid, geomid] + + # Stage 1 Filter: Coarse flex AABB vs Geom world AABB check + # Transform center to global frame + geom_center_global = geom_rot @ geom_center_local + geom_pos + + # Project local half-size onto world axes using absolute rotation matrix entries + geom_half_size_global = wp.vec3( + wp.abs(geom_rot[0, 0]) * geom_half_size_local[0] + + wp.abs(geom_rot[0, 1]) * geom_half_size_local[1] + + wp.abs(geom_rot[0, 2]) * geom_half_size_local[2], + wp.abs(geom_rot[1, 0]) * geom_half_size_local[0] + + wp.abs(geom_rot[1, 1]) * geom_half_size_local[1] + + wp.abs(geom_rot[1, 2]) * geom_half_size_local[2], + wp.abs(geom_rot[2, 0]) * geom_half_size_local[0] + + wp.abs(geom_rot[2, 1]) * geom_half_size_local[1] + + wp.abs(geom_rot[2, 2]) * geom_half_size_local[2], + ) + + inflate = wp.vec3(margin, margin, margin) + geom_box_min = geom_center_global - geom_half_size_global - inflate + geom_box_max = geom_center_global + geom_half_size_global + inflate + + if _flex_element_aabb_filter(geom_box_min, geom_box_max, flex_aabb_min_val, flex_aabb_max_val): + return + + # Stage 2 Filter: Element AABB vs Geom world AABB check + tri_min = wp.min(t1, wp.min(t2, t3)) - wp.vec3(tri_radius, tri_radius, tri_radius) + tri_max = wp.max(t1, wp.max(t2, t3)) + wp.vec3(tri_radius, tri_radius, tri_radius) + + if _flex_element_aabb_filter(geom_box_min, geom_box_max, tri_min, tri_max): + return + + # Stage 3 Filter: Project Geom onto triangle plane normal + normal = wp.normalize(wp.cross(t2 - t1, t3 - t1)) + signed_dist = wp.dot(geom_pos - t1, normal) + + r_extent = float(0.0) + if gtype == int(GeomType.SPHERE): + r_extent = geom_half_size_local[0] + elif gtype == int(GeomType.CAPSULE): + r_extent = geom_half_size_local[0] + geom_half_size_local[1] + elif gtype == int(GeomType.CYLINDER): + r_extent = wp.sqrt(geom_half_size_local[0] * geom_half_size_local[0] + geom_half_size_local[1] * geom_half_size_local[1]) + elif gtype == int(GeomType.BOX): + r_extent = wp.length(geom_half_size_local) + elif gtype == int(GeomType.MESH): + r_extent = wp.length(geom_half_size_local) + + if wp.abs(signed_dist) > r_extent + margin + tri_radius: + return + + # Overlap found! Save candidate. + idx = wp.atomic_add(ncollision_out, 0, 1) + if idx >= naconmax_in: + if opt_warn_overflow: + wp.printf("Collision buffer overflow in flex broadphase - please increase naconmax to %u\n", idx + 1) + wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.BROADPHASE)) + return + collision_pair_out[idx] = wp.vec2i(element_or_shell_id, geomid) + collision_worldid_out[idx] = worldid + + +@wp.kernel +def _flex_broadphase_unified( + # Model: + ngeom: int, + nflex: int, + opt_warn_overflow: bool, + geom_type: wp.array[int], + geom_size: wp.array2d[wp.vec3], + geom_aabb: wp.array3d[wp.vec3], + geom_rbound: wp.array2d[float], + geom_margin: wp.array2d[float], + flex_margin: wp.array[float], + flex_dim: wp.array[int], + flex_vertadr: wp.array[int], + flex_radius: wp.array[float], + # Data in: + geom_xpos_in: wp.array2d[wp.vec3], + geom_xmat_in: wp.array2d[wp.mat33], + flexvert_xpos_in: wp.array2d[wp.vec3], + naconmax_in: int, + flex_aabb_min_in: wp.array2d[wp.vec3], + flex_aabb_max_in: wp.array2d[wp.vec3], + # In: + triadr: wp.array[int], + tridataadr: wp.array[int], + tri: wp.array[int], + pairs_filtered: wp.array[wp.vec2i], + triflexid: wp.array[int], + # Data out: + ncollision_out: wp.array[int], + # Data out: + overflow_out: wp.array[int], + # Out: + collision_pair_out: wp.array[wp.vec2i], + collision_worldid_out: wp.array[int], +): + worldid, pairid = wp.tid() + + pair = pairs_filtered[pairid] + tri_id = pair[0] + geomid = pair[1] + + flexid = triflexid[tri_id] + + vert_adr = flex_vertadr[flexid] + tri_radius = flex_radius[flexid] + tri_margin = flex_margin[flexid] + + tri_data_idx = tridataadr[flexid] + (tri_id - triadr[flexid]) * 3 + v0_local = tri[tri_data_idx] + v1_local = tri[tri_data_idx + 1] + v2_local = tri[tri_data_idx + 2] + + t1 = flexvert_xpos_in[worldid, vert_adr + v0_local] + t2 = flexvert_xpos_in[worldid, vert_adr + v1_local] + t3 = flexvert_xpos_in[worldid, vert_adr + v2_local] + + _flex_triangle_geom_broadphase( + ngeom, + opt_warn_overflow, + geom_type, + geom_aabb, + geom_margin, + geom_xpos_in, + geom_xmat_in, + naconmax_in, + worldid, + tri_id, + flexid, + geomid, + t1, + t2, + t3, + tri_radius, + tri_margin, + flex_aabb_min_in[worldid, flexid], + flex_aabb_max_in[worldid, flexid], + # Data out: + ncollision_out, + overflow_out, + collision_pair_out, + collision_worldid_out, + ) + + +@wp.kernel +def _flex_broadphase_plane( + # Model: + ngeom: int, + opt_warn_overflow: bool, + geom_type: wp.array[int], + geom_margin: wp.array2d[float], + flex_margin: wp.array[float], + flex_vertadr: wp.array[int], + flex_radius: wp.array[float], + flexvert_geom_pair_filtered: wp.array[wp.vec2i], + flex_vertflexid: wp.array[int], + # Data in: + geom_xpos_in: wp.array2d[wp.vec3], + geom_xmat_in: wp.array2d[wp.mat33], + flexvert_xpos_in: wp.array2d[wp.vec3], + naconmax_in: int, + flex_aabb_min_in: wp.array2d[wp.vec3], + flex_aabb_max_in: wp.array2d[wp.vec3], + # Data out: + ncollision_out: wp.array[int], + # Data out: + overflow_out: wp.array[int], + # Out: + collision_pair_out: wp.array[wp.vec2i], + collision_worldid_out: wp.array[int], +): + worldid, pairid = wp.tid() + + pair = flexvert_geom_pair_filtered[pairid] + vertid = pair[0] + geomid = pair[1] + + flexid = flex_vertflexid[vertid] + radius = flex_radius[flexid] + flex_margin_val = flex_margin[flexid] + + vert = flexvert_xpos_in[worldid, vertid] + + flex_aabb_min = flex_aabb_min_in[worldid, flexid] + flex_aabb_max = flex_aabb_max_in[worldid, flexid] + + gtype = geom_type[geomid] + if gtype != int(GeomType.PLANE): + return + + margin = geom_margin[worldid % geom_margin.shape[0], geomid] + flex_margin_val + geom_pos = geom_xpos_in[worldid, geomid] + geom_rot = geom_xmat_in[worldid, geomid] + plane_normal = wp.vec3(geom_rot[0, 2], geom_rot[1, 2], geom_rot[2, 2]) + + # Stage 1 filter: Bounding box of flex vs plane + flex_center = 0.5 * (flex_aabb_min + flex_aabb_max) + flex_half_size = 0.5 * (flex_aabb_max - flex_aabb_min) + + proj_half = ( + wp.abs(flex_half_size[0] * plane_normal[0]) + + wp.abs(flex_half_size[1] * plane_normal[1]) + + wp.abs(flex_half_size[2] * plane_normal[2]) + ) + + diff_center = flex_center - geom_pos + dist_center = wp.dot(diff_center, plane_normal) + if dist_center - proj_half > margin: + return + + diff = vert - geom_pos + signed_dist = wp.dot(diff, plane_normal) + dist = signed_dist - radius + + if dist < margin: + # Append Candidate to Context + idx = wp.atomic_add(ncollision_out, 0, 1) + if idx >= naconmax_in: + if opt_warn_overflow: + wp.printf("Collision buffer overflow in flex plane broadphase - please increase naconmax to %u\n", idx + 1) + wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.BROADPHASE)) + return + collision_pair_out[idx] = wp.vec2i(vertid, geomid) + collision_worldid_out[idx] = worldid + + +@wp.func +def _write_flex_contact( + # Model: + geom_condim: wp.array[int], + geom_priority: wp.array[int], + geom_solmix: wp.array2d[float], + geom_solref: wp.array2d[wp.vec2], + geom_solimp: wp.array2d[vec5], + geom_friction: wp.array2d[wp.vec3], + geom_gap: wp.array2d[float], + flex_condim: wp.array[int], + flex_priority: wp.array[int], + flex_solmix: wp.array[float], + flex_solref: wp.array[wp.vec2], + flex_solimp: wp.array[vec5], + flex_friction: wp.array[wp.vec3], + flex_margin: wp.array[float], + flex_gap: wp.array[float], + # Data in: + naconmax_in: int, + # In: + collisionid: int, + dist: float, + pos: wp.vec3, + normal: wp.vec3, + margin: float, + geomid: int, + flexid: int, + elemid: int, + vertid: int, + worldid: int, + # Data out: + contact_dist_out: wp.array[float], + contact_pos_out: wp.array[wp.vec3], + contact_frame_out: wp.array[wp.mat33], + contact_includemargin_out: wp.array[float], + contact_friction_out: wp.array[vec5], + contact_solref_out: wp.array[wp.vec2], + contact_solreffriction_out: wp.array[wp.vec2], + contact_solimp_out: wp.array[vec5], + contact_dim_out: wp.array[int], + contact_geom_out: wp.array[wp.vec2i], + contact_flex_out: wp.array[wp.vec2i], + contact_vert_out: wp.array[wp.vec2i], + contact_worldid_out: wp.array[int], + contact_type_out: wp.array[int], + contact_geomcollisionid_out: wp.array[int], + nacon_out: wp.array[int], +): + condim, gap, solref, solimp, friction = _mix_flex_contact_params( + geom_condim[geomid], + geom_priority[geomid], + geom_solmix[worldid % geom_solmix.shape[0], geomid], + geom_solref[worldid % geom_solref.shape[0], geomid], + geom_solimp[worldid % geom_solimp.shape[0], geomid], + geom_friction[worldid % geom_friction.shape[0], geomid], + geom_gap[worldid % geom_gap.shape[0], geomid], + flex_condim[flexid], + flex_priority[flexid], + flex_solmix[flexid], + flex_solref[flexid], + flex_solimp[flexid], + flex_friction[flexid], + flex_gap[flexid], + ) + + c_idx = wp.atomic_add(nacon_out, 0, 1) + if c_idx < naconmax_in: + frame = make_frame(normal) + + contact_dist_out[c_idx] = dist + contact_pos_out[c_idx] = pos + contact_frame_out[c_idx] = frame + contact_includemargin_out[c_idx] = margin + contact_friction_out[c_idx] = friction + contact_solref_out[c_idx] = solref + contact_solreffriction_out[c_idx] = solref # Same for flex contacts + contact_solimp_out[c_idx] = solimp + contact_dim_out[c_idx] = condim + contact_geom_out[c_idx] = wp.vec2i(geomid, -1) + contact_flex_out[c_idx] = wp.vec2i(flexid, elemid) + contact_vert_out[c_idx] = wp.vec2i(vertid, -1) + contact_worldid_out[c_idx] = worldid + contact_type_out[c_idx] = ContactType.CONSTRAINT + contact_geomcollisionid_out[c_idx] = collisionid + + # TODO(team): generalize into a shared contact parameter mixing function # (mj_contactParam) that works for both geom-geom and geom-flex contacts. @wp.func @@ -117,69 +547,71 @@ def _mix_flex_contact_params( @wp.func -def _write_flex_contact( - # Data in: - naconmax_in: int, +def _write_candidate_contact( # In: + max_candidates: int, dist: float, pos: wp.vec3, - frame: wp.mat33, - margin: float, - condim: int, - friction: vec5, - solref: wp.vec2, - solimp: vec5, + nrm: wp.vec3, geom: int, flexid: int, + elemid: int, vertid: int, worldid: int, # Data out: - contact_dist_out: wp.array[float], - contact_pos_out: wp.array[wp.vec3], - contact_frame_out: wp.array[wp.mat33], - contact_includemargin_out: wp.array[float], - contact_friction_out: wp.array[vec5], - contact_solref_out: wp.array[wp.vec2], - contact_solreffriction_out: wp.array[wp.vec2], - contact_solimp_out: wp.array[vec5], - contact_dim_out: wp.array[int], - contact_geom_out: wp.array[wp.vec2i], - contact_flex_out: wp.array[wp.vec2i], - contact_vert_out: wp.array[wp.vec2i], - contact_worldid_out: wp.array[int], - contact_type_out: wp.array[int], - contact_geomcollisionid_out: wp.array[int], - nacon_out: wp.array[int], + overflow_out: wp.array[int], + # Out: + cand_dist_out: wp.array[float], + cand_pos_out: wp.array[wp.vec3], + cand_nrm_out: wp.array[wp.vec3], + cand_geom_out: wp.array[wp.vec2i], + cand_flex_out: wp.array[wp.vec2i], + cand_elem_out: wp.array[wp.vec2i], + cand_vert_out: wp.array[wp.vec2i], + cand_worldid_out: wp.array[int], + cand_type_out: wp.array[int], + cand_geomcollisionid_out: wp.array[int], + ncand_out: wp.array[int], ): - if dist >= margin or dist >= MJ_MAXVAL: + if dist >= MJ_MAXVAL: return - id_ = wp.atomic_add(nacon_out, 0, 1) - if id_ >= naconmax_in: + candid = wp.atomic_add(ncand_out, 0, 1) + if candid >= max_candidates: + wp.printf( + "flex candidate overflow - please increase naconmax to %u\n", + candid + 1, + ) + wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.BROADPHASE)) return - contact_dist_out[id_] = dist - contact_pos_out[id_] = pos - contact_frame_out[id_] = frame - contact_includemargin_out[id_] = margin - contact_friction_out[id_] = friction - contact_solref_out[id_] = solref - contact_solreffriction_out[id_] = wp.vec2(0.0, 0.0) - contact_solimp_out[id_] = solimp - contact_dim_out[id_] = condim - contact_geom_out[id_] = wp.vec2i(geom, -1) - contact_flex_out[id_] = wp.vec2i(-1, flexid) - contact_vert_out[id_] = wp.vec2i(-1, vertid) - contact_worldid_out[id_] = worldid - contact_type_out[id_] = 1 - contact_geomcollisionid_out[id_] = 0 + cand_dist_out[candid] = dist + cand_pos_out[candid] = pos + cand_nrm_out[candid] = nrm + if geom >= 0: + cand_geom_out[candid] = wp.vec2i(geom, -1) + cand_flex_out[candid] = wp.vec2i(-1, flexid) + cand_elem_out[candid] = wp.vec2i(-1, elemid) + cand_vert_out[candid] = wp.vec2i(-1, vertid) + elif geom == -2: + cand_geom_out[candid] = wp.vec2i(-1, -1) + cand_flex_out[candid] = wp.vec2i(flexid, flexid) + cand_elem_out[candid] = wp.vec2i(elemid, vertid) + cand_vert_out[candid] = wp.vec2i(-1, -1) + else: + cand_geom_out[candid] = wp.vec2i(-1, -1) + cand_flex_out[candid] = wp.vec2i(flexid, flexid) + cand_elem_out[candid] = wp.vec2i(-1, elemid) + cand_vert_out[candid] = wp.vec2i(vertid, -1) + cand_worldid_out[candid] = worldid + cand_type_out[candid] = 1 + cand_geomcollisionid_out[candid] = 0 @wp.func -def _collide_geom_triangle( - # Data in: - naconmax_in: int, +def _collide_geom_triangle_detect( # In: + max_candidates: int, gtype: int, pos: wp.vec3, rot: wp.mat33, @@ -189,66 +621,52 @@ def _collide_geom_triangle( t3: wp.vec3, tri_radius: float, margin: float, - condim: int, - friction: vec5, - solref: wp.vec2, - solimp: vec5, geomid: int, flexid: int, + elemid: int, vertex_id: int, worldid: int, # Data out: - contact_dist_out: wp.array[float], - contact_pos_out: wp.array[wp.vec3], - contact_frame_out: wp.array[wp.mat33], - contact_includemargin_out: wp.array[float], - contact_friction_out: wp.array[vec5], - contact_solref_out: wp.array[wp.vec2], - contact_solreffriction_out: wp.array[wp.vec2], - contact_solimp_out: wp.array[vec5], - contact_dim_out: wp.array[int], - contact_geom_out: wp.array[wp.vec2i], - contact_flex_out: wp.array[wp.vec2i], - contact_vert_out: wp.array[wp.vec2i], - contact_worldid_out: wp.array[int], - contact_type_out: wp.array[int], - contact_geomcollisionid_out: wp.array[int], - nacon_out: wp.array[int], + overflow_out: wp.array[int], + # Out: + cand_dist_out: wp.array[float], + cand_pos_out: wp.array[wp.vec3], + cand_nrm_out: wp.array[wp.vec3], + cand_geom_out: wp.array[wp.vec2i], + cand_flex_out: wp.array[wp.vec2i], + cand_elem_out: wp.array[wp.vec2i], + cand_vert_out: wp.array[wp.vec2i], + cand_worldid_out: wp.array[int], + cand_type_out: wp.array[int], + cand_geomcollisionid_out: wp.array[int], + ncand_out: wp.array[int], ): if gtype == int(GeomType.SPHERE): sphere_radius = size_val[0] dist, contact_pos, nrm = collision_primitive_core.sphere_triangle(pos, sphere_radius, t1, t2, t3, tri_radius) if dist < margin: - _write_flex_contact( - naconmax_in, + _write_candidate_contact( + max_candidates, dist, contact_pos, - make_frame(nrm), - margin, - condim, - friction, - solref, - solimp, + nrm, geomid, flexid, + elemid, vertex_id, worldid, - 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_flex_out, - contact_vert_out, - contact_worldid_out, - contact_type_out, - contact_geomcollisionid_out, - nacon_out, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, ) return @@ -278,73 +696,267 @@ def _collide_geom_triangle( if dists[0] < margin: p1 = wp.vec3(poss[0, 0], poss[0, 1], poss[0, 2]) n1 = wp.vec3(nrms[0, 0], nrms[0, 1], nrms[0, 2]) - _write_flex_contact( - naconmax_in, + _write_candidate_contact( + max_candidates, dists[0], p1, - make_frame(n1), - margin, - condim, - friction, - solref, - solimp, + n1, geomid, flexid, + elemid, vertex_id, worldid, - 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_flex_out, - contact_vert_out, - contact_worldid_out, - contact_type_out, - contact_geomcollisionid_out, - nacon_out, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, ) if dists[1] < margin: p2 = wp.vec3(poss[1, 0], poss[1, 1], poss[1, 2]) n2 = wp.vec3(nrms[1, 0], nrms[1, 1], nrms[1, 2]) - _write_flex_contact( - naconmax_in, + _write_candidate_contact( + max_candidates, dists[1], p2, - make_frame(n2), - margin, - condim, - friction, - solref, - solimp, + n2, geomid, flexid, + elemid, vertex_id, worldid, - 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_flex_out, - contact_vert_out, - contact_worldid_out, - contact_type_out, - contact_geomcollisionid_out, - nacon_out, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, ) +@wp.func +def _collide_mesh_triangle( + # Model: + mesh_vertadr: wp.array[int], + mesh_vertnum: wp.array[int], + mesh_graphadr: wp.array[int], + mesh_vert: wp.array[wp.vec3], + mesh_graph: wp.array[int], + mesh_pos: wp.array[wp.vec3], + mesh_polynormal: wp.array[wp.vec3], + mesh_polyvertadr: wp.array[int], + mesh_polyvert: wp.array[int], + mesh_polymapadr: wp.array[int], + mesh_polymapnum: wp.array[int], + mesh_polymap: wp.array[int], + # In: + max_candidates: int, + mesh_geom_pos: wp.vec3, + geom_rot: wp.mat33, + geom_size_val: wp.vec3, + t1: wp.vec3, + t2: wp.vec3, + t3: wp.vec3, + tri_radius: float, + margin: float, + geomid: int, + flexid: int, + elemid: int, + v0_local: int, + v1_local: int, + v2_local: int, + worldid: int, + did: int, + epa_vert: wp.array[wp.vec3], + epa_vert_index: wp.array[int], + epa_face: wp.array[int], + epa_pr: wp.array[wp.vec3], + epa_norm2: wp.array[float], + epa_horizon: wp.array[int], + tolerance: float, + ccd_iterations: int, + # Data out: + overflow_out: wp.array[int], + # Out: + cand_dist_out: wp.array[float], + cand_pos_out: wp.array[wp.vec3], + cand_nrm_out: wp.array[wp.vec3], + cand_geom_out: wp.array[wp.vec2i], + cand_flex_out: wp.array[wp.vec2i], + cand_elem_out: wp.array[wp.vec2i], + cand_vert_out: wp.array[wp.vec2i], + cand_worldid_out: wp.array[int], + cand_type_out: wp.array[int], + cand_geomcollisionid_out: wp.array[int], + ncand_out: wp.array[int], +): + # Construct Mesh Geom (geom1) + geom1 = Geom() + geom1.pos = mesh_geom_pos + geom1.rot = geom_rot + geom1.size = geom_size_val + geom1.margin = 0.0 + geom1.index = -1 + + geom1.vertadr = wp.where(did >= 0, mesh_vertadr[did], -1) + geom1.vertnum = wp.where(did >= 0, mesh_vertnum[did], -1) + geom1.graphadr = wp.where(did >= 0, mesh_graphadr[did], -1) + geom1.vert = mesh_vert + geom1.graph = mesh_graph + + # Construct Triangle Geom (geom2) + geom2 = Geom() + geom2.pos = wp.vec3(0.0, 0.0, 0.0) + geom2.rot = wp.mat33(t1[0], t1[1], t1[2], t2[0], t2[1], t2[2], t3[0], t3[1], t3[2]) + geom2.margin = 0.0 + geom2.index = -1 + + centroid = (t1 + t2 + t3) * (1.0 / 3.0) + r_geom = wp.length(geom_size_val) + d1 = wp.length(t1 - centroid) + d2 = wp.length(t2 - centroid) + d3 = wp.length(t3 - centroid) + r_tri = wp.max(d1, wp.max(d2, d3)) + + geom_center = mesh_geom_pos + if did >= 0: + geom_center = mesh_geom_pos + geom_rot @ mesh_pos[did] + + if wp.length(centroid - geom_center) <= r_geom + r_tri + margin + tri_radius + 0.04: + dist, ncontact, w1, w2, idx = ccd( + tolerance, + margin + tri_radius, + ccd_iterations, + ccd_iterations, + geom1, + geom2, + int(GeomType.MESH), + int(GeomType.TRIANGLE), + mesh_geom_pos, + centroid, + epa_vert, + epa_vert_index, + epa_face, + epa_pr, + epa_norm2, + epa_horizon, + ) + + if ncontact > 0 and dist < margin + tri_radius: + if not _inside_triangle(w2, t1, t2, t3, 0.2): + return + + if dist < 0.0: + gjk_normal = wp.normalize(w1 - w2) + else: + gjk_normal = wp.normalize(w2 - w1) + + # TODO(thowell): remove after resolving contact normal issue + best_normal = gjk_normal + if idx >= 0: + # Extract GJK/EPA support vertex on the mesh + f_verts = wp.vec3i(epa_face[idx] & 0x3FF, (epa_face[idx] >> 10) & 0x3FF, (epa_face[idx] >> 20) & 0x3FF) + sv0 = epa_vert_index[2 * f_verts[0]] + + # Transform w1 to mesh local frame + w1_local = wp.transpose(geom_rot) @ (w1 - mesh_geom_pos) + + min_plane_dist = float(1e10) + best_normal_local = wp.transpose(geom_rot) @ gjk_normal + best_poly_idx = int(-1) + + v_offset = mesh_vertadr[did] + v_global_idx = v_offset + sv0 + polymap_start = mesh_polymapadr[v_global_idx] + npolygons = mesh_polymapnum[v_global_idx] + + for k in range(npolygons): + poly_idx = mesh_polymap[polymap_start + k] + normal_local = mesh_polynormal[poly_idx] + + v0_local_idx = mesh_polyvert[mesh_polyvertadr[poly_idx]] + v0_mesh_local = mesh_vert[v_offset + v0_local_idx] + + dist_to_plane = wp.abs(wp.dot(w1_local - v0_mesh_local, normal_local)) + if dist_to_plane < min_plane_dist: + min_plane_dist = dist_to_plane + best_normal_local = normal_local + best_poly_idx = poly_idx + + if best_poly_idx >= 0: + vert_start = mesh_polyvertadr[best_poly_idx] + v0_idx = mesh_polyvert[vert_start] + v1_idx = mesh_polyvert[vert_start + 1] + v2_idx = mesh_polyvert[vert_start + 2] + + v0_mesh_local = mesh_vert[v_offset + v0_idx] + v1_mesh_local = mesh_vert[v_offset + v1_idx] + v2_mesh_local = mesh_vert[v_offset + v2_idx] + + # Filter out contacts that are too far from the plane or fall outside the face + if min_plane_dist > 0.005 or not _inside_triangle(w1_local, v0_mesh_local, v1_mesh_local, v2_mesh_local, 0.05): + return + + best_normal = wp.normalize(geom_rot @ best_normal_local) + + normal = wp.where(wp.dot(best_normal, gjk_normal) >= 0.0, best_normal, -best_normal) + contact_pos = 0.5 * (w1 + w2) + + dist_v0 = wp.dot(t1 - w1, normal) - tri_radius + dist_v1 = wp.dot(t2 - w1, normal) - tri_radius + dist_v2 = wp.dot(t3 - w1, normal) - tri_radius + + min_dist = wp.min(dist_v0, wp.min(dist_v1, dist_v2)) + if min_dist < margin: + deepest_vert = v0_local + pos = t1 - normal * (tri_radius + 0.5 * dist_v0) + if dist_v1 < dist_v0 and dist_v1 < dist_v2: + deepest_vert = v1_local + pos = t2 - normal * (tri_radius + 0.5 * dist_v1) + elif dist_v2 < dist_v0 and dist_v2 < dist_v1: + deepest_vert = v2_local + pos = t3 - normal * (tri_radius + 0.5 * dist_v2) + + _write_candidate_contact( + max_candidates, + min_dist, + pos, + normal, + geomid, + flexid, + elemid, + -1, + worldid, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, + ) + + return + + @wp.kernel def _flex_plane_narrowphase( # Model: @@ -376,6 +988,10 @@ def _flex_plane_narrowphase( flexvert_xpos_in: wp.array2d[wp.vec3], nworld_in: int, naconmax_in: int, + ncollision_in: wp.array[int], + # In: + collision_pair_in: wp.array[wp.vec2i], + collision_worldid_in: wp.array[int], # Data out: contact_dist_out: wp.array[float], contact_pos_out: wp.array[wp.vec3], @@ -394,7 +1010,19 @@ def _flex_plane_narrowphase( contact_geomcollisionid_out: wp.array[int], nacon_out: wp.array[int], ): - worldid, vertid = wp.tid() + collisionid = wp.tid() + if collisionid >= ncollision_in[0] or collisionid >= collision_pair_in.shape[0]: + return + + pair = collision_pair_in[collisionid] + geomid = pair[1] + + gtype = geom_type[geomid] + if gtype != int(GeomType.PLANE): + return + + vertid = pair[0] + worldid = collision_worldid_in[collisionid] flexid = flex_vertflexid[vertid] radius = flex_radius[flexid] @@ -404,101 +1032,362 @@ def _flex_plane_narrowphase( vert = flexvert_xpos_in[worldid, vertid] - # TODO: Add a broadphase - for geomid in range(ngeom): - gtype = geom_type[geomid] - if gtype != int(GeomType.PLANE): - continue + plane_pos = geom_xpos_in[worldid, geomid] + plane_rot = geom_xmat_in[worldid, geomid] + plane_normal = wp.vec3(plane_rot[0, 2], plane_rot[1, 2], plane_rot[2, 2]) - plane_pos = geom_xpos_in[worldid, geomid] - plane_rot = geom_xmat_in[worldid, geomid] - plane_normal = wp.vec3(plane_rot[0, 2], plane_rot[1, 2], plane_rot[2, 2]) + margin = geom_margin[worldid % geom_margin.shape[0], geomid] + flex_margin_val - margin = geom_margin[worldid % geom_margin.shape[0], geomid] + flex_margin_val + diff = vert - plane_pos + signed_dist = wp.dot(diff, plane_normal) + dist = signed_dist - radius - diff = vert - plane_pos - signed_dist = wp.dot(diff, plane_normal) - dist = signed_dist - radius - - if dist < margin: - condim, gap, solref, solimp, friction = _mix_flex_contact_params( - geom_condim[geomid], - geom_priority[geomid], - geom_solmix[worldid % geom_solmix.shape[0], geomid], - geom_solref[worldid % geom_solref.shape[0], geomid], - geom_solimp[worldid % geom_solimp.shape[0], geomid], - geom_friction[worldid % geom_friction.shape[0], geomid], - geom_gap[worldid % geom_gap.shape[0], geomid], - flex_condim[flexid], - flex_priority[flexid], - flex_solmix[flexid], - flex_solref[flexid], - flex_solimp[flexid], - flex_friction[flexid], - flex_gap[flexid], - ) - - contact_pos = vert - plane_normal * (dist * 0.5 + radius) - _write_flex_contact( - naconmax_in, - dist, - contact_pos, - make_frame(plane_normal), - margin - gap, - condim, - friction, - solref, - solimp, - geomid, - flexid, - local_vertid, - worldid, - 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_flex_out, - contact_vert_out, - contact_worldid_out, - contact_type_out, - contact_geomcollisionid_out, - nacon_out, - ) + if dist < margin: + contact_pos = vert - plane_normal * (dist * 0.5 + radius) + _write_flex_contact( + geom_condim, + geom_priority, + geom_solmix, + geom_solref, + geom_solimp, + geom_friction, + geom_gap, + flex_condim, + flex_priority, + flex_solmix, + flex_solref, + flex_solimp, + flex_friction, + flex_margin, + flex_gap, + naconmax_in, + collisionid, + dist, + contact_pos, + plane_normal, + margin, + geomid, + flexid, + -1, + local_vertid, + worldid, + 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_flex_out, + contact_vert_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.kernel -def _flex_narrowphase_dim2( +def _flex_geom_vertex_narrowphase_detect( # Model: ngeom: int, - nflex: int, + nflexvert: int, geom_type: wp.array[int], geom_contype: wp.array[int], geom_conaffinity: wp.array[int], - geom_condim: wp.array[int], - geom_priority: wp.array[int], - geom_solmix: wp.array2d[float], - geom_solref: wp.array2d[wp.vec2], - geom_solimp: wp.array2d[vec5], geom_size: wp.array2d[wp.vec3], - geom_friction: wp.array2d[wp.vec3], geom_margin: wp.array2d[float], - geom_gap: wp.array2d[float], flex_contype: wp.array[int], flex_conaffinity: wp.array[int], - flex_condim: wp.array[int], - flex_priority: wp.array[int], - flex_solmix: wp.array[float], - flex_solref: wp.array[wp.vec2], - flex_solimp: wp.array[vec5], - flex_friction: wp.array[wp.vec3], flex_margin: wp.array[float], - flex_gap: wp.array[float], + flex_dim: wp.array[int], + flex_vertadr: wp.array[int], + flex_radius: wp.array[float], + flex_vertflexid: wp.array[int], + # Data in: + geom_xpos_in: wp.array2d[wp.vec3], + geom_xmat_in: wp.array2d[wp.mat33], + flexvert_xpos_in: wp.array2d[wp.vec3], + nworld_in: int, + # In: + max_candidates: int, + # Data out: + overflow_out: wp.array[int], + # Out: + cand_dist_out: wp.array[float], + cand_pos_out: wp.array[wp.vec3], + cand_nrm_out: wp.array[wp.vec3], + cand_geom_out: wp.array[wp.vec2i], + cand_flex_out: wp.array[wp.vec2i], + cand_elem_out: wp.array[wp.vec2i], + cand_vert_out: wp.array[wp.vec2i], + cand_worldid_out: wp.array[int], + cand_type_out: wp.array[int], + cand_geomcollisionid_out: wp.array[int], + ncand_out: wp.array[int], +): + worldid, vertid = wp.tid() + + flexid = flex_vertflexid[vertid] + if flex_dim[flexid] >= 2: + return + + radius = flex_radius[flexid] + flex_margin_val = flex_margin[flexid] + local_vertid = vertid - flex_vertadr[flexid] + + v_pos = flexvert_xpos_in[worldid, vertid] + + for geomid in range(ngeom): + gtype = geom_type[geomid] + if ( + gtype != int(GeomType.SPHERE) + and gtype != int(GeomType.CAPSULE) + and gtype != int(GeomType.BOX) + and gtype != int(GeomType.CYLINDER) + ): + continue + + g_contype = geom_contype[geomid] + g_conaffinity = geom_conaffinity[geomid] + f_contype = flex_contype[flexid] + f_conaffinity = flex_conaffinity[flexid] + if not ((g_contype & f_conaffinity) or (f_contype & g_conaffinity)): + continue + + geom_margin_val = geom_margin[worldid % geom_margin.shape[0], geomid] + margin = geom_margin_val + flex_margin_val + + geom_pos = geom_xpos_in[worldid, geomid] + geom_rot = geom_xmat_in[worldid, geomid] + geom_size_val = geom_size[worldid % geom_size.shape[0], geomid] + + dist = collision_primitive_core.MJ_MAXVAL + contact_pos = wp.vec3(0.0) + nrm = wp.vec3(0.0) + + if gtype == int(GeomType.SPHERE): + sphere_radius = geom_size_val[0] + dist, contact_pos, nrm = collision_primitive_core.sphere_sphere(v_pos, radius, geom_pos, sphere_radius) + elif gtype == int(GeomType.CAPSULE): + cap_radius = geom_size_val[0] + cap_half_len = geom_size_val[1] + cap_axis = wp.vec3(geom_rot[0, 2], geom_rot[1, 2], geom_rot[2, 2]) + dist, contact_pos, nrm = collision_primitive_core.sphere_capsule( + v_pos, radius, geom_pos, cap_axis, cap_radius, cap_half_len + ) + elif gtype == int(GeomType.BOX): + dist, contact_pos, nrm = collision_primitive_core.sphere_box(v_pos, radius, geom_pos, geom_rot, geom_size_val) + elif gtype == int(GeomType.CYLINDER): + cyl_radius = geom_size_val[0] + cyl_half_height = geom_size_val[1] + cyl_axis = wp.vec3(geom_rot[0, 2], geom_rot[1, 2], geom_rot[2, 2]) + dist, contact_pos, nrm = collision_primitive_core.sphere_cylinder( + v_pos, radius, geom_pos, cyl_axis, cyl_radius, cyl_half_height + ) + + if dist < margin: + _write_candidate_contact( + max_candidates, + dist, + contact_pos, + nrm, + geomid, + flexid, + -1, + local_vertid, + worldid, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, + ) + + +@wp.func +def _sphere_tetrahedron( + # In: + sphere_pos: wp.vec3, + sphere_radius: float, + t0: wp.vec3, + t1: wp.vec3, + t2: wp.vec3, + t3: wp.vec3, + tri_radius: float, +) -> Tuple[float, wp.vec3, wp.vec3]: + d0, p0, n0 = collision_primitive_core.sphere_triangle(sphere_pos, sphere_radius, t0, t1, t2, tri_radius) + d1, p1, n1 = collision_primitive_core.sphere_triangle(sphere_pos, sphere_radius, t0, t2, t3, tri_radius) + d2, p2, n2 = collision_primitive_core.sphere_triangle(sphere_pos, sphere_radius, t0, t3, t1, tri_radius) + d3, p3, n3 = collision_primitive_core.sphere_triangle(sphere_pos, sphere_radius, t1, t3, t2, tri_radius) + + min_d = d0 + min_p = p0 + min_n = n0 + + if d1 < min_d: + min_d = d1 + min_p = p1 + min_n = n1 + if d2 < min_d: + min_d = d2 + min_p = p2 + min_n = n2 + if d3 < min_d: + min_d = d3 + min_p = p3 + min_n = n3 + + return min_d, min_p, min_n + + +@wp.func +def _plane_vertex( + # In: + pos_v: wp.vec3, + rad: float, + t0: wp.vec3, + t1: wp.vec3, + t2: wp.vec3, +) -> Tuple[bool, float, wp.vec3, wp.vec3]: + e1 = t1 - t0 + e2 = t2 - t0 + ev = pos_v - t0 + + nrm = wp.normalize(wp.cross(e1, e2)) + dst = wp.dot(ev, nrm) + if dst >= 0.0 or dst <= -2.0 * rad: + return False, 0.0, wp.vec3(0.0), wp.vec3(0.0) + + dist = -dst - 2.0 * rad + nrm_out = -nrm + contact_pos = pos_v - nrm * (0.5 * dst) + return True, dist, contact_pos, nrm_out + + +@wp.kernel(module="unique", enable_backward=False) +def _flex_internal_collisions_detect( + # Model: + nflex: int, + flex_margin: wp.array[float], + flex_internal: wp.array[int], + flex_dim: wp.array[int], + flex_vertadr: wp.array[int], + flex_elemdataadr: wp.array[int], + flex_evpairadr: wp.array[int], + flex_evpairnum: wp.array[int], + flex_elem: wp.array[int], + flex_evpair: wp.array[wp.vec2i], + flex_radius: wp.array[float], + flex_evpairflexid: wp.array[int], + # Data in: + flexvert_xpos_in: wp.array2d[wp.vec3], + # In: + max_candidates: int, + # Data out: + overflow_out: wp.array[int], + # Out: + cand_dist_out: wp.array[float], + cand_pos_out: wp.array[wp.vec3], + cand_nrm_out: wp.array[wp.vec3], + cand_geom_out: wp.array[wp.vec2i], + cand_flex_out: wp.array[wp.vec2i], + cand_elem_out: wp.array[wp.vec2i], + cand_vert_out: wp.array[wp.vec2i], + cand_worldid_out: wp.array[int], + cand_type_out: wp.array[int], + cand_geomcollisionid_out: wp.array[int], + ncand_out: wp.array[int], +): + worldid, pair_idx = wp.tid() + + flexid = flex_evpairflexid[pair_idx] + if flex_internal[flexid] == 0: + return + + ev = flex_evpair[pair_idx] + e = ev[0] + v = ev[1] + + dim = flex_dim[flexid] + radius = flex_radius[flexid] + margin = flex_margin[flexid] + vert_adr = flex_vertadr[flexid] + + sphere_pos = flexvert_xpos_in[worldid, vert_adr + v] + + elem_data_idx = flex_elemdataadr[flexid] + e * (dim + 1) + v0_local = flex_elem[elem_data_idx] + p0 = flexvert_xpos_in[worldid, vert_adr + v0_local] + + dist = float(MJ_MAXVAL) + contact_pos = wp.vec3(0.0) + nrm = wp.vec3(0.0) + + if dim == 1: + v1_local = flex_elem[elem_data_idx + 1] + p1 = flexvert_xpos_in[worldid, vert_adr + v1_local] + capsule_pos = 0.5 * (p0 + p1) + capsule_axis = wp.normalize(p1 - p0) + capsule_half_len = 0.5 * wp.length(p1 - p0) + dist, contact_pos, nrm = collision_primitive_core.sphere_capsule( + sphere_pos, radius, capsule_pos, capsule_axis, radius, capsule_half_len + ) + elif dim == 2: + v1_local = flex_elem[elem_data_idx + 1] + v2_local = flex_elem[elem_data_idx + 2] + p1 = flexvert_xpos_in[worldid, vert_adr + v1_local] + p2 = flexvert_xpos_in[worldid, vert_adr + v2_local] + dist, contact_pos, nrm = collision_primitive_core.sphere_triangle(sphere_pos, radius, p0, p1, p2, radius) + elif dim == 3: + v1_local = flex_elem[elem_data_idx + 1] + v2_local = flex_elem[elem_data_idx + 2] + v3_local = flex_elem[elem_data_idx + 3] + p1 = flexvert_xpos_in[worldid, vert_adr + v1_local] + p2 = flexvert_xpos_in[worldid, vert_adr + v2_local] + p3 = flexvert_xpos_in[worldid, vert_adr + v3_local] + dist, contact_pos, nrm = _sphere_tetrahedron(sphere_pos, radius, p0, p1, p2, p3, radius) + + if dist < margin: + _write_candidate_contact( + max_candidates, + dist, + contact_pos, + nrm, + -1, + flexid, + e, + v, + worldid, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, + ) + + +@wp.kernel(module="unique", enable_backward=False) +def _flex_tet_internal_collisions_detect( + # Model: + nflex: int, flex_dim: wp.array[int], flex_vertadr: wp.array[int], flex_elemadr: wp.array[int], @@ -506,147 +1395,511 @@ def _flex_narrowphase_dim2( flex_elemdataadr: wp.array[int], flex_elem: wp.array[int], flex_radius: wp.array[float], + flex_elemflexid: wp.array[int], # Data in: - geom_xpos_in: wp.array2d[wp.vec3], - geom_xmat_in: wp.array2d[wp.mat33], flexvert_xpos_in: wp.array2d[wp.vec3], - nworld_in: int, - naconmax_in: int, + # In: + max_candidates: int, # Data out: - contact_dist_out: wp.array[float], - contact_pos_out: wp.array[wp.vec3], - contact_frame_out: wp.array[wp.mat33], - contact_includemargin_out: wp.array[float], - contact_friction_out: wp.array[vec5], - contact_solref_out: wp.array[wp.vec2], - contact_solreffriction_out: wp.array[wp.vec2], - contact_solimp_out: wp.array[vec5], - contact_dim_out: wp.array[int], - contact_geom_out: wp.array[wp.vec2i], - contact_flex_out: wp.array[wp.vec2i], - contact_vert_out: wp.array[wp.vec2i], - contact_worldid_out: wp.array[int], - contact_type_out: wp.array[int], - contact_geomcollisionid_out: wp.array[int], - nacon_out: wp.array[int], + overflow_out: wp.array[int], + # Out: + cand_dist_out: wp.array[float], + cand_pos_out: wp.array[wp.vec3], + cand_nrm_out: wp.array[wp.vec3], + cand_geom_out: wp.array[wp.vec2i], + cand_flex_out: wp.array[wp.vec2i], + cand_elem_out: wp.array[wp.vec2i], + cand_vert_out: wp.array[wp.vec2i], + cand_worldid_out: wp.array[int], + cand_type_out: wp.array[int], + cand_geomcollisionid_out: wp.array[int], + ncand_out: wp.array[int], ): worldid, elemid = wp.tid() - flexid = int(-1) - for i in range(nflex): - if flex_dim[i] != 2: - continue - elem_adr = flex_elemadr[i] - elem_num = flex_elemnum[i] - if elemid >= elem_adr and elemid < elem_adr + elem_num: - flexid = i - break - - if flexid < 0: + flexid = flex_elemflexid[elemid] + if flex_dim[flexid] != 3: return + radius = flex_radius[flexid] vert_adr = flex_vertadr[flexid] - tri_radius = flex_radius[flexid] - tri_margin = flex_margin[flexid] - elem_data_idx = flex_elemdataadr[flexid] + (elemid - flex_elemadr[flexid]) * 3 - v0_local = flex_elem[elem_data_idx] - v1_local = flex_elem[elem_data_idx + 1] - v2_local = flex_elem[elem_data_idx + 2] + local_elemid = elemid - flex_elemadr[flexid] + elem_data_idx = flex_elemdataadr[flexid] + local_elemid * 4 - t1 = flexvert_xpos_in[worldid, vert_adr + v0_local] - t2 = flexvert_xpos_in[worldid, vert_adr + v1_local] - t3 = flexvert_xpos_in[worldid, vert_adr + v2_local] + v0 = flex_elem[elem_data_idx] + v1 = flex_elem[elem_data_idx + 1] + v2 = flex_elem[elem_data_idx + 2] + v3 = flex_elem[elem_data_idx + 3] - # TODO: Add a broadphase - for geomid in range(ngeom): - gtype = geom_type[geomid] - if ( - gtype != int(GeomType.SPHERE) - and gtype != int(GeomType.CAPSULE) - and gtype != int(GeomType.BOX) - and gtype != int(GeomType.CYLINDER) - ): - continue + p0 = flexvert_xpos_in[worldid, vert_adr + v0] + p1 = flexvert_xpos_in[worldid, vert_adr + v1] + p2 = flexvert_xpos_in[worldid, vert_adr + v2] + p3 = flexvert_xpos_in[worldid, vert_adr + v3] - g_contype = geom_contype[geomid] - g_conaffinity = geom_conaffinity[geomid] - f_contype = flex_contype[flexid] - f_conaffinity = flex_conaffinity[flexid] - if not ((g_contype & f_conaffinity) or (f_contype & g_conaffinity)): - continue - - geom_margin_val = geom_margin[worldid % geom_margin.shape[0], geomid] - margin = geom_margin_val + tri_margin - - geom_pos = geom_xpos_in[worldid, geomid] - geom_rot = geom_xmat_in[worldid, geomid] - geom_size_val = geom_size[worldid % geom_size.shape[0], geomid] - - condim, gap, solref, solimp, friction = _mix_flex_contact_params( - geom_condim[geomid], - geom_priority[geomid], - geom_solmix[worldid % geom_solmix.shape[0], geomid], - geom_solref[worldid % geom_solref.shape[0], geomid], - geom_solimp[worldid % geom_solimp.shape[0], geomid], - geom_friction[worldid % geom_friction.shape[0], geomid], - geom_gap[worldid % geom_gap.shape[0], geomid], - flex_condim[flexid], - flex_priority[flexid], - flex_solmix[flexid], - flex_solref[flexid], - flex_solimp[flexid], - flex_friction[flexid], - flex_gap[flexid], - ) - - _collide_geom_triangle( - naconmax_in, - gtype, - geom_pos, - geom_rot, - geom_size_val, - t1, - t2, - t3, - tri_radius, - margin, - condim, - friction, - solref, - solimp, - geomid, + # Test face (0,1,2) vs Vertex 3 + ok0, dist0, pos0, nrm0 = _plane_vertex(p3, radius, p0, p1, p2) + if ok0: + _write_candidate_contact( + max_candidates, + dist0, + pos0, + nrm0, + -1, flexid, - v0_local, + local_elemid, + v3, worldid, - 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_flex_out, - contact_vert_out, - contact_worldid_out, - contact_type_out, - contact_geomcollisionid_out, - nacon_out, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, ) + # Test face (0,2,3) vs Vertex 1 + ok1, dist1, pos1, nrm1 = _plane_vertex(p1, radius, p0, p2, p3) + if ok1: + _write_candidate_contact( + max_candidates, + dist1, + pos1, + nrm1, + -1, + flexid, + local_elemid, + v1, + worldid, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, + ) + + # Test face (0,3,1) vs Vertex 2 + ok2, dist2, pos2, nrm2 = _plane_vertex(p2, radius, p0, p3, p1) + if ok2: + _write_candidate_contact( + max_candidates, + dist2, + pos2, + nrm2, + -1, + flexid, + local_elemid, + v2, + worldid, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, + ) + + # Test face (1,3,2) vs Vertex 0 + ok3, dist3, pos3, nrm3 = _plane_vertex(p0, radius, p1, p3, p2) + if ok3: + _write_candidate_contact( + max_candidates, + dist3, + pos3, + nrm3, + -1, + flexid, + local_elemid, + v0, + worldid, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, + ) + + +@wp.func +def _inside_triangle( + # In: + P: wp.vec3, + A: wp.vec3, + B: wp.vec3, + C: wp.vec3, + tol: float, +) -> bool: + v0 = B - A + v1 = C - A + v2 = P - A + d00 = wp.dot(v0, v0) + d01 = wp.dot(v0, v1) + d11 = wp.dot(v1, v1) + d20 = wp.dot(v2, v0) + d21 = wp.dot(v2, v1) + denom = d00 * d11 - d01 * d01 + if wp.abs(denom) < 1e-12: + return False + v = (d11 * d20 - d01 * d21) / denom + w = (d00 * d21 - d01 * d20) / denom + u = 1.0 - v - w + return u >= -tol and v >= -tol and w >= -tol and u <= 1.0 + tol and v <= 1.0 + tol and w <= 1.0 + tol + + +@wp.func +def _exclude_self_collision( + # Model: + flex_vertbodyid: wp.array[int], + # In: + v1: wp.vec4i, + n1: int, + v2: wp.vec4i, + n2: int, + vert_adr: int, +) -> bool: + for i in range(n1): + idx1 = v1[i] + if idx1 >= 0: + b1 = flex_vertbodyid[vert_adr + idx1] + for j in range(n2): + idx2 = v2[j] + if idx1 == idx2: + return True + if idx2 >= 0 and b1 >= 0: + b2 = flex_vertbodyid[vert_adr + idx2] + if b1 == b2: + return True + return False + + +@wp.func +def _get_element_vertices( + # Model: + flex_elem: wp.array[int], + # In: + dim: int, + elem_data_idx: int, +) -> wp.vec4i: + v0 = flex_elem[elem_data_idx] + v1 = flex_elem[elem_data_idx + 1] + v2 = int(-1) + v3 = int(-1) + if dim >= 2: + v2 = flex_elem[elem_data_idx + 2] + if dim >= 3: + v3 = flex_elem[elem_data_idx + 3] + return wp.vec4i(v0, v1, v2, v3) + + +@wp.func +def _elements_overlap( + # Data in: + flexvert_xpos_in: wp.array2d[wp.vec3], + # In: + dim: int, + radius: float, + v1_indices: wp.vec4i, + v2_indices: wp.vec4i, + vert_adr: int, + worldid: int, +) -> bool: + p1_0 = flexvert_xpos_in[worldid, vert_adr + v1_indices[0]] + p1_1 = flexvert_xpos_in[worldid, vert_adr + v1_indices[1]] + + min1 = wp.min(p1_0, p1_1) + max1 = wp.max(p1_0, p1_1) + + if dim >= 2: + p1_2 = flexvert_xpos_in[worldid, vert_adr + v1_indices[2]] + min1 = wp.min(min1, p1_2) + max1 = wp.max(max1, p1_2) + if dim >= 3: + p1_3 = flexvert_xpos_in[worldid, vert_adr + v1_indices[3]] + min1 = wp.min(min1, p1_3) + max1 = wp.max(max1, p1_3) + + p2_0 = flexvert_xpos_in[worldid, vert_adr + v2_indices[0]] + p2_1 = flexvert_xpos_in[worldid, vert_adr + v2_indices[1]] + + min2 = wp.min(p2_0, p2_1) + max2 = wp.max(p2_0, p2_1) + + if dim >= 2: + p2_2 = flexvert_xpos_in[worldid, vert_adr + v2_indices[2]] + min2 = wp.min(min2, p2_2) + max2 = wp.max(max2, p2_2) + if dim >= 3: + p2_3 = flexvert_xpos_in[worldid, vert_adr + v2_indices[3]] + min2 = wp.min(min2, p2_3) + max2 = wp.max(max2, p2_3) + + rbound = 2.0 * radius + + if min1[0] - rbound > max2[0] or max1[0] + rbound < min2[0]: + return False + if min1[1] - rbound > max2[1] or max1[1] + rbound < min2[1]: + return False + if min1[2] - rbound > max2[2] or max1[2] + rbound < min2[2]: + return False + + return True + + +@wp.kernel(module="unique", enable_backward=False) +def _flex_active_element_collisions_detect( + # Model: + nflex: int, + opt_ccd_tolerance: wp.array[float], + flex_selfcollide: wp.array[int], + flex_dim: wp.array[int], + flex_vertadr: wp.array[int], + flex_elemadr: wp.array[int], + flex_elemnum: wp.array[int], + flex_elemdataadr: wp.array[int], + flex_vertbodyid: wp.array[int], + flex_elem: wp.array[int], + flex_radius: wp.array[float], + flex_elemflexid: wp.array[int], + # Data in: + flexvert_xpos_in: wp.array2d[wp.vec3], + # In: + max_candidates: int, + gjk_iterations: int, + epa_iterations: int, + n_total_elems: int, + # Data out: + overflow_out: wp.array[int], + # Out: + workspace_verts_out: wp.array[wp.vec3], + epa_vert_out: wp.array2d[wp.vec3], + epa_vert_index_out: wp.array2d[int], + epa_face_out: wp.array2d[int], + epa_pr_out: wp.array2d[wp.vec3], + epa_norm2_out: wp.array2d[float], + epa_horizon_out: wp.array2d[int], + cand_dist_out: wp.array[float], + cand_pos_out: wp.array[wp.vec3], + cand_nrm_out: wp.array[wp.vec3], + cand_geom_out: wp.array[wp.vec2i], + cand_flex_out: wp.array[wp.vec2i], + cand_elem_out: wp.array[wp.vec2i], + cand_vert_out: wp.array[wp.vec2i], + cand_worldid_out: wp.array[int], + cand_type_out: wp.array[int], + cand_geomcollisionid_out: wp.array[int], + ncand_out: wp.array[int], +): + worldid, elem1_global = wp.tid() + + flexid = flex_elemflexid[elem1_global] + if flex_selfcollide[flexid] == 0: + return + + radius = flex_radius[flexid] + dim = flex_dim[flexid] + vert_adr = flex_vertadr[flexid] + elem_adr = flex_elemadr[flexid] + elem_num = flex_elemnum[flexid] + + e1 = elem1_global - elem_adr + elem_data_idx1 = flex_elemdataadr[flexid] + e1 * (dim + 1) + + v1_indices = _get_element_vertices(flex_elem, dim, elem_data_idx1) + + unique_thread_id = worldid * n_total_elems + elem1_global + + offset1 = unique_thread_id * 8 + for idx in range(dim + 1): + workspace_verts_out[offset1 + idx] = flexvert_xpos_in[worldid, vert_adr + v1_indices[idx]] + + for e2 in range(e1 + 1, elem_num): + elem_data_idx2 = flex_elemdataadr[flexid] + e2 * (dim + 1) + v2_indices = _get_element_vertices(flex_elem, dim, elem_data_idx2) + + if _exclude_self_collision(flex_vertbodyid, v1_indices, dim + 1, v2_indices, dim + 1, vert_adr): + continue + + overlap = _elements_overlap(flexvert_xpos_in, dim, radius, v1_indices, v2_indices, vert_adr, worldid) + if not overlap: + continue + + if dim == 1: + p0 = workspace_verts_out[offset1] + p1 = workspace_verts_out[offset1 + 1] + cap1_pos = 0.5 * (p0 + p1) + cap1_axis = wp.normalize(p1 - p0) + cap1_half_len = 0.5 * wp.length(p1 - p0) + + p2_0 = flexvert_xpos_in[worldid, vert_adr + v2_indices[0]] + p2_1 = flexvert_xpos_in[worldid, vert_adr + v2_indices[1]] + cap2_pos = 0.5 * (p2_0 + p2_1) + cap2_axis = wp.normalize(p2_1 - p2_0) + cap2_half_len = 0.5 * wp.length(p2_1 - p2_0) + + margin = 0.0 + + contact_dist, contact_pos, contact_normal = collision_primitive_core.capsule_capsule( + cap1_pos, cap1_axis, radius, cap1_half_len, cap2_pos, cap2_axis, radius, cap2_half_len, margin + ) + + for c in range(2): + d_val = contact_dist[c] + if d_val < 0.0: + _write_candidate_contact( + max_candidates, + d_val, + contact_pos[c], + contact_normal[c], + -2, + flexid, + e1, + e2, + worldid, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, + ) + else: + offset2 = unique_thread_id * 8 + 4 + for idx in range(dim + 1): + workspace_verts_out[offset2 + idx] = flexvert_xpos_in[worldid, vert_adr + v2_indices[idx]] + + geom1 = Geom() + geom1.pos = wp.vec3(0.0) + geom1.rot = wp.mat33(1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0) + geom1.size = wp.vec3(0.0) + geom1.margin = 2.0 * radius + geom1.vert = workspace_verts_out + geom1.vertadr = offset1 + geom1.vertnum = dim + 1 + geom1.graphadr = -1 + geom1.index = -1 + + geom2 = Geom() + geom2.pos = wp.vec3(0.0) + geom2.rot = wp.mat33(1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0) + geom2.size = wp.vec3(0.0) + geom2.margin = 2.0 * radius + geom2.vert = workspace_verts_out + geom2.vertadr = offset2 + geom2.vertnum = dim + 1 + geom2.graphadr = -1 + geom2.index = -1 + + center1 = wp.vec3(0.0) + for idx in range(dim + 1): + center1 += workspace_verts_out[offset1 + idx] + center1 = center1 / float(dim + 1) + + center2 = wp.vec3(0.0) + for idx in range(dim + 1): + center2 += workspace_verts_out[offset2 + idx] + center2 = center2 / float(dim + 1) + + tol = opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]] + + dist, ncontact, w1, w2, _ = ccd( + tol, + 2.0 * radius, + gjk_iterations, + epa_iterations, + geom1, + geom2, + int(GeomType.MESH), + int(GeomType.MESH), + center1, + center2, + epa_vert_out[unique_thread_id], + epa_vert_index_out[unique_thread_id], + epa_face_out[unique_thread_id], + epa_pr_out[unique_thread_id], + epa_norm2_out[unique_thread_id], + epa_horizon_out[unique_thread_id], + ) + + phys_dist = dist + if ncontact > 0 and phys_dist < 0.0: + p1_0 = workspace_verts_out[offset1] + p1_1 = workspace_verts_out[offset1 + 1] + p1_2 = workspace_verts_out[offset1 + 2] + p2_0 = workspace_verts_out[offset2] + p2_1 = workspace_verts_out[offset2 + 1] + p2_2 = workspace_verts_out[offset2 + 2] + if not (_inside_triangle(w1, p1_0, p1_1, p1_2, 0.2) and _inside_triangle(w2, p2_0, p2_1, p2_2, 0.2)): + continue + + pos = 0.5 * (w1 + w2) + nrm = wp.normalize(w1 - w2) + _write_candidate_contact( + max_candidates, + phys_dist, + pos, + nrm, + -2, + flexid, + e1, + e2, + worldid, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, + ) + @wp.kernel -def _flex_narrowphase_dim3( +def _flex_narrowphase_unified( # Model: ngeom: int, nflex: int, + opt_ccd_tolerance: wp.array[float], + opt_warn_overflow: bool, geom_type: wp.array[int], - geom_contype: wp.array[int], - geom_conaffinity: wp.array[int], geom_condim: wp.array[int], + geom_dataid: wp.array2d[int], geom_priority: wp.array[int], geom_solmix: wp.array2d[float], geom_solref: wp.array2d[wp.vec2], @@ -655,8 +1908,6 @@ def _flex_narrowphase_dim3( geom_friction: wp.array2d[wp.vec3], geom_margin: wp.array2d[float], geom_gap: wp.array2d[float], - flex_contype: wp.array[int], - flex_conaffinity: wp.array[int], flex_condim: wp.array[int], flex_priority: wp.array[int], flex_solmix: wp.array[float], @@ -667,16 +1918,280 @@ def _flex_narrowphase_dim3( flex_gap: wp.array[float], flex_dim: wp.array[int], flex_vertadr: wp.array[int], - flex_shellnum: wp.array[int], - flex_shelldataadr: wp.array[int], - flex_shell: wp.array[int], flex_radius: wp.array[float], + mesh_vertadr: wp.array[int], + mesh_vertnum: wp.array[int], + mesh_graphadr: wp.array[int], + mesh_vert: wp.array[wp.vec3], + mesh_graph: wp.array[int], + mesh_pos: wp.array[wp.vec3], + mesh_polynormal: wp.array[wp.vec3], + mesh_polyvertadr: wp.array[int], + mesh_polyvert: wp.array[int], + mesh_polymapadr: wp.array[int], + mesh_polymapnum: wp.array[int], + mesh_polymap: wp.array[int], # Data in: geom_xpos_in: wp.array2d[wp.vec3], geom_xmat_in: wp.array2d[wp.mat33], flexvert_xpos_in: wp.array2d[wp.vec3], nworld_in: int, naconmax_in: int, + naccdmax_in: int, + ncollision_in: wp.array[int], + # In: + triadr: wp.array[int], + flex_tridataadr: wp.array[int], + flex_tri: wp.array[int], + flex_triflexid: wp.array[int], + collision_pair_in: wp.array[wp.vec2i], + collision_worldid_in: wp.array[int], + epa_vert: wp.array2d[wp.vec3], + epa_vert_index: wp.array2d[int], + epa_face: wp.array2d[int], + epa_pr: wp.array2d[wp.vec3], + epa_norm2: wp.array2d[float], + epa_horizon: wp.array2d[int], + nccd: wp.array[int], + ccd_iterations: int, + max_candidates: int, + # Data out: + overflow_out: wp.array[int], + # Out: + cand_dist_out: wp.array[float], + cand_pos_out: wp.array[wp.vec3], + cand_nrm_out: wp.array[wp.vec3], + cand_geom_out: wp.array[wp.vec2i], + cand_flex_out: wp.array[wp.vec2i], + cand_elem_out: wp.array[wp.vec2i], + cand_vert_out: wp.array[wp.vec2i], + cand_worldid_out: wp.array[int], + cand_type_out: wp.array[int], + cand_geomcollisionid_out: wp.array[int], + ncand_out: wp.array[int], +): + collisionid = wp.tid() + if collisionid >= ncollision_in[0] or collisionid >= collision_pair_in.shape[0]: + return + + pair = collision_pair_in[collisionid] + tri_id = pair[0] + geomid = pair[1] + + gtype = geom_type[geomid] + if ( + gtype != int(GeomType.SPHERE) + and gtype != int(GeomType.CAPSULE) + and gtype != int(GeomType.BOX) + and gtype != int(GeomType.CYLINDER) + and gtype != int(GeomType.MESH) + ): + return + + worldid = collision_worldid_in[collisionid] + + flexid = flex_triflexid[tri_id] + + vert_adr = flex_vertadr[flexid] + tri_radius = flex_radius[flexid] + tri_margin = flex_margin[flexid] + + local_tri_id = tri_id - triadr[flexid] + tri_data_idx = flex_tridataadr[flexid] + local_tri_id * 3 + v0_local = flex_tri[tri_data_idx] + v1_local = flex_tri[tri_data_idx + 1] + v2_local = flex_tri[tri_data_idx + 2] + + t1 = flexvert_xpos_in[worldid, vert_adr + v0_local] + t2 = flexvert_xpos_in[worldid, vert_adr + v1_local] + t3 = flexvert_xpos_in[worldid, vert_adr + v2_local] + + geom_margin_val = geom_margin[worldid % geom_margin.shape[0], geomid] + margin = geom_margin_val + tri_margin + + geom_pos = geom_xpos_in[worldid, geomid] + geom_rot = geom_xmat_in[worldid, geomid] + geom_size_val = geom_size[worldid % geom_size.shape[0], geomid] + + if gtype == int(GeomType.MESH): + ccdid = wp.atomic_add(nccd, 0, 1) + if ccdid >= naccdmax_in: + if opt_warn_overflow: + wp.printf("CCD overflow in flex narrowphase - please increase naccdmax to %u\n", ccdid) + wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.CCD)) + else: + did = geom_dataid[worldid % geom_dataid.shape[0], geomid] + tolerance = opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]] + _collide_mesh_triangle( + mesh_vertadr, + mesh_vertnum, + mesh_graphadr, + mesh_vert, + mesh_graph, + mesh_pos, + mesh_polynormal, + mesh_polyvertadr, + mesh_polyvert, + mesh_polymapadr, + mesh_polymapnum, + mesh_polymap, + max_candidates, + geom_pos, + geom_rot, + geom_size_val, + t1, + t2, + t3, + tri_radius, + margin, + geomid, + flexid, + local_tri_id, + v0_local, + v1_local, + v2_local, + worldid, + did, + epa_vert[ccdid], + epa_vert_index[ccdid], + epa_face[ccdid], + epa_pr[ccdid], + epa_norm2[ccdid], + epa_horizon[ccdid], + tolerance, + ccd_iterations, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, + ) + else: + _collide_geom_triangle_detect( + max_candidates, + gtype, + geom_pos, + geom_rot, + geom_size_val, + t1, + t2, + t3, + tri_radius, + margin, + geomid, + flexid, + local_tri_id, + -1, + worldid, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, + ) + + +@wp.kernel +def _filter_flex_candidates( + # In: + max_candidates: int, + ncand: wp.array[int], + epsilon: float, + cand_dist: wp.array[float], + cand_pos: wp.array[wp.vec3], + cand_geom: wp.array[wp.vec2i], + cand_flex: wp.array[wp.vec2i], + cand_worldid: wp.array[int], + # Out: + cand_active_out: wp.array[int], +): + i = wp.tid() + limit = ncand[0] + if i >= limit: + return + + geom_i = cand_geom[i][0] + flex_i = cand_flex[i][1] + world_i = cand_worldid[i] + pos_i = cand_pos[i] + dist_i = cand_dist[i] + + keep = int(1) + for j in range(max_candidates): + if j >= limit: + break + if j == i: + continue + geom_j = cand_geom[j][0] + if (geom_i >= 0 and geom_j >= 0 and geom_j == geom_i) or (geom_i < 0 and geom_j < 0): + flex_j = cand_flex[j][1] + if flex_j == flex_i: + world_j = cand_worldid[j] + if world_j == world_i: + pos_j = cand_pos[j] + dist_j = cand_dist[j] + + diff = pos_i - pos_j + if wp.dot(diff, diff) < epsilon * epsilon: + if dist_j < dist_i: + keep = 0 + elif dist_j == dist_i and j < i: + keep = 0 + + cand_active_out[i] = keep + + +@wp.kernel +def _write_filtered_contacts( + # Model: + opt_warn_overflow: bool, + geom_type: wp.array[int], + geom_condim: wp.array[int], + geom_priority: wp.array[int], + geom_solmix: wp.array2d[float], + geom_solref: wp.array2d[wp.vec2], + geom_solimp: wp.array2d[vec5], + geom_friction: wp.array2d[wp.vec3], + geom_margin: wp.array2d[float], + geom_gap: wp.array2d[float], + flex_condim: wp.array[int], + flex_priority: wp.array[int], + flex_solmix: wp.array[float], + flex_solref: wp.array[wp.vec2], + flex_solimp: wp.array[vec5], + flex_friction: wp.array[wp.vec3], + flex_margin: wp.array[float], + flex_gap: wp.array[float], + flex_dim: wp.array[int], + # Data in: + naconmax_in: int, + # In: + ncand: wp.array[int], + cand_dist: wp.array[float], + cand_pos: wp.array[wp.vec3], + cand_nrm: wp.array[wp.vec3], + cand_geom: wp.array[wp.vec2i], + cand_flex: wp.array[wp.vec2i], + cand_elem: wp.array[wp.vec2i], + cand_vert: wp.array[wp.vec2i], + cand_worldid: wp.array[int], + cand_type: wp.array[int], + cand_geomcollisionid: wp.array[int], + cand_active: wp.array[int], # Data out: contact_dist_out: wp.array[float], contact_pos_out: wp.array[wp.vec3], @@ -689,69 +2204,38 @@ def _flex_narrowphase_dim3( contact_dim_out: wp.array[int], contact_geom_out: wp.array[wp.vec2i], contact_flex_out: wp.array[wp.vec2i], + contact_elem_out: wp.array[wp.vec2i], contact_vert_out: wp.array[wp.vec2i], contact_worldid_out: wp.array[int], contact_type_out: wp.array[int], contact_geomcollisionid_out: wp.array[int], nacon_out: wp.array[int], + # Data out: + overflow_out: wp.array[int], ): - worldid, shellid = wp.tid() - - flexid = int(-1) - shell_offset = int(0) - for i in range(nflex): - if flex_dim[i] != 3: - continue - shell_num = flex_shellnum[i] - if shellid >= shell_offset and shellid < shell_offset + shell_num: - flexid = i - break - shell_offset += shell_num - - if flexid < 0: + i = wp.tid() + if i >= ncand[0]: return - vert_adr = flex_vertadr[flexid] - tri_radius = flex_radius[flexid] - tri_margin = flex_margin[flexid] + if cand_active[i] == 0: + return - shell_adr = flex_shelldataadr[flexid] - local_shellid = shellid - shell_offset - shell_data_idx = shell_adr + local_shellid * 3 + geomid = cand_geom[i][0] + worldid = cand_worldid[i] - v0_local = flex_shell[shell_data_idx] - v1_local = flex_shell[shell_data_idx + 1] - v2_local = flex_shell[shell_data_idx + 2] - - t1 = flexvert_xpos_in[worldid, vert_adr + v0_local] - t2 = flexvert_xpos_in[worldid, vert_adr + v1_local] - t3 = flexvert_xpos_in[worldid, vert_adr + v2_local] - - # TODO: Add a broadphase - for geomid in range(ngeom): - gtype = geom_type[geomid] - if ( - gtype != int(GeomType.SPHERE) - and gtype != int(GeomType.CAPSULE) - and gtype != int(GeomType.BOX) - and gtype != int(GeomType.CYLINDER) - ): - continue - - g_contype = geom_contype[geomid] - g_conaffinity = geom_conaffinity[geomid] - f_contype = flex_contype[flexid] - f_conaffinity = flex_conaffinity[flexid] - if not ((g_contype & f_conaffinity) or (f_contype & g_conaffinity)): - continue + condim = int(0) + margin = float(0.0) + gap = float(0.0) + solref = wp.vec2(0.0, 0.0) + solimp = vec5(0.0, 0.0, 0.0, 0.0, 0.0) + friction = vec5(0.0, 0.0, 0.0, 0.0, 0.0) + if geomid >= 0: + flexid = cand_flex[i][1] geom_margin_val = geom_margin[worldid % geom_margin.shape[0], geomid] + tri_margin = flex_margin[flexid] margin = geom_margin_val + tri_margin - geom_pos = geom_xpos_in[worldid, geomid] - geom_rot = geom_xmat_in[worldid, geomid] - geom_size_val = geom_size[worldid % geom_size.shape[0], geomid] - condim, gap, solref, solimp, friction = _mix_flex_contact_params( geom_condim[geomid], geom_priority[geomid], @@ -768,71 +2252,669 @@ def _flex_narrowphase_dim3( flex_friction[flexid], flex_gap[flexid], ) + else: + flex1 = cand_flex[i][0] + flex2 = cand_flex[i][1] + margin = 0.0 + gap = 0.0 - _collide_geom_triangle( - naconmax_in, - gtype, - geom_pos, - geom_rot, - geom_size_val, - t1, - t2, - t3, - tri_radius, - margin, - condim, - friction, - solref, - solimp, - geomid, - flexid, - v0_local, - worldid, - 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_flex_out, - contact_vert_out, - contact_worldid_out, - contact_type_out, - contact_geomcollisionid_out, - nacon_out, + mixed_condim, _, solref, solimp, friction = _mix_flex_contact_params( + flex_condim[flex1], + flex_priority[flex1], + flex_solmix[flex1], + flex_solref[flex1], + flex_solimp[flex1], + flex_friction[flex1], + 0.0, + flex_condim[flex2], + flex_priority[flex2], + flex_solmix[flex2], + flex_solref[flex2], + flex_solimp[flex2], + flex_friction[flex2], + 0.0, ) + if cand_vert[i][0] >= 0 and flex_dim[flex1] == 3: + condim = 1 + else: + condim = mixed_condim + + if cand_dist[i] >= margin: + return + + id_ = wp.atomic_add(nacon_out, 0, 1) + if id_ >= naconmax_in: + if opt_warn_overflow: + wp.printf( + "flex contact overflow - please increase naconmax to %u\n", + id_ + 1, + ) + wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.NARROWPHASE)) + return + + contact_dist_out[id_] = cand_dist[i] + contact_pos_out[id_] = cand_pos[i] + contact_frame_out[id_] = make_frame(cand_nrm[i]) + if geomid >= 0 and geom_type[geomid] == int(GeomType.PLANE): + contact_includemargin_out[id_] = margin - gap + else: + contact_includemargin_out[id_] = margin + contact_friction_out[id_] = friction + contact_solref_out[id_] = solref + contact_solreffriction_out[id_] = solref + contact_solimp_out[id_] = solimp + contact_dim_out[id_] = condim + contact_geom_out[id_] = cand_geom[i] + contact_flex_out[id_] = cand_flex[i] + contact_elem_out[id_] = cand_elem[i] + contact_vert_out[id_] = cand_vert[i] + contact_worldid_out[id_] = cand_worldid[i] + contact_type_out[id_] = cand_type[i] + contact_geomcollisionid_out[id_] = cand_geomcollisionid[i] + + +def flex_broadphase(m: Model, d: Data): + """Precompute dynamic flex object bounding boxes.""" + wp.launch( + _flex_broadphase_bounds, + dim=(d.nworld, m.nflex), + inputs=[ + m.flex_margin, + m.flex_gap, + m.flex_vertadr, + m.flex_vertnum, + m.flex_radius, + d.flexvert_xpos, + ], + outputs=[ + d.flex_aabb_min, + d.flex_aabb_max, + ], + ) + @event_scope -def flex_narrowphase(m: Model, d: Data): +def flex_collision(m: Model, d: Data, ctx): """Runs collision detection between geoms and flex elements.""" if m.nflex == 0: return + # Deduplicated candidate buffers (allocated on GPU) + cand_dist = wp.empty(d.naconmax, dtype=float) + cand_pos = wp.empty(d.naconmax, dtype=wp.vec3) + cand_nrm = wp.empty(d.naconmax, dtype=wp.vec3) + cand_geom = wp.empty(d.naconmax, dtype=wp.vec2i) + cand_flex = wp.empty(d.naconmax, dtype=wp.vec2i) + cand_elem = wp.empty(d.naconmax, dtype=wp.vec2i) + cand_vert = wp.empty(d.naconmax, dtype=wp.vec2i) + cand_worldid = wp.empty(d.naconmax, dtype=int) + cand_type = wp.empty(d.naconmax, dtype=int) + cand_geomcollisionid = wp.empty(d.naconmax, dtype=int) + + ncand = wp.zeros(1, dtype=int) + + # EPA workspaces if mesh or self collisions are possible + epa_iterations = m.opt.ccd_iterations + if m.nmesh > 0: + mesh_epa_vert = wp.empty(shape=(d.naccdmax, 10 + 2 * epa_iterations), dtype=wp.vec3) + mesh_epa_vert_index = wp.empty(shape=(d.naccdmax, 10 + 2 * epa_iterations), dtype=int) + mesh_epa_face = wp.empty(shape=(d.naccdmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=int) + mesh_epa_pr = wp.empty(shape=(d.naccdmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=wp.vec3) + mesh_epa_norm2 = wp.empty(shape=(d.naccdmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=float) + mesh_epa_horizon = wp.empty(shape=(d.naccdmax, MJ_MAX_EPAHORIZON), dtype=int) + mesh_nccd = wp.zeros(1, dtype=int) + else: + mesh_epa_vert = wp.empty(shape=(1, 1), dtype=wp.vec3) + mesh_epa_vert_index = wp.empty(shape=(1, 1), dtype=int) + mesh_epa_face = wp.empty(shape=(1, 1), dtype=int) + mesh_epa_pr = wp.empty(shape=(1, 1), dtype=wp.vec3) + mesh_epa_norm2 = wp.empty(shape=(1, 1), dtype=float) + mesh_epa_horizon = wp.empty(shape=(1, 1), dtype=int) + mesh_nccd = wp.zeros(1, dtype=int) + + ncollision_dim2 = wp.zeros(1, dtype=int) + ncollision_dim3 = wp.zeros(1, dtype=int) + ncollision_plane = wp.zeros(1, dtype=int) + + # Update dynamic flex object bounding boxes + flex_broadphase(m, d) + + # 2D Flex Element Collisions + if m.flexelem_geom_pair_filtered.shape[0] > 0: + wp.launch( + _flex_broadphase_unified, + dim=(d.nworld, m.flexelem_geom_pair_filtered.shape[0]), + inputs=[ + m.ngeom, + m.nflex, + m.opt.warn_overflow, + m.geom_type, + m.geom_size, + m.geom_aabb, + m.geom_rbound, + m.geom_margin, + m.flex_margin, + m.flex_dim, + m.flex_vertadr, + m.flex_radius, + d.geom_xpos, + d.geom_xmat, + d.flexvert_xpos, + d.naconmax, + d.flex_aabb_min, + d.flex_aabb_max, + m.flex_elemadr, + m.flex_elemdataadr, + m.flex_elem, + m.flexelem_geom_pair_filtered, + m.flex_elemflexid, + ], + outputs=[ + ncollision_dim2, + d.overflow, + ctx.collision_pair, + ctx.collision_worldid, + ], + ) + + wp.launch( + _flex_narrowphase_unified, + dim=d.naconmax, + inputs=[ + m.ngeom, + m.nflex, + m.opt.ccd_tolerance, + m.opt.warn_overflow, + m.geom_type, + m.geom_condim, + m.geom_dataid, + m.geom_priority, + m.geom_solmix, + m.geom_solref, + m.geom_solimp, + m.geom_size, + m.geom_friction, + m.geom_margin, + m.geom_gap, + m.flex_condim, + m.flex_priority, + m.flex_solmix, + m.flex_solref, + m.flex_solimp, + m.flex_friction, + m.flex_margin, + m.flex_gap, + m.flex_dim, + m.flex_vertadr, + m.flex_radius, + m.mesh_vertadr, + m.mesh_vertnum, + m.mesh_graphadr, + m.mesh_vert, + m.mesh_graph, + m.mesh_pos, + m.mesh_polynormal, + m.mesh_polyvertadr, + m.mesh_polyvert, + m.mesh_polymapadr, + m.mesh_polymapnum, + m.mesh_polymap, + d.geom_xpos, + d.geom_xmat, + d.flexvert_xpos, + d.nworld, + d.naconmax, + d.naccdmax, + ncollision_dim2, + m.flex_elemadr, + m.flex_elemdataadr, + m.flex_elem, + m.flex_elemflexid, + ctx.collision_pair, + ctx.collision_worldid, + mesh_epa_vert, + mesh_epa_vert_index, + mesh_epa_face, + mesh_epa_pr, + mesh_epa_norm2, + mesh_epa_horizon, + mesh_nccd, + epa_iterations, + d.naconmax, + ], + outputs=[ + d.overflow, + cand_dist, + cand_pos, + cand_nrm, + cand_geom, + cand_flex, + cand_elem, + cand_vert, + cand_worldid, + cand_type, + cand_geomcollisionid, + ncand, + ], + ) + + # 3D Flex Element Collisions + if m.flexshell_geom_pair_filtered.shape[0] > 0: + wp.launch( + _flex_broadphase_unified, + dim=(d.nworld, m.flexshell_geom_pair_filtered.shape[0]), + inputs=[ + m.ngeom, + m.nflex, + m.opt.warn_overflow, + m.geom_type, + m.geom_size, + m.geom_aabb, + m.geom_rbound, + m.geom_margin, + m.flex_margin, + m.flex_dim, + m.flex_vertadr, + m.flex_radius, + d.geom_xpos, + d.geom_xmat, + d.flexvert_xpos, + d.naconmax, + d.flex_aabb_min, + d.flex_aabb_max, + m.flex_shelladr, + m.flex_shelldataadr, + m.flex_shell, + m.flexshell_geom_pair_filtered, + m.flex_shellflexid, + ], + outputs=[ + ncollision_dim3, + d.overflow, + ctx.collision_pair, + ctx.collision_worldid, + ], + ) + + wp.launch( + _flex_narrowphase_unified, + dim=d.naconmax, + inputs=[ + m.ngeom, + m.nflex, + m.opt.ccd_tolerance, + m.opt.warn_overflow, + m.geom_type, + m.geom_condim, + m.geom_dataid, + m.geom_priority, + m.geom_solmix, + m.geom_solref, + m.geom_solimp, + m.geom_size, + m.geom_friction, + m.geom_margin, + m.geom_gap, + m.flex_condim, + m.flex_priority, + m.flex_solmix, + m.flex_solref, + m.flex_solimp, + m.flex_friction, + m.flex_margin, + m.flex_gap, + m.flex_dim, + m.flex_vertadr, + m.flex_radius, + m.mesh_vertadr, + m.mesh_vertnum, + m.mesh_graphadr, + m.mesh_vert, + m.mesh_graph, + m.mesh_pos, + m.mesh_polynormal, + m.mesh_polyvertadr, + m.mesh_polyvert, + m.mesh_polymapadr, + m.mesh_polymapnum, + m.mesh_polymap, + d.geom_xpos, + d.geom_xmat, + d.flexvert_xpos, + d.nworld, + d.naconmax, + d.naccdmax, + ncollision_dim3, + m.flex_shelladr, + m.flex_shelldataadr, + m.flex_shell, + m.flex_shellflexid, + ctx.collision_pair, + ctx.collision_worldid, + mesh_epa_vert, + mesh_epa_vert_index, + mesh_epa_face, + mesh_epa_pr, + mesh_epa_norm2, + mesh_epa_horizon, + mesh_nccd, + epa_iterations, + d.naconmax, + ], + outputs=[ + d.overflow, + cand_dist, + cand_pos, + cand_nrm, + cand_geom, + cand_flex, + cand_elem, + cand_vert, + cand_worldid, + cand_type, + cand_geomcollisionid, + ncand, + ], + ) + + # Plane Vertex Collisions + if m.flexvert_geom_pair_filtered.shape[0] > 0: + wp.launch( + _flex_broadphase_plane, + dim=(d.nworld, m.flexvert_geom_pair_filtered.shape[0]), + inputs=[ + m.ngeom, + m.opt.warn_overflow, + m.geom_type, + m.geom_margin, + m.flex_margin, + m.flex_vertadr, + m.flex_radius, + m.flexvert_geom_pair_filtered, + m.flex_vertflexid, + d.geom_xpos, + d.geom_xmat, + d.flexvert_xpos, + d.naconmax, + d.flex_aabb_min, + d.flex_aabb_max, + ], + outputs=[ + ncollision_plane, + d.overflow, + ctx.collision_pair, + ctx.collision_worldid, + ], + ) + + wp.launch( + _flex_plane_narrowphase, + dim=d.naconmax, + inputs=[ + m.ngeom, + m.nflexvert, + m.geom_type, + m.geom_condim, + m.geom_priority, + m.geom_solmix, + m.geom_solref, + m.geom_solimp, + m.geom_friction, + m.geom_margin, + m.geom_gap, + m.flex_condim, + m.flex_priority, + m.flex_solmix, + m.flex_solref, + m.flex_solimp, + m.flex_friction, + m.flex_margin, + m.flex_gap, + m.flex_vertadr, + m.flex_radius, + m.flex_vertflexid, + d.geom_xpos, + d.geom_xmat, + d.flexvert_xpos, + d.nworld, + d.naconmax, + ncollision_plane, + ctx.collision_pair, + ctx.collision_worldid, + ], + outputs=[ + d.contact.dist, + d.contact.pos, + d.contact.frame, + d.contact.includemargin, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + d.contact.dim, + d.contact.geom, + d.contact.flex, + d.contact.vert, + d.contact.worldid, + d.contact.type, + d.contact.geomcollisionid, + d.nacon, + ], + ) + + # Geom Vertex Collisions (Candidate-based deduplicated narrowphase) + if m.nflexvert > 0: + wp.launch( + _flex_geom_vertex_narrowphase_detect, + dim=(d.nworld, m.nflexvert), + inputs=[ + m.ngeom, + m.nflexvert, + m.geom_type, + m.geom_contype, + m.geom_conaffinity, + m.geom_size, + m.geom_margin, + m.flex_contype, + m.flex_conaffinity, + m.flex_margin, + m.flex_dim, + m.flex_vertadr, + m.flex_radius, + m.flex_vertflexid, + d.geom_xpos, + d.geom_xmat, + d.flexvert_xpos, + d.nworld, + d.naconmax, + ], + outputs=[ + d.overflow, + cand_dist, + cand_pos, + cand_nrm, + cand_geom, + cand_flex, + cand_elem, + cand_vert, + cand_worldid, + cand_type, + cand_geomcollisionid, + ncand, + ], + ) + + if m.nflexevpair > 0: + wp.launch( + _flex_internal_collisions_detect, + dim=(d.nworld, m.nflexevpair), + inputs=[ + m.nflex, + m.flex_margin, + m.flex_internal, + m.flex_dim, + m.flex_vertadr, + m.flex_elemdataadr, + m.flex_evpairadr, + m.flex_evpairnum, + m.flex_elem, + m.flex_evpair, + m.flex_radius, + m.flex_evpairflexid, + d.flexvert_xpos, + d.naconmax, + ], + outputs=[ + d.overflow, + cand_dist, + cand_pos, + cand_nrm, + cand_geom, + cand_flex, + cand_elem, + cand_vert, + cand_worldid, + cand_type, + cand_geomcollisionid, + ncand, + ], + ) + + if m.nflexelem > 0: + wp.launch( + _flex_tet_internal_collisions_detect, + dim=(d.nworld, m.nflexelem), + inputs=[ + m.nflex, + m.flex_dim, + m.flex_vertadr, + m.flex_elemadr, + m.flex_elemnum, + m.flex_elemdataadr, + m.flex_elem, + m.flex_radius, + m.flex_elemflexid, + d.flexvert_xpos, + d.naconmax, + ], + outputs=[ + d.overflow, + cand_dist, + cand_pos, + cand_nrm, + cand_geom, + cand_flex, + cand_elem, + cand_vert, + cand_worldid, + cand_type, + cand_geomcollisionid, + ncand, + ], + ) + + selfcollide_enabled = m.has_flex_selfcollide + + if selfcollide_enabled and m.nflexelem > 0: + workspace_verts = wp.empty(d.nworld * m.nflexelem * 8, dtype=wp.vec3) + + epa_iterations = m.opt.ccd_iterations + if m.max_flex_dim > 1: + epa_vert = wp.empty(shape=(d.nworld * m.nflexelem, 10 + 2 * epa_iterations), dtype=wp.vec3) + epa_vert_index = wp.empty(shape=(d.nworld * m.nflexelem, 10 + 2 * epa_iterations), dtype=int) + epa_face = wp.empty(shape=(d.nworld * m.nflexelem, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=int) + epa_pr = wp.empty(shape=(d.nworld * m.nflexelem, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=wp.vec3) + epa_norm2 = wp.empty(shape=(d.nworld * m.nflexelem, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=float) + epa_horizon = wp.empty(shape=(d.nworld * m.nflexelem, MJ_MAX_EPAHORIZON), dtype=int) + else: + epa_vert = wp.empty(shape=(1, 1), dtype=wp.vec3) + epa_vert_index = wp.empty(shape=(1, 1), dtype=int) + epa_face = wp.empty(shape=(1, 1), dtype=int) + epa_pr = wp.empty(shape=(1, 1), dtype=wp.vec3) + epa_norm2 = wp.empty(shape=(1, 1), dtype=float) + epa_horizon = wp.empty(shape=(1, 1), dtype=int) + + wp.launch( + _flex_active_element_collisions_detect, + dim=(d.nworld, m.nflexelem), + inputs=[ + m.nflex, + m.opt.ccd_tolerance, + m.flex_selfcollide, + m.flex_dim, + m.flex_vertadr, + m.flex_elemadr, + m.flex_elemnum, + m.flex_elemdataadr, + m.flex_vertbodyid, + m.flex_elem, + m.flex_radius, + m.flex_elemflexid, + d.flexvert_xpos, + d.naconmax, + m.opt.ccd_iterations, + epa_iterations, + m.nflexelem, + ], + outputs=[ + d.overflow, + workspace_verts, + epa_vert, + epa_vert_index, + epa_face, + epa_pr, + epa_norm2, + epa_horizon, + cand_dist, + cand_pos, + cand_nrm, + cand_geom, + cand_flex, + cand_elem, + cand_vert, + cand_worldid, + cand_type, + cand_geomcollisionid, + ncand, + ], + ) + + # Filter duplicate contacts (e.g. from shared vertices or edges) + cand_active = wp.empty(d.naconmax, dtype=int) wp.launch( - _flex_narrowphase_dim2, - dim=(d.nworld, m.nflexelem), + _filter_flex_candidates, + dim=d.naconmax, inputs=[ - m.ngeom, - m.nflex, + d.naconmax, + ncand, + 1e-3, # epsilon + cand_dist, + cand_pos, + cand_geom, + cand_flex, + cand_worldid, + ], + outputs=[ + cand_active, + ], + ) + + # Copy filtered contacts to the main d.contact array, computing contact parameters on-the-fly + wp.launch( + _write_filtered_contacts, + dim=d.naconmax, + inputs=[ + m.opt.warn_overflow, m.geom_type, - m.geom_contype, - m.geom_conaffinity, m.geom_condim, m.geom_priority, m.geom_solmix, m.geom_solref, m.geom_solimp, - m.geom_size, m.geom_friction, m.geom_margin, m.geom_gap, - m.flex_contype, - m.flex_conaffinity, m.flex_condim, m.flex_priority, m.flex_solmix, @@ -842,129 +2924,19 @@ def flex_narrowphase(m: Model, d: Data): m.flex_margin, m.flex_gap, m.flex_dim, - m.flex_vertadr, - m.flex_elemadr, - m.flex_elemnum, - m.flex_elemdataadr, - m.flex_elem, - m.flex_radius, - d.geom_xpos, - d.geom_xmat, - d.flexvert_xpos, - d.nworld, - d.naconmax, - ], - outputs=[ - d.contact.dist, - d.contact.pos, - d.contact.frame, - d.contact.includemargin, - d.contact.friction, - d.contact.solref, - d.contact.solreffriction, - d.contact.solimp, - d.contact.dim, - d.contact.geom, - d.contact.flex, - d.contact.vert, - d.contact.worldid, - d.contact.type, - d.contact.geomcollisionid, - d.nacon, - ], - ) - - wp.launch( - _flex_narrowphase_dim3, - dim=(d.nworld, m.nflexshelldata // 3), - inputs=[ - m.ngeom, - m.nflex, - m.geom_type, - m.geom_contype, - m.geom_conaffinity, - m.geom_condim, - m.geom_priority, - m.geom_solmix, - m.geom_solref, - m.geom_solimp, - m.geom_size, - m.geom_friction, - m.geom_margin, - m.geom_gap, - m.flex_contype, - m.flex_conaffinity, - m.flex_condim, - m.flex_priority, - m.flex_solmix, - m.flex_solref, - m.flex_solimp, - m.flex_friction, - m.flex_margin, - m.flex_gap, - m.flex_dim, - m.flex_vertadr, - m.flex_shellnum, - m.flex_shelldataadr, - m.flex_shell, - m.flex_radius, - d.geom_xpos, - d.geom_xmat, - d.flexvert_xpos, - d.nworld, - d.naconmax, - ], - outputs=[ - d.contact.dist, - d.contact.pos, - d.contact.frame, - d.contact.includemargin, - d.contact.friction, - d.contact.solref, - d.contact.solreffriction, - d.contact.solimp, - d.contact.dim, - d.contact.geom, - d.contact.flex, - d.contact.vert, - d.contact.worldid, - d.contact.type, - d.contact.geomcollisionid, - d.nacon, - ], - ) - - wp.launch( - _flex_plane_narrowphase, - dim=(d.nworld, m.nflexvert), - inputs=[ - m.ngeom, - m.nflexvert, - m.geom_type, - m.geom_condim, - m.geom_priority, - m.geom_solmix, - m.geom_solref, - m.geom_solimp, - m.geom_friction, - m.geom_margin, - m.geom_gap, - m.flex_condim, - m.flex_priority, - m.flex_solmix, - m.flex_solref, - m.flex_solimp, - m.flex_friction, - m.flex_margin, - m.flex_gap, - m.flex_vertadr, - m.flex_radius, - m.flex_vertflexid, - d.geom_xpos, - d.geom_xmat, - d.flexvert_xpos, - d.nworld, d.naconmax, + ncand, + cand_dist, + cand_pos, + cand_nrm, + cand_geom, + cand_flex, + cand_elem, + cand_vert, + cand_worldid, + cand_type, + cand_geomcollisionid, + cand_active, ], outputs=[ d.contact.dist, @@ -978,10 +2950,12 @@ def flex_narrowphase(m: Model, d: Data): d.contact.dim, d.contact.geom, d.contact.flex, + d.contact.elem, d.contact.vert, d.contact.worldid, d.contact.type, d.contact.geomcollisionid, d.nacon, + d.overflow, ], ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py index f299cd78..ea16b361 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py @@ -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): diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive_core.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive_core.py index 8a85d241..73946675 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive_core.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive_core.py @@ -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]), diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py index bcb1acf5..42476838 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py @@ -13,6 +13,8 @@ # limitations under the License. # ============================================================================== +from typing import Tuple + import warp as wp from mujoco.mjx.third_party.mujoco_warp._src import math @@ -36,6 +38,7 @@ def _zero_constraint_counts( nf_out: wp.array[int], nl_out: wp.array[int], nefc_out: wp.array[int], + efc_jtdaj_nblock_out: wp.array[int], # Out: efc_nnz_out: wp.array[int], ): @@ -46,6 +49,7 @@ def _zero_constraint_counts( nf_out[worldid] = 0 nl_out[worldid] = 0 nefc_out[worldid] = 0 + efc_jtdaj_nblock_out[worldid] = 0 efc_nnz_out[worldid] = 0 @@ -121,323 +125,476 @@ def _efc_row( id_out[worldid, efcid] = id -@wp.kernel -def _equality_connect( - # Model: - nv: int, - nsite: int, - opt_timestep: wp.array[float], - opt_disableflags: int, - body_parentid: wp.array[int], - body_rootid: wp.array[int], - body_weldid: wp.array[int], - body_dofnum: wp.array[int], - body_dofadr: wp.array[int], - body_invweight0: wp.array2d[wp.vec2], - jnt_type: wp.array[int], - jnt_dofadr: wp.array[int], - dof_bodyid: wp.array[int], - dof_jntid: wp.array[int], - dof_parentid: wp.array[int], - site_bodyid: wp.array[int], - eq_obj1id: wp.array[int], - eq_obj2id: wp.array[int], - eq_objtype: wp.array[int], - eq_solref: wp.array2d[wp.vec2], - eq_solimp: wp.array2d[vec5], - eq_data: wp.array2d[vec11], - is_sparse: bool, - body_isdofancestor: wp.array2d[int], - eq_connect_adr: wp.array[int], - # Data in: - qvel_in: wp.array2d[float], - eq_active_in: wp.array2d[bool], - xpos_in: wp.array2d[wp.vec3], - xmat_in: wp.array2d[wp.mat33], - site_xpos_in: wp.array2d[wp.vec3], - subtree_com_in: wp.array2d[wp.vec3], - cdof_in: wp.array2d[wp.spatial_vector], - cvel_in: wp.array2d[wp.spatial_vector], - cdof_dot_in: wp.array2d[wp.spatial_vector], - subtree_linvel_in: wp.array2d[wp.vec3], - njmax_in: int, - njmax_nnz_in: int, - # Data out: - ne_out: wp.array[int], - nefc_out: wp.array[int], - efc_type_out: wp.array2d[int], - efc_id_out: wp.array2d[int], - efc_J_rownnz_out: wp.array2d[int], - efc_J_rowadr_out: wp.array2d[int], - efc_J_colind_out: wp.array3d[int], - efc_J_out: wp.array3d[float], - efc_pos_out: wp.array2d[float], - efc_margin_out: wp.array2d[float], - efc_D_out: wp.array2d[float], - efc_vel_out: wp.array2d[float], - efc_aref_out: wp.array2d[float], - efc_frictionloss_out: wp.array2d[float], - # Out: - efc_nnz_out: wp.array[int], -): - """Calculates constraint rows for connect equality constraints.""" - worldid, eqconnectid = wp.tid() - eqid = eq_connect_adr[eqconnectid] +@cache_kernel +def _equality_connect(is_sparse: bool, newton: bool): + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + nv: int, + nsite: int, + opt_timestep: wp.array[float], + opt_disableflags: int, + body_parentid: wp.array[int], + body_rootid: wp.array[int], + body_weldid: wp.array[int], + body_dofnum: wp.array[int], + body_dofadr: wp.array[int], + body_invweight0: wp.array2d[wp.vec2], + jnt_type: wp.array[int], + jnt_dofadr: wp.array[int], + dof_bodyid: wp.array[int], + dof_jntid: wp.array[int], + dof_parentid: wp.array[int], + site_bodyid: wp.array[int], + eq_obj1id: wp.array[int], + eq_obj2id: wp.array[int], + eq_objtype: wp.array[int], + eq_solref: wp.array2d[wp.vec2], + eq_solimp: wp.array2d[vec5], + eq_data: wp.array2d[vec11], + body_isdofancestor: wp.array2d[int], + eq_connect_adr: wp.array[int], + # Data in: + qvel_in: wp.array2d[float], + eq_active_in: wp.array2d[bool], + xpos_in: wp.array2d[wp.vec3], + xmat_in: wp.array2d[wp.mat33], + site_xpos_in: wp.array2d[wp.vec3], + subtree_com_in: wp.array2d[wp.vec3], + cdof_in: wp.array2d[wp.spatial_vector], + cvel_in: wp.array2d[wp.spatial_vector], + cdof_dot_in: wp.array2d[wp.spatial_vector], + subtree_linvel_in: wp.array2d[wp.vec3], + njmax_in: int, + njmax_nnz_in: int, + # Data out: + ne_out: wp.array[int], + nefc_out: wp.array[int], + efc_type_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_jtdaj_adr_out: wp.array2d[int], + efc_jtdaj_nrow_out: wp.array2d[int], + efc_jtdaj_nblock_out: wp.array[int], + efc_J_rownnz_out: wp.array2d[int], + efc_J_rowadr_out: wp.array2d[int], + efc_J_colind_out: wp.array3d[int], + efc_J_out: wp.array3d[float], + efc_pos_out: wp.array2d[float], + efc_margin_out: wp.array2d[float], + efc_D_out: wp.array2d[float], + efc_vel_out: wp.array2d[float], + efc_aref_out: wp.array2d[float], + efc_frictionloss_out: wp.array2d[float], + # Out: + efc_nnz_out: wp.array[int], + ): + """Calculates constraint rows for connect equality constraints.""" + worldid, eqconnectid = wp.tid() + eqid = eq_connect_adr[eqconnectid] - if not eq_active_in[worldid, eqid]: - return - - wp.atomic_add(ne_out, worldid, 3) - efcid = wp.atomic_add(nefc_out, worldid, 3) - - if efcid >= njmax_in - 3: - return - - efcid0 = efcid + 0 - efcid1 = efcid + 1 - efcid2 = efcid + 2 - - data = eq_data[worldid % eq_data.shape[0], eqid] - anchor1 = wp.vec3f(data[0], data[1], data[2]) - anchor2 = wp.vec3f(data[3], data[4], data[5]) - - obj1id = eq_obj1id[eqid] - obj2id = eq_obj2id[eqid] - - if nsite > 0 and eq_objtype[eqid] == types.ObjType.SITE: - body1 = site_bodyid[obj1id] - body2 = site_bodyid[obj2id] - pos1 = site_xpos_in[worldid, obj1id] - pos2 = site_xpos_in[worldid, obj2id] - else: - body1 = obj1id - body2 = obj2id - pos1 = xpos_in[worldid, body1] + xmat_in[worldid, body1] @ anchor1 - pos2 = xpos_in[worldid, body2] + xmat_in[worldid, body2] @ anchor2 - - # error is difference in global positions - pos = pos1 - pos2 - - # compute Jacobian difference (opposite of contact: 0 - 1) - Jqvel = wp.vec3f(0.0, 0.0, 0.0) - Jdotv = wp.vec3f(0.0, 0.0, 0.0) - - if is_sparse: - # TODO(team): pre-compute number of non-zeros - body1 = body_weldid[body1] - body2 = body_weldid[body2] - - da1 = int(body_dofadr[body1] + body_dofnum[body1] - 1) - da2 = int(body_dofadr[body2] + body_dofnum[body2] - 1) - - # count non-zeros - pda1 = da1 - pda2 = da2 - rownnz = int(0) - while pda1 >= 0 or pda2 >= 0: - da = wp.max(pda1, pda2) - if pda1 == da: - pda1 = dof_parentid[pda1] - if pda2 == da: - pda2 = dof_parentid[pda2] - rownnz += 1 - - # get rowadr - rowadr = wp.atomic_add(efc_nnz_out, worldid, 3 * rownnz) - if rowadr + 3 * rownnz > njmax_nnz_in: + if not eq_active_in[worldid, eqid]: return - efc_J_rowadr_out[worldid, efcid0] = rowadr - efc_J_rowadr_out[worldid, efcid1] = rowadr + rownnz - efc_J_rowadr_out[worldid, efcid2] = rowadr + 2 * rownnz - efc_J_rownnz_out[worldid, efcid0] = rownnz - efc_J_rownnz_out[worldid, efcid1] = rownnz - efc_J_rownnz_out[worldid, efcid2] = rownnz + wp.atomic_add(ne_out, worldid, 3) + efcid = wp.atomic_add(nefc_out, worldid, 3) - # compute J and colind - nnz = int(0) - while da1 >= 0 or da2 >= 0: - da = wp.max(da1, da2) - if da1 == da: - da1 = dof_parentid[da1] - if da2 == da: - da2 = dof_parentid[da2] + if efcid >= njmax_in - 3: + return - jacp1, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - body_isdofancestor, - subtree_com_in, - cdof_in, - pos1, - body1, - da, + if wp.static(is_sparse and newton): + jgid = wp.atomic_add(efc_jtdaj_nblock_out, worldid, 1) + efc_jtdaj_adr_out[worldid, jgid] = efcid + efc_jtdaj_nrow_out[worldid, jgid] = 3 + + efcid0 = efcid + 0 + efcid1 = efcid + 1 + efcid2 = efcid + 2 + + data = eq_data[worldid % eq_data.shape[0], eqid] + anchor1 = wp.vec3f(data[0], data[1], data[2]) + anchor2 = wp.vec3f(data[3], data[4], data[5]) + + obj1id = eq_obj1id[eqid] + obj2id = eq_obj2id[eqid] + + if nsite > 0 and eq_objtype[eqid] == types.ObjType.SITE: + body1 = site_bodyid[obj1id] + body2 = site_bodyid[obj2id] + pos1 = site_xpos_in[worldid, obj1id] + pos2 = site_xpos_in[worldid, obj2id] + else: + body1 = obj1id + body2 = obj2id + pos1 = xpos_in[worldid, body1] + xmat_in[worldid, body1] @ anchor1 + pos2 = xpos_in[worldid, body2] + xmat_in[worldid, body2] @ anchor2 + + # error is difference in global positions + pos = pos1 - pos2 + + # compute Jacobian difference (opposite of contact: 0 - 1) + Jqvel = wp.vec3f(0.0, 0.0, 0.0) + Jdotv = wp.vec3f(0.0, 0.0, 0.0) + + if wp.static(is_sparse): + # TODO(team): pre-compute number of non-zeros + body1 = body_weldid[body1] + body2 = body_weldid[body2] + + da1 = int(body_dofadr[body1] + body_dofnum[body1] - 1) + da2 = int(body_dofadr[body2] + body_dofnum[body2] - 1) + + # count non-zeros + pda1 = da1 + pda2 = da2 + rownnz = int(0) + while pda1 >= 0 or pda2 >= 0: + da = wp.max(pda1, pda2) + if pda1 == da: + pda1 = dof_parentid[pda1] + if pda2 == da: + pda2 = dof_parentid[pda2] + rownnz += 1 + + # get rowadr + rowadr = wp.atomic_add(efc_nnz_out, worldid, 3 * rownnz) + if rowadr + 3 * rownnz > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid0] = rowadr + efc_J_rowadr_out[worldid, efcid1] = rowadr + rownnz + efc_J_rowadr_out[worldid, efcid2] = rowadr + 2 * rownnz + + efc_J_rownnz_out[worldid, efcid0] = rownnz + efc_J_rownnz_out[worldid, efcid1] = rownnz + efc_J_rownnz_out[worldid, efcid2] = rownnz + + # compute J and colind + nnz = int(0) + while da1 >= 0 or da2 >= 0: + da = wp.max(da1, da2) + if da1 == da: + da1 = dof_parentid[da1] + if da2 == da: + da2 = dof_parentid[da2] + + jacp1, _ = support.jac_dof( + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + pos1, + body1, + da, + worldid, + ) + jacp2, _ = support.jac_dof( + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + pos2, + body2, + da, + worldid, + ) + j1mj2 = jacp1 - jacp2 + + jacp1_dot, _ = support.jac_dot_dof( + body_parentid, + body_rootid, + jnt_type, + jnt_dofadr, + dof_bodyid, + dof_jntid, + body_isdofancestor, + subtree_com_in, + cdof_in, + cvel_in, + cdof_dot_in, + pos1, + body1, + da, + worldid, + ) + jacp2_dot, _ = support.jac_dot_dof( + body_parentid, + body_rootid, + jnt_type, + jnt_dofadr, + dof_bodyid, + dof_jntid, + body_isdofancestor, + subtree_com_in, + cdof_in, + cvel_in, + cdof_dot_in, + pos2, + body2, + da, + worldid, + ) + j1mj2_dot = jacp1_dot - jacp2_dot + + sparseid0 = rowadr + nnz + sparseid1 = rowadr + rownnz + nnz + sparseid2 = rowadr + 2 * rownnz + nnz + + efc_J_colind_out[worldid, 0, sparseid0] = da + efc_J_colind_out[worldid, 0, sparseid1] = da + efc_J_colind_out[worldid, 0, sparseid2] = da + + efc_J_out[worldid, 0, sparseid0] = j1mj2[0] + efc_J_out[worldid, 0, sparseid1] = j1mj2[1] + efc_J_out[worldid, 0, sparseid2] = j1mj2[2] + + qvel = qvel_in[worldid, da] + Jqvel += j1mj2 * qvel + Jdotv += j1mj2_dot * qvel + + nnz += 1 + else: + # TODO(team): dof tree traversal + for dofid in range(nv): + jacp1, _ = support.jac_dof( + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + pos1, + body1, + dofid, + worldid, + ) + jacp2, _ = support.jac_dof( + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + pos2, + body2, + dofid, + worldid, + ) + j1mj2 = jacp1 - jacp2 + + jacp1_dot, _ = support.jac_dot_dof( + body_parentid, + body_rootid, + jnt_type, + jnt_dofadr, + dof_bodyid, + dof_jntid, + body_isdofancestor, + subtree_com_in, + cdof_in, + cvel_in, + cdof_dot_in, + pos1, + body1, + dofid, + worldid, + ) + jacp2_dot, _ = support.jac_dot_dof( + body_parentid, + body_rootid, + jnt_type, + jnt_dofadr, + dof_bodyid, + dof_jntid, + body_isdofancestor, + subtree_com_in, + cdof_in, + cvel_in, + cdof_dot_in, + pos2, + body2, + dofid, + worldid, + ) + j1mj2_dot = jacp1_dot - jacp2_dot + + efc_J_out[worldid, efcid0, dofid] = j1mj2[0] + efc_J_out[worldid, efcid1, dofid] = j1mj2[1] + efc_J_out[worldid, efcid2, dofid] = j1mj2[2] + + qvel = qvel_in[worldid, dofid] + Jqvel += j1mj2 * qvel + Jdotv += j1mj2_dot * qvel + + body_invweight0_id = worldid % body_invweight0.shape[0] + invweight = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, body2][0] + pos_imp = wp.length(pos) + + solref = eq_solref[worldid % eq_solref.shape[0], eqid] + solimp = eq_solimp[worldid % eq_solimp.shape[0], eqid] + timestep = opt_timestep[worldid % opt_timestep.shape[0]] + + for i in range(3): + efcidi = efcid + i + + _efc_row( + opt_disableflags, worldid, + timestep, + efcidi, + pos[i], + pos_imp, + invweight, + solref, + solimp, + 0.0, + Jqvel[i], + 0.0, + ConstraintType.EQUALITY, + eqid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) - jacp2, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - body_isdofancestor, - subtree_com_in, - cdof_in, - pos2, - body2, - da, - worldid, - ) - j1mj2 = jacp1 - jacp2 - jacp1_dot, _ = support.jac_dot_dof( - body_parentid, - body_rootid, - jnt_type, - jnt_dofadr, - dof_bodyid, - dof_jntid, - body_isdofancestor, - subtree_com_in, - cdof_in, - cvel_in, - cdof_dot_in, - pos1, - body1, - da, - worldid, - ) - jacp2_dot, _ = support.jac_dot_dof( - body_parentid, - body_rootid, - jnt_type, - jnt_dofadr, - dof_bodyid, - dof_jntid, - body_isdofancestor, - subtree_com_in, - cdof_in, - cvel_in, - cdof_dot_in, - pos2, - body2, - da, - worldid, - ) - j1mj2_dot = jacp1_dot - jacp2_dot + efc_aref_out[worldid, efcidi] -= Jdotv[i] - sparseid0 = rowadr + nnz - sparseid1 = rowadr + rownnz + nnz - sparseid2 = rowadr + 2 * rownnz + nnz + return kernel - efc_J_colind_out[worldid, 0, sparseid0] = da - efc_J_colind_out[worldid, 0, sparseid1] = da - efc_J_colind_out[worldid, 0, sparseid2] = da - efc_J_out[worldid, 0, sparseid0] = j1mj2[0] - efc_J_out[worldid, 0, sparseid1] = j1mj2[1] - efc_J_out[worldid, 0, sparseid2] = j1mj2[2] +@cache_kernel +def _equality_joint(is_sparse: bool, newton: bool): + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + nv: int, + opt_timestep: wp.array[float], + opt_disableflags: int, + qpos0: wp.array2d[float], + jnt_qposadr: wp.array[int], + jnt_dofadr: wp.array[int], + dof_invweight0: wp.array2d[float], + eq_obj1id: wp.array[int], + eq_obj2id: wp.array[int], + eq_solref: wp.array2d[wp.vec2], + eq_solimp: wp.array2d[vec5], + eq_data: wp.array2d[vec11], + eq_jnt_adr: wp.array[int], + # Data in: + qpos_in: wp.array2d[float], + qvel_in: wp.array2d[float], + eq_active_in: wp.array2d[bool], + njmax_in: int, + njmax_nnz_in: int, + # Data out: + ne_out: wp.array[int], + nefc_out: wp.array[int], + efc_type_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_jtdaj_adr_out: wp.array2d[int], + efc_jtdaj_nrow_out: wp.array2d[int], + efc_jtdaj_nblock_out: wp.array[int], + efc_J_rownnz_out: wp.array2d[int], + efc_J_rowadr_out: wp.array2d[int], + efc_J_colind_out: wp.array3d[int], + efc_J_out: wp.array3d[float], + efc_pos_out: wp.array2d[float], + efc_margin_out: wp.array2d[float], + efc_D_out: wp.array2d[float], + efc_vel_out: wp.array2d[float], + efc_aref_out: wp.array2d[float], + efc_frictionloss_out: wp.array2d[float], + # Out: + efc_nnz_out: wp.array[int], + ): + worldid, eqjntid = wp.tid() + eqid = eq_jnt_adr[eqjntid] - qvel = qvel_in[worldid, da] - Jqvel += j1mj2 * qvel - Jdotv += j1mj2_dot * qvel + if not eq_active_in[worldid, eqid]: + return - nnz += 1 - else: - # TODO(team): dof tree traversal - for dofid in range(nv): - jacp1, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - body_isdofancestor, - subtree_com_in, - cdof_in, - pos1, - body1, - dofid, - worldid, - ) - jacp2, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - body_isdofancestor, - subtree_com_in, - cdof_in, - pos2, - body2, - dofid, - worldid, - ) - j1mj2 = jacp1 - jacp2 + wp.atomic_add(ne_out, worldid, 1) + efcid = wp.atomic_add(nefc_out, worldid, 1) - jacp1_dot, _ = support.jac_dot_dof( - body_parentid, - body_rootid, - jnt_type, - jnt_dofadr, - dof_bodyid, - dof_jntid, - body_isdofancestor, - subtree_com_in, - cdof_in, - cvel_in, - cdof_dot_in, - pos1, - body1, - dofid, - worldid, - ) - jacp2_dot, _ = support.jac_dot_dof( - body_parentid, - body_rootid, - jnt_type, - jnt_dofadr, - dof_bodyid, - dof_jntid, - body_isdofancestor, - subtree_com_in, - cdof_in, - cvel_in, - cdof_dot_in, - pos2, - body2, - dofid, - worldid, - ) - j1mj2_dot = jacp1_dot - jacp2_dot + if efcid >= njmax_in: + return - efc_J_out[worldid, efcid0, dofid] = j1mj2[0] - efc_J_out[worldid, efcid1, dofid] = j1mj2[1] - efc_J_out[worldid, efcid2, dofid] = j1mj2[2] + if wp.static(is_sparse and newton): + jgid = wp.atomic_add(efc_jtdaj_nblock_out, worldid, 1) + efc_jtdaj_adr_out[worldid, jgid] = efcid + efc_jtdaj_nrow_out[worldid, jgid] = 1 - qvel = qvel_in[worldid, dofid] - Jqvel += j1mj2 * qvel - Jdotv += j1mj2_dot * qvel + jntid_1 = eq_obj1id[eqid] + jntid_2 = eq_obj2id[eqid] + data = eq_data[worldid % eq_data.shape[0], eqid] + dofadr1 = jnt_dofadr[jntid_1] + qposadr1 = jnt_qposadr[jntid_1] + qpos0_id = worldid % qpos0.shape[0] + dof_invweight0_id = worldid % dof_invweight0.shape[0] - body_invweight0_id = worldid % body_invweight0.shape[0] - invweight = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, body2][0] - pos_imp = wp.length(pos) + if wp.static(is_sparse): + if jntid_2 > -1: + rownnz = 2 + else: + rownnz = 1 + efc_J_rownnz_out[worldid, efcid] = rownnz + rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz) + if rowadr + rownnz > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid] = rowadr + efc_J_colind_out[worldid, 0, rowadr] = dofadr1 + efc_J_out[worldid, 0, rowadr] = 1.0 + else: + for i in range(nv): + efc_J_out[worldid, efcid, i] = 0.0 + efc_J_out[worldid, efcid, dofadr1] = 1.0 - solref = eq_solref[worldid % eq_solref.shape[0], eqid] - solimp = eq_solimp[worldid % eq_solimp.shape[0], eqid] - timestep = opt_timestep[worldid % opt_timestep.shape[0]] + if jntid_2 > -1: + # Two joint constraint + qposadr2 = jnt_qposadr[jntid_2] + dofadr2 = jnt_dofadr[jntid_2] + dif = qpos_in[worldid, qposadr2] - qpos0[qpos0_id, qposadr2] - for i in range(3): - efcidi = efcid + i + # Horner's method for polynomials + rhs = data[0] + dif * (data[1] + dif * (data[2] + dif * (data[3] + dif * data[4]))) + deriv_2 = data[1] + dif * (2.0 * data[2] + dif * (3.0 * data[3] + dif * 4.0 * data[4])) + pos = qpos_in[worldid, qposadr1] - qpos0[qpos0_id, qposadr1] - rhs + Jqvel = qvel_in[worldid, dofadr1] - qvel_in[worldid, dofadr2] * deriv_2 + invweight = dof_invweight0[dof_invweight0_id, dofadr1] + dof_invweight0[dof_invweight0_id, dofadr2] + + if wp.static(is_sparse): + sparseid = rowadr + 1 + efc_J_colind_out[worldid, 0, sparseid] = dofadr2 + efc_J_out[worldid, 0, sparseid] = -deriv_2 + else: + efc_J_out[worldid, efcid, dofadr2] = -deriv_2 + else: + # Single joint constraint + pos = qpos_in[worldid, qposadr1] - qpos0[qpos0_id, qposadr1] - data[0] + Jqvel = qvel_in[worldid, dofadr1] + invweight = dof_invweight0[dof_invweight0_id, dofadr1] + + # Update constraint parameters _efc_row( opt_disableflags, worldid, - timestep, - efcidi, - pos[i], - pos_imp, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + pos, + pos, invweight, - solref, - solimp, + eq_solref[worldid % eq_solref.shape[0], eqid], + eq_solimp[worldid % eq_solimp.shape[0], eqid], 0.0, - Jqvel[i], + Jqvel, 0.0, ConstraintType.EQUALITY, eqid, @@ -451,320 +608,200 @@ def _equality_connect( efc_frictionloss_out, ) - efc_aref_out[worldid, efcidi] -= Jdotv[i] - - -@wp.kernel -def _equality_joint( - # Model: - nv: int, - opt_timestep: wp.array[float], - opt_disableflags: int, - qpos0: wp.array2d[float], - jnt_qposadr: wp.array[int], - jnt_dofadr: wp.array[int], - dof_invweight0: wp.array2d[float], - eq_obj1id: wp.array[int], - eq_obj2id: wp.array[int], - eq_solref: wp.array2d[wp.vec2], - eq_solimp: wp.array2d[vec5], - eq_data: wp.array2d[vec11], - is_sparse: bool, - eq_jnt_adr: wp.array[int], - # Data in: - qpos_in: wp.array2d[float], - qvel_in: wp.array2d[float], - eq_active_in: wp.array2d[bool], - njmax_in: int, - njmax_nnz_in: int, - # Data out: - ne_out: wp.array[int], - nefc_out: wp.array[int], - efc_type_out: wp.array2d[int], - efc_id_out: wp.array2d[int], - efc_J_rownnz_out: wp.array2d[int], - efc_J_rowadr_out: wp.array2d[int], - efc_J_colind_out: wp.array3d[int], - efc_J_out: wp.array3d[float], - efc_pos_out: wp.array2d[float], - efc_margin_out: wp.array2d[float], - efc_D_out: wp.array2d[float], - efc_vel_out: wp.array2d[float], - efc_aref_out: wp.array2d[float], - efc_frictionloss_out: wp.array2d[float], - # Out: - efc_nnz_out: wp.array[int], -): - worldid, eqjntid = wp.tid() - eqid = eq_jnt_adr[eqjntid] - - if not eq_active_in[worldid, eqid]: - return - - wp.atomic_add(ne_out, worldid, 1) - efcid = wp.atomic_add(nefc_out, worldid, 1) - - if efcid >= njmax_in: - return - - jntid_1 = eq_obj1id[eqid] - jntid_2 = eq_obj2id[eqid] - data = eq_data[worldid % eq_data.shape[0], eqid] - dofadr1 = jnt_dofadr[jntid_1] - qposadr1 = jnt_qposadr[jntid_1] - qpos0_id = worldid % qpos0.shape[0] - dof_invweight0_id = worldid % dof_invweight0.shape[0] - - if is_sparse: - if jntid_2 > -1: - rownnz = 2 - else: - rownnz = 1 - efc_J_rownnz_out[worldid, efcid] = rownnz - rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz) - if rowadr + rownnz > njmax_nnz_in: - return - efc_J_rowadr_out[worldid, efcid] = rowadr - efc_J_colind_out[worldid, 0, rowadr] = dofadr1 - efc_J_out[worldid, 0, rowadr] = 1.0 - else: - for i in range(nv): - efc_J_out[worldid, efcid, i] = 0.0 - efc_J_out[worldid, efcid, dofadr1] = 1.0 - - if jntid_2 > -1: - # Two joint constraint - qposadr2 = jnt_qposadr[jntid_2] - dofadr2 = jnt_dofadr[jntid_2] - dif = qpos_in[worldid, qposadr2] - qpos0[qpos0_id, qposadr2] - - # Horner's method for polynomials - rhs = data[0] + dif * (data[1] + dif * (data[2] + dif * (data[3] + dif * data[4]))) - deriv_2 = data[1] + dif * (2.0 * data[2] + dif * (3.0 * data[3] + dif * 4.0 * data[4])) - - pos = qpos_in[worldid, qposadr1] - qpos0[qpos0_id, qposadr1] - rhs - Jqvel = qvel_in[worldid, dofadr1] - qvel_in[worldid, dofadr2] * deriv_2 - invweight = dof_invweight0[dof_invweight0_id, dofadr1] + dof_invweight0[dof_invweight0_id, dofadr2] - - if is_sparse: - sparseid = rowadr + 1 - efc_J_colind_out[worldid, 0, sparseid] = dofadr2 - efc_J_out[worldid, 0, sparseid] = -deriv_2 - else: - efc_J_out[worldid, efcid, dofadr2] = -deriv_2 - else: - # Single joint constraint - pos = qpos_in[worldid, qposadr1] - qpos0[qpos0_id, qposadr1] - data[0] - Jqvel = qvel_in[worldid, dofadr1] - invweight = dof_invweight0[dof_invweight0_id, dofadr1] - - # Update constraint parameters - _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - pos, - pos, - invweight, - eq_solref[worldid % eq_solref.shape[0], eqid], - eq_solimp[worldid % eq_solimp.shape[0], eqid], - 0.0, - Jqvel, - 0.0, - ConstraintType.EQUALITY, - eqid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, - ) - - -@wp.kernel -def _equality_tendon( - # Model: - nv: int, - opt_timestep: wp.array[float], - opt_disableflags: int, - eq_obj1id: wp.array[int], - eq_obj2id: wp.array[int], - eq_solref: wp.array2d[wp.vec2], - eq_solimp: wp.array2d[vec5], - eq_data: wp.array2d[vec11], - ten_J_rownnz: wp.array[int], - ten_J_rowadr: wp.array[int], - ten_J_colind: wp.array[int], - tendon_length0: wp.array2d[float], - tendon_invweight0: wp.array2d[float], - is_sparse: bool, - eq_ten_adr: wp.array[int], - # Data in: - qvel_in: wp.array2d[float], - eq_active_in: wp.array2d[bool], - ten_J_in: wp.array2d[float], - ten_length_in: wp.array2d[float], - njmax_in: int, - njmax_nnz_in: int, - # Data out: - ne_out: wp.array[int], - nefc_out: wp.array[int], - efc_type_out: wp.array2d[int], - efc_id_out: wp.array2d[int], - efc_J_rownnz_out: wp.array2d[int], - efc_J_rowadr_out: wp.array2d[int], - efc_J_colind_out: wp.array3d[int], - efc_J_out: wp.array3d[float], - efc_pos_out: wp.array2d[float], - efc_margin_out: wp.array2d[float], - efc_D_out: wp.array2d[float], - efc_vel_out: wp.array2d[float], - efc_aref_out: wp.array2d[float], - efc_frictionloss_out: wp.array2d[float], - # Out: - efc_nnz_out: wp.array[int], -): - worldid, eqtenid = wp.tid() - eqid = eq_ten_adr[eqtenid] - - if not eq_active_in[worldid, eqid]: - return - - wp.atomic_add(ne_out, worldid, 1) - efcid = wp.atomic_add(nefc_out, worldid, 1) - - if efcid >= njmax_in: - return - - obj1id = eq_obj1id[eqid] - obj2id = eq_obj2id[eqid] - - data = eq_data[worldid % eq_data.shape[0], eqid] - solref = eq_solref[worldid % eq_solref.shape[0], eqid] - solimp = eq_solimp[worldid % eq_solimp.shape[0], eqid] - tendon_length0_id = worldid % tendon_length0.shape[0] - tendon_invweight0_id = worldid % tendon_invweight0.shape[0] - pos1 = ten_length_in[worldid, obj1id] - tendon_length0[tendon_length0_id, obj1id] - - if obj2id > -1: - invweight = tendon_invweight0[tendon_invweight0_id, obj1id] + tendon_invweight0[tendon_invweight0_id, obj2id] - - pos2 = ten_length_in[worldid, obj2id] - tendon_length0[tendon_length0_id, obj2id] - - dif = pos2 - dif2 = dif * dif - dif3 = dif2 * dif - dif4 = dif3 * dif - - pos = pos1 - (data[0] + data[1] * dif + data[2] * dif2 + data[3] * dif3 + data[4] * dif4) - deriv = data[1] + 2.0 * data[2] * dif + 3.0 * data[3] * dif2 + 4.0 * data[4] * dif3 - else: - invweight = tendon_invweight0[tendon_invweight0_id, obj1id] - pos = pos1 - data[0] - deriv = 0.0 - - rownnz1 = ten_J_rownnz[obj1id] - rowadr1 = ten_J_rowadr[obj1id] - rownnz2 = 0 - rowadr2 = 0 - - if deriv != 0.0: - rownnz2 = ten_J_rownnz[obj2id] - rowadr2 = ten_J_rowadr[obj2id] - - if is_sparse: - # TODO(team): pre-compute rownnz - # count unique dofs - p1, p2 = int(0), int(0) - rownnz = int(0) - while p1 < rownnz1 or p2 < rownnz2: - col1 = nv - col2 = nv - if p1 < rownnz1: - col1 = ten_J_colind[rowadr1 + p1] - if p2 < rownnz2: - col2 = ten_J_colind[rowadr2 + p2] - if col1 <= col2: - p1 += 1 - if col2 <= col1: - p2 += 1 - rownnz += 1 - - rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz) - if rowadr + rownnz > njmax_nnz_in: - return - efc_J_rowadr_out[worldid, efcid] = rowadr - - ptr1 = int(0) - ptr2 = int(0) - - Jqvel = float(0.0) - - nnz = int(0) - for i in range(nv): - J1 = float(0.0) - if ptr1 < rownnz1: - sparseid1 = rowadr1 + ptr1 - if ten_J_colind[sparseid1] == i: - J1 = ten_J_in[worldid, sparseid1] - ptr1 += 1 - - J = J1 - if deriv != 0.0: - J2 = float(0.0) - if ptr2 < rownnz2: - sparseid2 = rowadr2 + ptr2 - if ten_J_colind[sparseid2] == i: - J2 = ten_J_in[worldid, sparseid2] - ptr2 += 1 - J += J2 * -deriv - - if is_sparse: - if J != 0.0: - sparseid = rowadr + nnz - efc_J_colind_out[worldid, 0, sparseid] = i - efc_J_out[worldid, 0, sparseid] = J - nnz += 1 - else: - efc_J_out[worldid, efcid, i] = J - - Jqvel += J * qvel_in[worldid, i] - - if is_sparse: - efc_J_rownnz_out[worldid, efcid] = nnz - - _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - pos, - pos, - invweight, - solref, - solimp, - 0.0, - Jqvel, - 0.0, - ConstraintType.EQUALITY, - eqid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, - ) + return kernel @cache_kernel -def _equality_flex(is_sparse: bool): +def _equality_tendon(is_sparse: bool, newton: bool): + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + nv: int, + opt_timestep: wp.array[float], + opt_disableflags: int, + eq_obj1id: wp.array[int], + eq_obj2id: wp.array[int], + eq_solref: wp.array2d[wp.vec2], + eq_solimp: wp.array2d[vec5], + eq_data: wp.array2d[vec11], + ten_J_rownnz: wp.array[int], + ten_J_rowadr: wp.array[int], + ten_J_colind: wp.array[int], + tendon_length0: wp.array2d[float], + tendon_invweight0: wp.array2d[float], + eq_ten_adr: wp.array[int], + # Data in: + qvel_in: wp.array2d[float], + eq_active_in: wp.array2d[bool], + ten_J_in: wp.array2d[float], + ten_length_in: wp.array2d[float], + njmax_in: int, + njmax_nnz_in: int, + # Data out: + ne_out: wp.array[int], + nefc_out: wp.array[int], + efc_type_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_jtdaj_adr_out: wp.array2d[int], + efc_jtdaj_nrow_out: wp.array2d[int], + efc_jtdaj_nblock_out: wp.array[int], + efc_J_rownnz_out: wp.array2d[int], + efc_J_rowadr_out: wp.array2d[int], + efc_J_colind_out: wp.array3d[int], + efc_J_out: wp.array3d[float], + efc_pos_out: wp.array2d[float], + efc_margin_out: wp.array2d[float], + efc_D_out: wp.array2d[float], + efc_vel_out: wp.array2d[float], + efc_aref_out: wp.array2d[float], + efc_frictionloss_out: wp.array2d[float], + # Out: + efc_nnz_out: wp.array[int], + ): + worldid, eqtenid = wp.tid() + eqid = eq_ten_adr[eqtenid] + + if not eq_active_in[worldid, eqid]: + return + + wp.atomic_add(ne_out, worldid, 1) + efcid = wp.atomic_add(nefc_out, worldid, 1) + + if efcid >= njmax_in: + return + + if wp.static(is_sparse and newton): + jgid = wp.atomic_add(efc_jtdaj_nblock_out, worldid, 1) + efc_jtdaj_adr_out[worldid, jgid] = efcid + efc_jtdaj_nrow_out[worldid, jgid] = 1 + + obj1id = eq_obj1id[eqid] + obj2id = eq_obj2id[eqid] + + data = eq_data[worldid % eq_data.shape[0], eqid] + solref = eq_solref[worldid % eq_solref.shape[0], eqid] + solimp = eq_solimp[worldid % eq_solimp.shape[0], eqid] + tendon_length0_id = worldid % tendon_length0.shape[0] + tendon_invweight0_id = worldid % tendon_invweight0.shape[0] + pos1 = ten_length_in[worldid, obj1id] - tendon_length0[tendon_length0_id, obj1id] + + if obj2id > -1: + invweight = tendon_invweight0[tendon_invweight0_id, obj1id] + tendon_invweight0[tendon_invweight0_id, obj2id] + + pos2 = ten_length_in[worldid, obj2id] - tendon_length0[tendon_length0_id, obj2id] + + dif = pos2 + dif2 = dif * dif + dif3 = dif2 * dif + dif4 = dif3 * dif + + pos = pos1 - (data[0] + data[1] * dif + data[2] * dif2 + data[3] * dif3 + data[4] * dif4) + deriv = data[1] + 2.0 * data[2] * dif + 3.0 * data[3] * dif2 + 4.0 * data[4] * dif3 + else: + invweight = tendon_invweight0[tendon_invweight0_id, obj1id] + pos = pos1 - data[0] + deriv = 0.0 + + rownnz1 = ten_J_rownnz[obj1id] + rowadr1 = ten_J_rowadr[obj1id] + rownnz2 = 0 + rowadr2 = 0 + + if deriv != 0.0: + rownnz2 = ten_J_rownnz[obj2id] + rowadr2 = ten_J_rowadr[obj2id] + + if wp.static(is_sparse): + # TODO(team): pre-compute rownnz + # count unique dofs + p1, p2 = int(0), int(0) + rownnz = int(0) + while p1 < rownnz1 or p2 < rownnz2: + col1 = nv + col2 = nv + if p1 < rownnz1: + col1 = ten_J_colind[rowadr1 + p1] + if p2 < rownnz2: + col2 = ten_J_colind[rowadr2 + p2] + if col1 <= col2: + p1 += 1 + if col2 <= col1: + p2 += 1 + rownnz += 1 + + rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz) + if rowadr + rownnz > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid] = rowadr + + ptr1 = int(0) + ptr2 = int(0) + + Jqvel = float(0.0) + + nnz = int(0) + for i in range(nv): + J1 = float(0.0) + if ptr1 < rownnz1: + sparseid1 = rowadr1 + ptr1 + if ten_J_colind[sparseid1] == i: + J1 = ten_J_in[worldid, sparseid1] + ptr1 += 1 + + J = J1 + if deriv != 0.0: + J2 = float(0.0) + if ptr2 < rownnz2: + sparseid2 = rowadr2 + ptr2 + if ten_J_colind[sparseid2] == i: + J2 = ten_J_in[worldid, sparseid2] + ptr2 += 1 + J += J2 * -deriv + + if wp.static(is_sparse): + if J != 0.0: + sparseid = rowadr + nnz + efc_J_colind_out[worldid, 0, sparseid] = i + efc_J_out[worldid, 0, sparseid] = J + nnz += 1 + else: + efc_J_out[worldid, efcid, i] = J + + Jqvel += J * qvel_in[worldid, i] + + if wp.static(is_sparse): + efc_J_rownnz_out[worldid, efcid] = nnz + + _efc_row( + opt_disableflags, + worldid, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + pos, + pos, + invweight, + solref, + solimp, + 0.0, + Jqvel, + 0.0, + ConstraintType.EQUALITY, + eqid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + return kernel + + +@cache_kernel +def _equality_flex(is_sparse: bool, newton: bool): @wp.kernel(module="unique", enable_backward=False) def kernel( # Model: @@ -794,6 +831,9 @@ def _equality_flex(is_sparse: bool): nefc_out: wp.array[int], efc_type_out: wp.array2d[int], efc_id_out: wp.array2d[int], + efc_jtdaj_adr_out: wp.array2d[int], + efc_jtdaj_nrow_out: wp.array2d[int], + efc_jtdaj_nblock_out: wp.array[int], efc_J_rownnz_out: wp.array2d[int], efc_J_rowadr_out: wp.array2d[int], efc_J_colind_out: wp.array3d[int], @@ -823,6 +863,11 @@ def _equality_flex(is_sparse: bool): if efcid >= njmax_in: return + if wp.static(is_sparse and newton): + jgid = wp.atomic_add(efc_jtdaj_nblock_out, worldid, 1) + efc_jtdaj_adr_out[worldid, jgid] = efcid + efc_jtdaj_nrow_out[worldid, jgid] = 1 + pos = flexedge_length_in[worldid, edgeid] - flexedge_length0[edgeid] solref = eq_solref[worldid % eq_solref.shape[0], eqid] solimp = eq_solimp[worldid % eq_solimp.shape[0], eqid] @@ -884,769 +929,571 @@ def _equality_flex(is_sparse: bool): return kernel -@wp.kernel -def _equality_weld( - # Model: - nv: int, - nsite: int, - opt_timestep: wp.array[float], - opt_disableflags: int, - body_parentid: wp.array[int], - body_rootid: wp.array[int], - body_weldid: wp.array[int], - body_dofnum: wp.array[int], - body_dofadr: wp.array[int], - body_invweight0: wp.array2d[wp.vec2], - jnt_type: wp.array[int], - jnt_dofadr: wp.array[int], - dof_bodyid: wp.array[int], - dof_jntid: wp.array[int], - dof_parentid: wp.array[int], - site_bodyid: wp.array[int], - site_quat: wp.array2d[wp.quat], - eq_obj1id: wp.array[int], - eq_obj2id: wp.array[int], - eq_objtype: wp.array[int], - eq_solref: wp.array2d[wp.vec2], - eq_solimp: wp.array2d[vec5], - eq_data: wp.array2d[vec11], - is_sparse: bool, - body_isdofancestor: wp.array2d[int], - eq_wld_adr: wp.array[int], - # Data in: - qvel_in: wp.array2d[float], - eq_active_in: wp.array2d[bool], - xpos_in: wp.array2d[wp.vec3], - xquat_in: wp.array2d[wp.quat], - xmat_in: wp.array2d[wp.mat33], - site_xpos_in: wp.array2d[wp.vec3], - subtree_com_in: wp.array2d[wp.vec3], - cdof_in: wp.array2d[wp.spatial_vector], - cvel_in: wp.array2d[wp.spatial_vector], - cdof_dot_in: wp.array2d[wp.spatial_vector], - subtree_linvel_in: wp.array2d[wp.vec3], - njmax_in: int, - njmax_nnz_in: int, - # Data out: - ne_out: wp.array[int], - nefc_out: wp.array[int], - efc_type_out: wp.array2d[int], - efc_id_out: wp.array2d[int], - efc_J_rownnz_out: wp.array2d[int], - efc_J_rowadr_out: wp.array2d[int], - efc_J_colind_out: wp.array3d[int], - efc_J_out: wp.array3d[float], - efc_pos_out: wp.array2d[float], - efc_margin_out: wp.array2d[float], - efc_D_out: wp.array2d[float], - efc_vel_out: wp.array2d[float], - efc_aref_out: wp.array2d[float], - efc_frictionloss_out: wp.array2d[float], - # Out: - efc_nnz_out: wp.array[int], -): - worldid, eqweldid = wp.tid() - eqid = eq_wld_adr[eqweldid] +@cache_kernel +def _equality_weld(is_sparse: bool, newton: bool): + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + nv: int, + nsite: int, + opt_timestep: wp.array[float], + opt_disableflags: int, + body_parentid: wp.array[int], + body_rootid: wp.array[int], + body_weldid: wp.array[int], + body_dofnum: wp.array[int], + body_dofadr: wp.array[int], + body_invweight0: wp.array2d[wp.vec2], + jnt_type: wp.array[int], + jnt_dofadr: wp.array[int], + dof_bodyid: wp.array[int], + dof_jntid: wp.array[int], + dof_parentid: wp.array[int], + site_bodyid: wp.array[int], + site_quat: wp.array2d[wp.quat], + eq_obj1id: wp.array[int], + eq_obj2id: wp.array[int], + eq_objtype: wp.array[int], + eq_solref: wp.array2d[wp.vec2], + eq_solimp: wp.array2d[vec5], + eq_data: wp.array2d[vec11], + body_isdofancestor: wp.array2d[int], + eq_wld_adr: wp.array[int], + # Data in: + qvel_in: wp.array2d[float], + eq_active_in: wp.array2d[bool], + xpos_in: wp.array2d[wp.vec3], + xquat_in: wp.array2d[wp.quat], + xmat_in: wp.array2d[wp.mat33], + site_xpos_in: wp.array2d[wp.vec3], + subtree_com_in: wp.array2d[wp.vec3], + cdof_in: wp.array2d[wp.spatial_vector], + cvel_in: wp.array2d[wp.spatial_vector], + cdof_dot_in: wp.array2d[wp.spatial_vector], + subtree_linvel_in: wp.array2d[wp.vec3], + njmax_in: int, + njmax_nnz_in: int, + # Data out: + ne_out: wp.array[int], + nefc_out: wp.array[int], + efc_type_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_jtdaj_adr_out: wp.array2d[int], + efc_jtdaj_nrow_out: wp.array2d[int], + efc_jtdaj_nblock_out: wp.array[int], + efc_J_rownnz_out: wp.array2d[int], + efc_J_rowadr_out: wp.array2d[int], + efc_J_colind_out: wp.array3d[int], + efc_J_out: wp.array3d[float], + efc_pos_out: wp.array2d[float], + efc_margin_out: wp.array2d[float], + efc_D_out: wp.array2d[float], + efc_vel_out: wp.array2d[float], + efc_aref_out: wp.array2d[float], + efc_frictionloss_out: wp.array2d[float], + # Out: + efc_nnz_out: wp.array[int], + ): + worldid, eqweldid = wp.tid() + eqid = eq_wld_adr[eqweldid] - if not eq_active_in[worldid, eqid]: - return - - wp.atomic_add(ne_out, worldid, 6) - efcid = wp.atomic_add(nefc_out, worldid, 6) - - if efcid >= njmax_in - 6: - return - - efcid0 = efcid + 0 - efcid1 = efcid + 1 - efcid2 = efcid + 2 - efcid3 = efcid + 3 - efcid4 = efcid + 4 - efcid5 = efcid + 5 - - is_site = eq_objtype[eqid] == types.ObjType.SITE and nsite > 0 - - obj1id = eq_obj1id[eqid] - obj2id = eq_obj2id[eqid] - - data = eq_data[worldid % eq_data.shape[0], eqid] - anchor1 = wp.vec3(data[0], data[1], data[2]) - anchor2 = wp.vec3(data[3], data[4], data[5]) - relpose = wp.quat(data[6], data[7], data[8], data[9]) - torquescale = data[10] - - if is_site: - body1 = site_bodyid[obj1id] - body2 = site_bodyid[obj2id] - pos1 = site_xpos_in[worldid, obj1id] - pos2 = site_xpos_in[worldid, obj2id] - - site_quat_id = worldid % site_quat.shape[0] - quat = math.mul_quat(xquat_in[worldid, body1], site_quat[site_quat_id, obj1id]) - quat1 = math.quat_inv(math.mul_quat(xquat_in[worldid, body2], site_quat[site_quat_id, obj2id])) - - else: - body1 = obj1id - body2 = obj2id - pos1 = xpos_in[worldid, body1] + xmat_in[worldid, body1] @ anchor2 - pos2 = xpos_in[worldid, body2] + xmat_in[worldid, body2] @ anchor1 - - quat = math.mul_quat(xquat_in[worldid, body1], relpose) - quat1 = math.quat_inv(xquat_in[worldid, body2]) - - # quat1 = quat_inv(xquat_in[worldid, body2]) - q2 = xquat_in[worldid, body2] - quat1 = wp.quat(q2[0], -q2[1], -q2[2], -q2[3]) - - # compute rotational Jdotv helper quaternions - omega1 = wp.spatial_top(cvel_in[worldid, body1]) - omega2 = wp.spatial_top(cvel_in[worldid, body2]) - domega = omega1 - omega2 - - omega1_q = wp.quat(0.0, omega1[0], omega1[1], omega1[2]) - omega2_q = wp.quat(0.0, omega2[0], omega2[1], omega2[2]) - domega_q = wp.quat(0.0, domega[0], domega[1], domega[2]) - - if is_site: - qdot0r = math.mul_quat(omega1_q, quat) * 0.5 - qfull1 = math.mul_quat(xquat_in[worldid, body2], site_quat[site_quat_id, obj2id]) - qdot1 = math.mul_quat(omega2_q, qfull1) * 0.5 - - negqdot1 = wp.quat(-qdot1[0], -qdot1[1], -qdot1[2], -qdot1[3]) - negq1 = wp.quat(qfull1[0], -qfull1[1], -qfull1[2], -qfull1[3]) - - else: - # qdot0 = mul_quat(xquat_in[worldid, body1], omega1_q) * 0.5 - u7 = xquat_in[worldid, body1] - qdot0 = math.mul_quat(omega1_q, xquat_in[worldid, body1]) * 0.5 - qdot0r = math.mul_quat(qdot0, relpose) - q1_non_site = xquat_in[worldid, body2] - qdot1 = math.mul_quat(omega2_q, q1_non_site) * 0.5 - - negqdot1 = wp.quat(-qdot1[0], -qdot1[1], -qdot1[2], -qdot1[3]) - negq1 = wp.quat(q1_non_site[0], -q1_non_site[1], -q1_non_site[2], -q1_non_site[3]) - - # compute Jacobian difference (opposite of contact: 0 - 1) - Jqvelp = wp.vec3f(0.0, 0.0, 0.0) - Jqvelr = wp.vec3f(0.0, 0.0, 0.0) - Jdotv_p = wp.vec3f(0.0, 0.0, 0.0) - Jdotv_r0 = wp.vec3f(0.0, 0.0, 0.0) - - if is_sparse: - # TODO(team): pre-compute number of non-zeros - body1 = body_weldid[body1] - body2 = body_weldid[body2] - - da1 = int(body_dofadr[body1] + body_dofnum[body1] - 1) - da2 = int(body_dofadr[body2] + body_dofnum[body2] - 1) - - # count non-zeros - pda1 = da1 - pda2 = da2 - rownnz = int(0) - while pda1 >= 0 or pda2 >= 0: - da = wp.max(pda1, pda2) - if pda1 == da: - pda1 = dof_parentid[da] - if pda2 == da: - pda2 = dof_parentid[da] - rownnz += 1 - - # get rowadr - rowadr = wp.atomic_add(efc_nnz_out, worldid, 6 * rownnz) - if rowadr + 6 * rownnz > njmax_nnz_in: + if not eq_active_in[worldid, eqid]: return - efc_J_rowadr_out[worldid, efcid0] = rowadr - efc_J_rowadr_out[worldid, efcid1] = rowadr + rownnz - efc_J_rowadr_out[worldid, efcid2] = rowadr + 2 * rownnz - efc_J_rowadr_out[worldid, efcid3] = rowadr + 3 * rownnz - efc_J_rowadr_out[worldid, efcid4] = rowadr + 4 * rownnz - efc_J_rowadr_out[worldid, efcid5] = rowadr + 5 * rownnz - efc_J_rownnz_out[worldid, efcid0] = rownnz - efc_J_rownnz_out[worldid, efcid1] = rownnz - efc_J_rownnz_out[worldid, efcid2] = rownnz - efc_J_rownnz_out[worldid, efcid3] = rownnz - efc_J_rownnz_out[worldid, efcid4] = rownnz - efc_J_rownnz_out[worldid, efcid5] = rownnz + wp.atomic_add(ne_out, worldid, 6) + efcid = wp.atomic_add(nefc_out, worldid, 6) - # compute J and colind - nnz = int(0) - while da1 >= 0 or da2 >= 0: - da = wp.max(da1, da2) - if da1 == da: - da1 = dof_parentid[da] - if da2 == da: - da2 = dof_parentid[da] - - jacp1, jacr1 = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - body_isdofancestor, - subtree_com_in, - cdof_in, - pos1, - body1, - da, - worldid, - ) - jacp2, jacr2 = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - body_isdofancestor, - subtree_com_in, - cdof_in, - pos2, - body2, - da, - worldid, - ) - - jacdifp = jacp1 - jacp2 - - jacdifr = (jacr1 - jacr2) * torquescale - jacdifrq = math.mul_quat(math.quat_mul_axis(quat1, jacdifr), quat) - jacdifr = 0.5 * wp.vec3(jacdifrq[1], jacdifrq[2], jacdifrq[3]) - - jacp1_dot, jacr1_dot = support.jac_dot_dof( - body_parentid, - body_rootid, - jnt_type, - jnt_dofadr, - dof_bodyid, - dof_jntid, - body_isdofancestor, - subtree_com_in, - cdof_in, - cvel_in, - cdof_dot_in, - pos1, - body1, - da, - worldid, - ) - jacp2_dot, jacr2_dot = support.jac_dot_dof( - body_parentid, - body_rootid, - jnt_type, - jnt_dofadr, - dof_bodyid, - dof_jntid, - body_isdofancestor, - subtree_com_in, - cdof_in, - cvel_in, - cdof_dot_in, - pos2, - body2, - da, - worldid, - ) - - jacdifp_dot = jacp1_dot - jacp2_dot - jacdifr_dot = jacr1_dot - jacr2_dot - - sparseid0 = rowadr + nnz - sparseid1 = rowadr + rownnz + nnz - sparseid2 = rowadr + 2 * rownnz + nnz - sparseid3 = rowadr + 3 * rownnz + nnz - sparseid4 = rowadr + 4 * rownnz + nnz - sparseid5 = rowadr + 5 * rownnz + nnz - - efc_J_colind_out[worldid, 0, sparseid0] = da - efc_J_colind_out[worldid, 0, sparseid1] = da - efc_J_colind_out[worldid, 0, sparseid2] = da - efc_J_colind_out[worldid, 0, sparseid3] = da - efc_J_colind_out[worldid, 0, sparseid4] = da - efc_J_colind_out[worldid, 0, sparseid5] = da - - efc_J_out[worldid, 0, sparseid0] = jacdifp[0] - efc_J_out[worldid, 0, sparseid1] = jacdifp[1] - efc_J_out[worldid, 0, sparseid2] = jacdifp[2] - efc_J_out[worldid, 0, sparseid3] = jacdifr[0] - efc_J_out[worldid, 0, sparseid4] = jacdifr[1] - efc_J_out[worldid, 0, sparseid5] = jacdifr[2] - - Jqvelp += jacdifp * qvel_in[worldid, da] - Jqvelr += jacdifr * qvel_in[worldid, da] - Jdotv_p += jacdifp_dot * qvel_in[worldid, da] - Jdotv_r0 += jacdifr_dot * qvel_in[worldid, da] - - nnz += 1 - else: - for dofid in range(nv): - jacp1, jacr1 = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - body_isdofancestor, - subtree_com_in, - cdof_in, - pos1, - body1, - dofid, - worldid, - ) - jacp2, jacr2 = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - body_isdofancestor, - subtree_com_in, - cdof_in, - pos2, - body2, - dofid, - worldid, - ) - - jacdifp = jacp1 - jacp2 - - efc_J_out[worldid, efcid0, dofid] = jacdifp[0] - efc_J_out[worldid, efcid1, dofid] = jacdifp[1] - efc_J_out[worldid, efcid2, dofid] = jacdifp[2] - - jacdifr = (jacr1 - jacr2) * torquescale - jacdifrq = math.mul_quat(math.quat_mul_axis(quat1, jacdifr), quat) - jacdifr = 0.5 * wp.vec3(jacdifrq[1], jacdifrq[2], jacdifrq[3]) - - jacp1_dot, jacr1_dot = support.jac_dot_dof( - body_parentid, - body_rootid, - jnt_type, - jnt_dofadr, - dof_bodyid, - dof_jntid, - body_isdofancestor, - subtree_com_in, - cdof_in, - cvel_in, - cdof_dot_in, - pos1, - body1, - dofid, - worldid, - ) - jacp2_dot, jacr2_dot = support.jac_dot_dof( - body_parentid, - body_rootid, - jnt_type, - jnt_dofadr, - dof_bodyid, - dof_jntid, - body_isdofancestor, - subtree_com_in, - cdof_in, - cvel_in, - cdof_dot_in, - pos2, - body2, - dofid, - worldid, - ) - - jacdifp_dot = jacp1_dot - jacp2_dot - jacdifr_dot = jacr1_dot - jacr2_dot - - efc_J_out[worldid, efcid3, dofid] = jacdifr[0] - efc_J_out[worldid, efcid4, dofid] = jacdifr[1] - efc_J_out[worldid, efcid5, dofid] = jacdifr[2] - - Jqvelp += jacdifp * qvel_in[worldid, dofid] - Jqvelr += jacdifr * qvel_in[worldid, dofid] - Jdotv_p += jacdifp_dot * qvel_in[worldid, dofid] - Jdotv_r0 += jacdifr_dot * qvel_in[worldid, dofid] - - # error is difference in global position and orientation - cpos = pos1 - pos2 - - crotq = math.mul_quat(quat1, quat) # copy axis components - crot = wp.vec3(crotq[1], crotq[2], crotq[3]) * torquescale - - body_invweight0_id = worldid % body_invweight0.shape[0] - invweight_t = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, body2][0] - - pos_imp = wp.sqrt(wp.length_sq(cpos) + wp.length_sq(crot)) - - solref = eq_solref[worldid % eq_solref.shape[0], eqid] - solimp = eq_solimp[worldid % eq_solimp.shape[0], eqid] - - timestep = opt_timestep[worldid % opt_timestep.shape[0]] - - djrdv_q = wp.quat(0.0, Jdotv_r0[0], Jdotv_r0[1], Jdotv_r0[2]) - - # Term 1: negqdot1 * domega * q0r - t1a = math.mul_quat(negqdot1, domega_q) - t1 = math.mul_quat(t1a, quat) - - # Term 2: negq1 * djrdv * q0r - t2a = math.mul_quat(negq1, djrdv_q) - t2 = math.mul_quat(t2a, quat) - - # Term 3: negq1 * domega * qdot0r - t3a = math.mul_quat(negq1, domega_q) - t3 = math.mul_quat(t3a, qdot0r) - - Jdotv_r = wp.vec3(t1[1] + t2[1] + t3[1], t1[2] + t2[2] + t3[2], t1[3] + t2[3] + t3[3]) * 0.5 * torquescale - - for i in range(3): - _efc_row( - opt_disableflags, - worldid, - timestep, - efcid + i, - cpos[i], - pos_imp, - invweight_t, - solref, - solimp, - 0.0, - Jqvelp[i], - 0.0, - ConstraintType.EQUALITY, - eqid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, - ) - - efc_aref_out[worldid, efcid + i] -= Jdotv_p[i] - - invweight_r = body_invweight0[body_invweight0_id, body1][1] + body_invweight0[body_invweight0_id, body2][1] - - for i in range(3): - _efc_row( - opt_disableflags, - worldid, - timestep, - efcid + 3 + i, - crot[i], - pos_imp, - invweight_r, - solref, - solimp, - 0.0, - Jqvelr[i], - 0.0, - ConstraintType.EQUALITY, - eqid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, - ) - - efc_aref_out[worldid, efcid + 3 + i] -= Jdotv_r[i] - - -@wp.kernel -def _friction_dof( - # Model: - nv: int, - opt_timestep: wp.array[float], - opt_disableflags: int, - dof_solref: wp.array2d[wp.vec2], - dof_solimp: wp.array2d[vec5], - dof_frictionloss: wp.array2d[float], - dof_invweight0: wp.array2d[float], - is_sparse: bool, - # Data in: - qvel_in: wp.array2d[float], - njmax_in: int, - njmax_nnz_in: int, - # Data out: - nf_out: wp.array[int], - nefc_out: wp.array[int], - efc_type_out: wp.array2d[int], - efc_id_out: wp.array2d[int], - efc_J_rownnz_out: wp.array2d[int], - efc_J_rowadr_out: wp.array2d[int], - efc_J_colind_out: wp.array3d[int], - efc_J_out: wp.array3d[float], - efc_pos_out: wp.array2d[float], - efc_margin_out: wp.array2d[float], - efc_D_out: wp.array2d[float], - efc_vel_out: wp.array2d[float], - efc_aref_out: wp.array2d[float], - efc_frictionloss_out: wp.array2d[float], - # Out: - efc_nnz_out: wp.array[int], -): - worldid, dofid = wp.tid() - - dof_frictionloss_id = worldid % dof_frictionloss.shape[0] - - if dof_frictionloss[dof_frictionloss_id, dofid] <= 0.0: - return - - wp.atomic_add(nf_out, worldid, 1) - efcid = wp.atomic_add(nefc_out, worldid, 1) - - if efcid >= njmax_in: - return - - if is_sparse: - efc_J_rownnz_out[worldid, efcid] = 1 - rowadr = wp.atomic_add(efc_nnz_out, worldid, 1) - if rowadr + 1 > njmax_nnz_in: + if efcid >= njmax_in - 6: return - efc_J_rowadr_out[worldid, efcid] = rowadr - efc_J_colind_out[worldid, 0, rowadr] = dofid - efc_J_out[worldid, 0, rowadr] = 1.0 - else: - for i in range(nv): - efc_J_out[worldid, efcid, i] = 0.0 - efc_J_out[worldid, efcid, dofid] = 1.0 - Jqvel = qvel_in[worldid, dofid] + if wp.static(is_sparse and newton): + jgid = wp.atomic_add(efc_jtdaj_nblock_out, worldid, 1) + efc_jtdaj_adr_out[worldid, jgid] = efcid + efc_jtdaj_nrow_out[worldid, jgid] = 6 - dof_invweight0_id = worldid % dof_invweight0.shape[0] - dof_solref_id = worldid % dof_solref.shape[0] - dof_solimp_id = worldid % dof_solimp.shape[0] - _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - 0.0, - 0.0, - dof_invweight0[dof_invweight0_id, dofid], - dof_solref[dof_solref_id, dofid], - dof_solimp[dof_solimp_id, dofid], - 0.0, - Jqvel, - dof_frictionloss[dof_frictionloss_id, dofid], - ConstraintType.FRICTION_DOF, - dofid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, - ) + efcid0 = efcid + 0 + efcid1 = efcid + 1 + efcid2 = efcid + 2 + efcid3 = efcid + 3 + efcid4 = efcid + 4 + efcid5 = efcid + 5 + is_site = eq_objtype[eqid] == types.ObjType.SITE and nsite > 0 -@wp.kernel -def _friction_tendon( - # Model: - nv: int, - opt_timestep: wp.array[float], - opt_disableflags: int, - ten_J_rownnz: wp.array[int], - ten_J_rowadr: wp.array[int], - ten_J_colind: wp.array[int], - tendon_solref_fri: wp.array2d[wp.vec2], - tendon_solimp_fri: wp.array2d[vec5], - tendon_frictionloss: wp.array2d[float], - tendon_invweight0: wp.array2d[float], - is_sparse: bool, - # Data in: - qvel_in: wp.array2d[float], - ten_J_in: wp.array2d[float], - njmax_in: int, - njmax_nnz_in: int, - # Data out: - nf_out: wp.array[int], - nefc_out: wp.array[int], - efc_type_out: wp.array2d[int], - efc_id_out: wp.array2d[int], - efc_J_rownnz_out: wp.array2d[int], - efc_J_rowadr_out: wp.array2d[int], - efc_J_colind_out: wp.array3d[int], - efc_J_out: wp.array3d[float], - efc_pos_out: wp.array2d[float], - efc_margin_out: wp.array2d[float], - efc_D_out: wp.array2d[float], - efc_vel_out: wp.array2d[float], - efc_aref_out: wp.array2d[float], - efc_frictionloss_out: wp.array2d[float], - # Out: - efc_nnz_out: wp.array[int], -): - worldid, tenid = wp.tid() + obj1id = eq_obj1id[eqid] + obj2id = eq_obj2id[eqid] - tendon_frictionloss_id = worldid % tendon_frictionloss.shape[0] + data = eq_data[worldid % eq_data.shape[0], eqid] + anchor1 = wp.vec3(data[0], data[1], data[2]) + anchor2 = wp.vec3(data[3], data[4], data[5]) + relpose = wp.quat(data[6], data[7], data[8], data[9]) + torquescale = data[10] - frictionloss = tendon_frictionloss[tendon_frictionloss_id, tenid] - if frictionloss <= 0.0: - return + if is_site: + body1 = site_bodyid[obj1id] + body2 = site_bodyid[obj2id] + pos1 = site_xpos_in[worldid, obj1id] + pos2 = site_xpos_in[worldid, obj2id] - wp.atomic_add(nf_out, worldid, 1) - efcid = wp.atomic_add(nefc_out, worldid, 1) + site_quat_id = worldid % site_quat.shape[0] + quat = math.mul_quat(xquat_in[worldid, body1], site_quat[site_quat_id, obj1id]) + quat1 = math.quat_inv(math.mul_quat(xquat_in[worldid, body2], site_quat[site_quat_id, obj2id])) - if efcid >= njmax_in: - return + else: + body1 = obj1id + body2 = obj2id + pos1 = xpos_in[worldid, body1] + xmat_in[worldid, body1] @ anchor2 + pos2 = xpos_in[worldid, body2] + xmat_in[worldid, body2] @ anchor1 - Jqvel = float(0.0) + quat = math.mul_quat(xquat_in[worldid, body1], relpose) + quat1 = math.quat_inv(xquat_in[worldid, body2]) - rownnz_tenJ = ten_J_rownnz[tenid] - rowadr_tenJ = ten_J_rowadr[tenid] - if is_sparse: - efc_J_rownnz_out[worldid, efcid] = rownnz_tenJ - rowadr_efc = wp.atomic_add(efc_nnz_out, worldid, rownnz_tenJ) - if rowadr_efc + rownnz_tenJ > njmax_nnz_in: - return - efc_J_rowadr_out[worldid, efcid] = rowadr_efc + # quat1 = quat_inv(xquat_in[worldid, body2]) + q2 = xquat_in[worldid, body2] + quat1 = wp.quat(q2[0], -q2[1], -q2[2], -q2[3]) + + # compute rotational Jdotv helper quaternions + omega1 = wp.spatial_top(cvel_in[worldid, body1]) + omega2 = wp.spatial_top(cvel_in[worldid, body2]) + domega = omega1 - omega2 + + omega1_q = wp.quat(0.0, omega1[0], omega1[1], omega1[2]) + omega2_q = wp.quat(0.0, omega2[0], omega2[1], omega2[2]) + domega_q = wp.quat(0.0, domega[0], domega[1], domega[2]) + + if is_site: + qdot0r = math.mul_quat(omega1_q, quat) * 0.5 + qfull1 = math.mul_quat(xquat_in[worldid, body2], site_quat[site_quat_id, obj2id]) + qdot1 = math.mul_quat(omega2_q, qfull1) * 0.5 + + negqdot1 = wp.quat(-qdot1[0], -qdot1[1], -qdot1[2], -qdot1[3]) + negq1 = wp.quat(qfull1[0], -qfull1[1], -qfull1[2], -qfull1[3]) + + else: + # qdot0 = mul_quat(xquat_in[worldid, body1], omega1_q) * 0.5 + u7 = xquat_in[worldid, body1] + qdot0 = math.mul_quat(omega1_q, xquat_in[worldid, body1]) * 0.5 + qdot0r = math.mul_quat(qdot0, relpose) + q1_non_site = xquat_in[worldid, body2] + qdot1 = math.mul_quat(omega2_q, q1_non_site) * 0.5 + + negqdot1 = wp.quat(-qdot1[0], -qdot1[1], -qdot1[2], -qdot1[3]) + negq1 = wp.quat(q1_non_site[0], -q1_non_site[1], -q1_non_site[2], -q1_non_site[3]) + + # compute Jacobian difference (opposite of contact: 0 - 1) + Jqvelp = wp.vec3f(0.0, 0.0, 0.0) + Jqvelr = wp.vec3f(0.0, 0.0, 0.0) + Jdotv_p = wp.vec3f(0.0, 0.0, 0.0) + Jdotv_r0 = wp.vec3f(0.0, 0.0, 0.0) + + if wp.static(is_sparse): + # TODO(team): pre-compute number of non-zeros + body1 = body_weldid[body1] + body2 = body_weldid[body2] + + da1 = int(body_dofadr[body1] + body_dofnum[body1] - 1) + da2 = int(body_dofadr[body2] + body_dofnum[body2] - 1) + + # count non-zeros + pda1 = da1 + pda2 = da2 + rownnz = int(0) + while pda1 >= 0 or pda2 >= 0: + da = wp.max(pda1, pda2) + if pda1 == da: + pda1 = dof_parentid[da] + if pda2 == da: + pda2 = dof_parentid[da] + rownnz += 1 + + # get rowadr + rowadr = wp.atomic_add(efc_nnz_out, worldid, 6 * rownnz) + if rowadr + 6 * rownnz > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid0] = rowadr + efc_J_rowadr_out[worldid, efcid1] = rowadr + rownnz + efc_J_rowadr_out[worldid, efcid2] = rowadr + 2 * rownnz + efc_J_rowadr_out[worldid, efcid3] = rowadr + 3 * rownnz + efc_J_rowadr_out[worldid, efcid4] = rowadr + 4 * rownnz + efc_J_rowadr_out[worldid, efcid5] = rowadr + 5 * rownnz + + efc_J_rownnz_out[worldid, efcid0] = rownnz + efc_J_rownnz_out[worldid, efcid1] = rownnz + efc_J_rownnz_out[worldid, efcid2] = rownnz + efc_J_rownnz_out[worldid, efcid3] = rownnz + efc_J_rownnz_out[worldid, efcid4] = rownnz + efc_J_rownnz_out[worldid, efcid5] = rownnz + + # compute J and colind + nnz = int(0) + while da1 >= 0 or da2 >= 0: + da = wp.max(da1, da2) + if da1 == da: + da1 = dof_parentid[da] + if da2 == da: + da2 = dof_parentid[da] + + jacp1, jacr1 = support.jac_dof( + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + pos1, + body1, + da, + worldid, + ) + jacp2, jacr2 = support.jac_dof( + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + pos2, + body2, + da, + worldid, + ) + + jacdifp = jacp1 - jacp2 + + jacdifr = (jacr1 - jacr2) * torquescale + jacdifrq = math.mul_quat(math.quat_mul_axis(quat1, jacdifr), quat) + jacdifr = 0.5 * wp.vec3(jacdifrq[1], jacdifrq[2], jacdifrq[3]) + + jacp1_dot, jacr1_dot = support.jac_dot_dof( + body_parentid, + body_rootid, + jnt_type, + jnt_dofadr, + dof_bodyid, + dof_jntid, + body_isdofancestor, + subtree_com_in, + cdof_in, + cvel_in, + cdof_dot_in, + pos1, + body1, + da, + worldid, + ) + jacp2_dot, jacr2_dot = support.jac_dot_dof( + body_parentid, + body_rootid, + jnt_type, + jnt_dofadr, + dof_bodyid, + dof_jntid, + body_isdofancestor, + subtree_com_in, + cdof_in, + cvel_in, + cdof_dot_in, + pos2, + body2, + da, + worldid, + ) + + jacdifp_dot = jacp1_dot - jacp2_dot + jacdifr_dot = jacr1_dot - jacr2_dot + + sparseid0 = rowadr + nnz + sparseid1 = rowadr + rownnz + nnz + sparseid2 = rowadr + 2 * rownnz + nnz + sparseid3 = rowadr + 3 * rownnz + nnz + sparseid4 = rowadr + 4 * rownnz + nnz + sparseid5 = rowadr + 5 * rownnz + nnz + + efc_J_colind_out[worldid, 0, sparseid0] = da + efc_J_colind_out[worldid, 0, sparseid1] = da + efc_J_colind_out[worldid, 0, sparseid2] = da + efc_J_colind_out[worldid, 0, sparseid3] = da + efc_J_colind_out[worldid, 0, sparseid4] = da + efc_J_colind_out[worldid, 0, sparseid5] = da + + efc_J_out[worldid, 0, sparseid0] = jacdifp[0] + efc_J_out[worldid, 0, sparseid1] = jacdifp[1] + efc_J_out[worldid, 0, sparseid2] = jacdifp[2] + efc_J_out[worldid, 0, sparseid3] = jacdifr[0] + efc_J_out[worldid, 0, sparseid4] = jacdifr[1] + efc_J_out[worldid, 0, sparseid5] = jacdifr[2] + + Jqvelp += jacdifp * qvel_in[worldid, da] + Jqvelr += jacdifr * qvel_in[worldid, da] + Jdotv_p += jacdifp_dot * qvel_in[worldid, da] + Jdotv_r0 += jacdifr_dot * qvel_in[worldid, da] - for i in range(rownnz_tenJ): - sparseid_ten = rowadr_tenJ + i - sparseid_efc = rowadr_efc + i - colind = ten_J_colind[sparseid_ten] - J = ten_J_in[worldid, sparseid_ten] - efc_J_colind_out[worldid, 0, sparseid_efc] = colind - efc_J_out[worldid, 0, sparseid_efc] = J - Jqvel += J * qvel_in[worldid, colind] - else: - nnz = int(0) - colind = ten_J_colind[rowadr_tenJ] - for i in range(nv): - if nnz < rownnz_tenJ and i == colind: - J = ten_J_in[worldid, rowadr_tenJ + nnz] - efc_J_out[worldid, efcid, i] = J - Jqvel += J * qvel_in[worldid, i] nnz += 1 - if nnz < rownnz_tenJ: - colind = ten_J_colind[rowadr_tenJ + nnz] - else: - efc_J_out[worldid, efcid, i] = 0.0 + else: + for dofid in range(nv): + jacp1, jacr1 = support.jac_dof( + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + pos1, + body1, + dofid, + worldid, + ) + jacp2, jacr2 = support.jac_dof( + body_parentid, + body_rootid, + dof_bodyid, + body_isdofancestor, + subtree_com_in, + cdof_in, + pos2, + body2, + dofid, + worldid, + ) - tendon_invweight0_id = worldid % tendon_invweight0.shape[0] - tendon_solref_fri_id = worldid % tendon_solref_fri.shape[0] - tendon_solimp_fri_id = worldid % tendon_solimp_fri.shape[0] - _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - 0.0, - 0.0, - tendon_invweight0[tendon_invweight0_id, tenid], - tendon_solref_fri[tendon_solref_fri_id, tenid], - tendon_solimp_fri[tendon_solimp_fri_id, tenid], - 0.0, - Jqvel, - frictionloss, - ConstraintType.FRICTION_TENDON, - tenid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, - ) + jacdifp = jacp1 - jacp2 + + efc_J_out[worldid, efcid0, dofid] = jacdifp[0] + efc_J_out[worldid, efcid1, dofid] = jacdifp[1] + efc_J_out[worldid, efcid2, dofid] = jacdifp[2] + + jacdifr = (jacr1 - jacr2) * torquescale + jacdifrq = math.mul_quat(math.quat_mul_axis(quat1, jacdifr), quat) + jacdifr = 0.5 * wp.vec3(jacdifrq[1], jacdifrq[2], jacdifrq[3]) + + jacp1_dot, jacr1_dot = support.jac_dot_dof( + body_parentid, + body_rootid, + jnt_type, + jnt_dofadr, + dof_bodyid, + dof_jntid, + body_isdofancestor, + subtree_com_in, + cdof_in, + cvel_in, + cdof_dot_in, + pos1, + body1, + dofid, + worldid, + ) + jacp2_dot, jacr2_dot = support.jac_dot_dof( + body_parentid, + body_rootid, + jnt_type, + jnt_dofadr, + dof_bodyid, + dof_jntid, + body_isdofancestor, + subtree_com_in, + cdof_in, + cvel_in, + cdof_dot_in, + pos2, + body2, + dofid, + worldid, + ) + + jacdifp_dot = jacp1_dot - jacp2_dot + jacdifr_dot = jacr1_dot - jacr2_dot + + efc_J_out[worldid, efcid3, dofid] = jacdifr[0] + efc_J_out[worldid, efcid4, dofid] = jacdifr[1] + efc_J_out[worldid, efcid5, dofid] = jacdifr[2] + + Jqvelp += jacdifp * qvel_in[worldid, dofid] + Jqvelr += jacdifr * qvel_in[worldid, dofid] + Jdotv_p += jacdifp_dot * qvel_in[worldid, dofid] + Jdotv_r0 += jacdifr_dot * qvel_in[worldid, dofid] + + # error is difference in global position and orientation + cpos = pos1 - pos2 + + crotq = math.mul_quat(quat1, quat) # copy axis components + crot = wp.vec3(crotq[1], crotq[2], crotq[3]) * torquescale + + body_invweight0_id = worldid % body_invweight0.shape[0] + invweight_t = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, body2][0] + + pos_imp = wp.sqrt(wp.length_sq(cpos) + wp.length_sq(crot)) + + solref = eq_solref[worldid % eq_solref.shape[0], eqid] + solimp = eq_solimp[worldid % eq_solimp.shape[0], eqid] + + timestep = opt_timestep[worldid % opt_timestep.shape[0]] + + djrdv_q = wp.quat(0.0, Jdotv_r0[0], Jdotv_r0[1], Jdotv_r0[2]) + + # Term 1: negqdot1 * domega * q0r + t1a = math.mul_quat(negqdot1, domega_q) + t1 = math.mul_quat(t1a, quat) + + # Term 2: negq1 * djrdv * q0r + t2a = math.mul_quat(negq1, djrdv_q) + t2 = math.mul_quat(t2a, quat) + + # Term 3: negq1 * domega * qdot0r + t3a = math.mul_quat(negq1, domega_q) + t3 = math.mul_quat(t3a, qdot0r) + + Jdotv_r = wp.vec3(t1[1] + t2[1] + t3[1], t1[2] + t2[2] + t3[2], t1[3] + t2[3] + t3[3]) * 0.5 * torquescale + + for i in range(3): + _efc_row( + opt_disableflags, + worldid, + timestep, + efcid + i, + cpos[i], + pos_imp, + invweight_t, + solref, + solimp, + 0.0, + Jqvelp[i], + 0.0, + ConstraintType.EQUALITY, + eqid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + efc_aref_out[worldid, efcid + i] -= Jdotv_p[i] + + invweight_r = body_invweight0[body_invweight0_id, body1][1] + body_invweight0[body_invweight0_id, body2][1] + + for i in range(3): + _efc_row( + opt_disableflags, + worldid, + timestep, + efcid + 3 + i, + crot[i], + pos_imp, + invweight_r, + solref, + solimp, + 0.0, + Jqvelr[i], + 0.0, + ConstraintType.EQUALITY, + eqid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + efc_aref_out[worldid, efcid + 3 + i] -= Jdotv_r[i] + + return kernel -@wp.kernel -def _limit_slide_hinge( - # Model: - nv: int, - opt_timestep: wp.array[float], - opt_disableflags: int, - jnt_qposadr: wp.array[int], - jnt_dofadr: wp.array[int], - jnt_solref: wp.array2d[wp.vec2], - jnt_solimp: wp.array2d[vec5], - jnt_range: wp.array2d[wp.vec2], - jnt_margin: wp.array2d[float], - dof_invweight0: wp.array2d[float], - is_sparse: bool, - jnt_limited_slide_hinge_adr: wp.array[int], - # Data in: - qpos_in: wp.array2d[float], - qvel_in: wp.array2d[float], - njmax_in: int, - njmax_nnz_in: int, - # Data out: - nl_out: wp.array[int], - nefc_out: wp.array[int], - efc_type_out: wp.array2d[int], - efc_id_out: wp.array2d[int], - efc_J_rownnz_out: wp.array2d[int], - efc_J_rowadr_out: wp.array2d[int], - efc_J_colind_out: wp.array3d[int], - efc_J_out: wp.array3d[float], - efc_pos_out: wp.array2d[float], - efc_margin_out: wp.array2d[float], - efc_D_out: wp.array2d[float], - efc_vel_out: wp.array2d[float], - efc_aref_out: wp.array2d[float], - efc_frictionloss_out: wp.array2d[float], - # Out: - efc_nnz_out: wp.array[int], -): - worldid, jntlimitedid = wp.tid() - jntid = jnt_limited_slide_hinge_adr[jntlimitedid] - jnt_range_id = worldid % jnt_range.shape[0] - jntrange = jnt_range[jnt_range_id, jntid] +@cache_kernel +def _friction_dof(is_sparse: bool, newton: bool): + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + nv: int, + opt_timestep: wp.array[float], + opt_disableflags: int, + dof_solref: wp.array2d[wp.vec2], + dof_solimp: wp.array2d[vec5], + dof_frictionloss: wp.array2d[float], + dof_invweight0: wp.array2d[float], + # Data in: + qvel_in: wp.array2d[float], + njmax_in: int, + njmax_nnz_in: int, + # Data out: + nf_out: wp.array[int], + nefc_out: wp.array[int], + efc_type_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_jtdaj_adr_out: wp.array2d[int], + efc_jtdaj_nrow_out: wp.array2d[int], + efc_jtdaj_nblock_out: wp.array[int], + efc_J_rownnz_out: wp.array2d[int], + efc_J_rowadr_out: wp.array2d[int], + efc_J_colind_out: wp.array3d[int], + efc_J_out: wp.array3d[float], + efc_pos_out: wp.array2d[float], + efc_margin_out: wp.array2d[float], + efc_D_out: wp.array2d[float], + efc_vel_out: wp.array2d[float], + efc_aref_out: wp.array2d[float], + efc_frictionloss_out: wp.array2d[float], + # Out: + efc_nnz_out: wp.array[int], + ): + worldid, dofid = wp.tid() - qpos = qpos_in[worldid, jnt_qposadr[jntid]] - jnt_margin_id = worldid % jnt_margin.shape[0] - jntmargin = jnt_margin[jnt_margin_id, jntid] - dist_min, dist_max = qpos - jntrange[0], jntrange[1] - qpos - pos = wp.min(dist_min, dist_max) - jntmargin - active = pos < 0 + dof_frictionloss_id = worldid % dof_frictionloss.shape[0] - if active: - wp.atomic_add(nl_out, worldid, 1) + if dof_frictionloss[dof_frictionloss_id, dofid] <= 0.0: + return + + wp.atomic_add(nf_out, worldid, 1) efcid = wp.atomic_add(nefc_out, worldid, 1) if efcid >= njmax_in: return - dofadr = jnt_dofadr[jntid] + if wp.static(is_sparse and newton): + jgid = wp.atomic_add(efc_jtdaj_nblock_out, worldid, 1) + efc_jtdaj_adr_out[worldid, jgid] = efcid + efc_jtdaj_nrow_out[worldid, jgid] = 1 - J = float(dist_min < dist_max) * 2.0 - 1.0 - - if is_sparse: + if wp.static(is_sparse): efc_J_rownnz_out[worldid, efcid] = 1 rowadr = wp.atomic_add(efc_nnz_out, worldid, 1) if rowadr + 1 > njmax_nnz_in: return efc_J_rowadr_out[worldid, efcid] = rowadr - efc_J_colind_out[worldid, 0, rowadr] = dofadr - efc_J_out[worldid, 0, rowadr] = J + efc_J_colind_out[worldid, 0, rowadr] = dofid + efc_J_out[worldid, 0, rowadr] = 1.0 else: for i in range(nv): efc_J_out[worldid, efcid, i] = 0.0 - efc_J_out[worldid, efcid, dofadr] = J + efc_J_out[worldid, efcid, dofid] = 1.0 - Jqvel = J * qvel_in[worldid, dofadr] + Jqvel = qvel_in[worldid, dofid] dof_invweight0_id = worldid % dof_invweight0.shape[0] - jnt_solref_id = worldid % jnt_solref.shape[0] - jnt_solimp_id = worldid % jnt_solimp.shape[0] + dof_solref_id = worldid % dof_solref.shape[0] + dof_solimp_id = worldid % dof_solimp.shape[0] _efc_row( opt_disableflags, worldid, opt_timestep[worldid % opt_timestep.shape[0]], efcid, - pos, - pos, - dof_invweight0[dof_invweight0_id, dofadr], - jnt_solref[jnt_solref_id, jntid], - jnt_solimp[jnt_solimp_id, jntid], - jntmargin, - Jqvel, 0.0, - ConstraintType.LIMIT_JOINT, - jntid, + 0.0, + dof_invweight0[dof_invweight0_id, dofid], + dof_solref[dof_solref_id, dofid], + dof_solimp[dof_solimp_id, dofid], + 0.0, + Jqvel, + dof_frictionloss[dof_frictionloss_id, dofid], + ConstraintType.FRICTION_DOF, + dofid, efc_type_out, efc_id_out, efc_pos_out, @@ -1657,197 +1504,74 @@ def _limit_slide_hinge( efc_frictionloss_out, ) + return kernel -@wp.kernel -def _limit_ball( - # Model: - nv: int, - opt_timestep: wp.array[float], - opt_disableflags: int, - jnt_qposadr: wp.array[int], - jnt_dofadr: wp.array[int], - jnt_solref: wp.array2d[wp.vec2], - jnt_solimp: wp.array2d[vec5], - jnt_range: wp.array2d[wp.vec2], - jnt_margin: wp.array2d[float], - dof_invweight0: wp.array2d[float], - is_sparse: bool, - jnt_limited_ball_adr: wp.array[int], - # Data in: - qpos_in: wp.array2d[float], - qvel_in: wp.array2d[float], - njmax_in: int, - njmax_nnz_in: int, - # Data out: - nl_out: wp.array[int], - nefc_out: wp.array[int], - efc_type_out: wp.array2d[int], - efc_id_out: wp.array2d[int], - efc_J_rownnz_out: wp.array2d[int], - efc_J_rowadr_out: wp.array2d[int], - efc_J_colind_out: wp.array3d[int], - efc_J_out: wp.array3d[float], - efc_pos_out: wp.array2d[float], - efc_margin_out: wp.array2d[float], - efc_D_out: wp.array2d[float], - efc_vel_out: wp.array2d[float], - efc_aref_out: wp.array2d[float], - efc_frictionloss_out: wp.array2d[float], - # Out: - efc_nnz_out: wp.array[int], -): - worldid, jntlimitedid = wp.tid() - jntid = jnt_limited_ball_adr[jntlimitedid] - qposadr = jnt_qposadr[jntid] - qpos = qpos_in[worldid] - jnt_quat = wp.quat(qpos[qposadr + 0], qpos[qposadr + 1], qpos[qposadr + 2], qpos[qposadr + 3]) - jnt_quat = wp.normalize(jnt_quat) - axis_angle = math.quat_to_vel(jnt_quat) - jnt_range_id = worldid % jnt_range.shape[0] - jntrange = jnt_range[jnt_range_id, jntid] - axis, angle = math.normalize_with_norm(axis_angle) - jnt_margin_id = worldid % jnt_margin.shape[0] - jntmargin = jnt_margin[jnt_margin_id, jntid] +@cache_kernel +def _friction_tendon(is_sparse: bool, newton: bool): + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + nv: int, + opt_timestep: wp.array[float], + opt_disableflags: int, + ten_J_rownnz: wp.array[int], + ten_J_rowadr: wp.array[int], + ten_J_colind: wp.array[int], + tendon_solref_fri: wp.array2d[wp.vec2], + tendon_solimp_fri: wp.array2d[vec5], + tendon_frictionloss: wp.array2d[float], + tendon_invweight0: wp.array2d[float], + # Data in: + qvel_in: wp.array2d[float], + ten_J_in: wp.array2d[float], + njmax_in: int, + njmax_nnz_in: int, + # Data out: + nf_out: wp.array[int], + nefc_out: wp.array[int], + efc_type_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_jtdaj_adr_out: wp.array2d[int], + efc_jtdaj_nrow_out: wp.array2d[int], + efc_jtdaj_nblock_out: wp.array[int], + efc_J_rownnz_out: wp.array2d[int], + efc_J_rowadr_out: wp.array2d[int], + efc_J_colind_out: wp.array3d[int], + efc_J_out: wp.array3d[float], + efc_pos_out: wp.array2d[float], + efc_margin_out: wp.array2d[float], + efc_D_out: wp.array2d[float], + efc_vel_out: wp.array2d[float], + efc_aref_out: wp.array2d[float], + efc_frictionloss_out: wp.array2d[float], + # Out: + efc_nnz_out: wp.array[int], + ): + worldid, tenid = wp.tid() - pos = wp.max(jntrange[0], jntrange[1]) - angle - jntmargin - active = pos < 0 + tendon_frictionloss_id = worldid % tendon_frictionloss.shape[0] - if active: - wp.atomic_add(nl_out, worldid, 1) + frictionloss = tendon_frictionloss[tendon_frictionloss_id, tenid] + if frictionloss <= 0.0: + return + + wp.atomic_add(nf_out, worldid, 1) efcid = wp.atomic_add(nefc_out, worldid, 1) if efcid >= njmax_in: return - dofadr = jnt_dofadr[jntid] - dof0 = dofadr + 0 - dof1 = dofadr + 1 - dof2 = dofadr + 2 - - if is_sparse: - efc_J_rownnz_out[worldid, efcid] = 3 - rowadr = wp.atomic_add(efc_nnz_out, worldid, 3) - if rowadr + 3 > njmax_nnz_in: - return - efc_J_rowadr_out[worldid, efcid] = rowadr - - sparseid0 = rowadr + 0 - sparseid1 = rowadr + 1 - sparseid2 = rowadr + 2 - - efc_J_colind_out[worldid, 0, sparseid0] = dof0 - efc_J_colind_out[worldid, 0, sparseid1] = dof1 - efc_J_colind_out[worldid, 0, sparseid2] = dof2 - - efc_J_out[worldid, 0, sparseid0] = -axis[0] - efc_J_out[worldid, 0, sparseid1] = -axis[1] - efc_J_out[worldid, 0, sparseid2] = -axis[2] - else: - for i in range(nv): - efc_J_out[worldid, efcid, i] = 0.0 - efc_J_out[worldid, efcid, dof0] = -axis[0] - efc_J_out[worldid, efcid, dof1] = -axis[1] - efc_J_out[worldid, efcid, dof2] = -axis[2] - - Jqvel = -axis[0] * qvel_in[worldid, dof0] - Jqvel -= axis[1] * qvel_in[worldid, dof1] - Jqvel -= axis[2] * qvel_in[worldid, dof2] - - dof_invweight0_id = worldid % dof_invweight0.shape[0] - jnt_solref_id = worldid % jnt_solref.shape[0] - jnt_solimp_id = worldid % jnt_solimp.shape[0] - _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - pos, - pos, - dof_invweight0[dof_invweight0_id, dofadr], - jnt_solref[jnt_solref_id, jntid], - jnt_solimp[jnt_solimp_id, jntid], - jntmargin, - Jqvel, - 0.0, - ConstraintType.LIMIT_JOINT, - jntid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, - ) - - -@wp.kernel -def _limit_tendon( - # Model: - nv: int, - opt_timestep: wp.array[float], - opt_disableflags: int, - ten_J_rownnz: wp.array[int], - ten_J_rowadr: wp.array[int], - ten_J_colind: wp.array[int], - tendon_solref_lim: wp.array2d[wp.vec2], - tendon_solimp_lim: wp.array2d[vec5], - tendon_range: wp.array2d[wp.vec2], - tendon_margin: wp.array2d[float], - tendon_invweight0: wp.array2d[float], - is_sparse: bool, - tendon_limited_adr: wp.array[int], - # Data in: - qvel_in: wp.array2d[float], - ten_J_in: wp.array2d[float], - ten_length_in: wp.array2d[float], - njmax_in: int, - njmax_nnz_in: int, - # Data out: - nl_out: wp.array[int], - nefc_out: wp.array[int], - efc_type_out: wp.array2d[int], - efc_id_out: wp.array2d[int], - efc_J_rownnz_out: wp.array2d[int], - efc_J_rowadr_out: wp.array2d[int], - efc_J_colind_out: wp.array3d[int], - efc_J_out: wp.array3d[float], - efc_pos_out: wp.array2d[float], - efc_margin_out: wp.array2d[float], - efc_D_out: wp.array2d[float], - efc_vel_out: wp.array2d[float], - efc_aref_out: wp.array2d[float], - efc_frictionloss_out: wp.array2d[float], - # Out: - efc_nnz_out: wp.array[int], -): - worldid, tenlimitedid = wp.tid() - tenid = tendon_limited_adr[tenlimitedid] - - tendon_range_id = worldid % tendon_range.shape[0] - tenrange = tendon_range[tendon_range_id, tenid] - length = ten_length_in[worldid, tenid] - dist_min, dist_max = length - tenrange[0], tenrange[1] - length - tendon_margin_id = worldid % tendon_margin.shape[0] - tenmargin = tendon_margin[tendon_margin_id, tenid] - pos = wp.min(dist_min, dist_max) - tenmargin - active = pos < 0 - - if active: - wp.atomic_add(nl_out, worldid, 1) - efcid = wp.atomic_add(nefc_out, worldid, 1) - - if efcid >= njmax_in: - return + if wp.static(is_sparse and newton): + jgid = wp.atomic_add(efc_jtdaj_nblock_out, worldid, 1) + efc_jtdaj_adr_out[worldid, jgid] = efcid + efc_jtdaj_nrow_out[worldid, jgid] = 1 Jqvel = float(0.0) - scl = float(dist_min < dist_max) * 2.0 - 1.0 rownnz_tenJ = ten_J_rownnz[tenid] rowadr_tenJ = ten_J_rowadr[tenid] - if is_sparse: + if wp.static(is_sparse): efc_J_rownnz_out[worldid, efcid] = rownnz_tenJ rowadr_efc = wp.atomic_add(efc_nnz_out, worldid, rownnz_tenJ) if rowadr_efc + rownnz_tenJ > njmax_nnz_in: @@ -1858,7 +1582,7 @@ def _limit_tendon( sparseid_ten = rowadr_tenJ + i sparseid_efc = rowadr_efc + i colind = ten_J_colind[sparseid_ten] - J = scl * ten_J_in[worldid, sparseid_ten] + J = ten_J_in[worldid, sparseid_ten] efc_J_colind_out[worldid, 0, sparseid_efc] = colind efc_J_out[worldid, 0, sparseid_efc] = J Jqvel += J * qvel_in[worldid, colind] @@ -1867,7 +1591,7 @@ def _limit_tendon( colind = ten_J_colind[rowadr_tenJ] for i in range(nv): if nnz < rownnz_tenJ and i == colind: - J = scl * ten_J_in[worldid, rowadr_tenJ + nnz] + J = ten_J_in[worldid, rowadr_tenJ + nnz] efc_J_out[worldid, efcid, i] = J Jqvel += J * qvel_in[worldid, i] nnz += 1 @@ -1877,22 +1601,22 @@ def _limit_tendon( efc_J_out[worldid, efcid, i] = 0.0 tendon_invweight0_id = worldid % tendon_invweight0.shape[0] - tendon_solref_lim_id = worldid % tendon_solref_lim.shape[0] - tendon_solimp_lim_id = worldid % tendon_solimp_lim.shape[0] + tendon_solref_fri_id = worldid % tendon_solref_fri.shape[0] + tendon_solimp_fri_id = worldid % tendon_solimp_fri.shape[0] _efc_row( opt_disableflags, worldid, opt_timestep[worldid % opt_timestep.shape[0]], efcid, - pos, - pos, - tendon_invweight0[tendon_invweight0_id, tenid], - tendon_solref_lim[tendon_solref_lim_id, tenid], - tendon_solimp_lim[tendon_solimp_lim_id, tenid], - tenmargin, - Jqvel, 0.0, - ConstraintType.LIMIT_TENDON, + 0.0, + tendon_invweight0[tendon_invweight0_id, tenid], + tendon_solref_fri[tendon_solref_fri_id, tenid], + tendon_solimp_fri[tendon_solimp_fri_id, tenid], + 0.0, + Jqvel, + frictionloss, + ConstraintType.FRICTION_TENDON, tenid, efc_type_out, efc_id_out, @@ -1904,9 +1628,505 @@ def _limit_tendon( efc_frictionloss_out, ) + return kernel + @cache_kernel -def _efc_contact_init(cone_type: types.ConeType, is_sparse: bool): +def _limit_slide_hinge(is_sparse: bool, newton: bool): + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + nv: int, + opt_timestep: wp.array[float], + opt_disableflags: int, + jnt_qposadr: wp.array[int], + jnt_dofadr: wp.array[int], + jnt_solref: wp.array2d[wp.vec2], + jnt_solimp: wp.array2d[vec5], + jnt_range: wp.array2d[wp.vec2], + jnt_margin: wp.array2d[float], + dof_invweight0: wp.array2d[float], + jnt_limited_slide_hinge_adr: wp.array[int], + # Data in: + qpos_in: wp.array2d[float], + qvel_in: wp.array2d[float], + njmax_in: int, + njmax_nnz_in: int, + # Data out: + nl_out: wp.array[int], + nefc_out: wp.array[int], + efc_type_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_jtdaj_adr_out: wp.array2d[int], + efc_jtdaj_nrow_out: wp.array2d[int], + efc_jtdaj_nblock_out: wp.array[int], + efc_J_rownnz_out: wp.array2d[int], + efc_J_rowadr_out: wp.array2d[int], + efc_J_colind_out: wp.array3d[int], + efc_J_out: wp.array3d[float], + efc_pos_out: wp.array2d[float], + efc_margin_out: wp.array2d[float], + efc_D_out: wp.array2d[float], + efc_vel_out: wp.array2d[float], + efc_aref_out: wp.array2d[float], + efc_frictionloss_out: wp.array2d[float], + # Out: + efc_nnz_out: wp.array[int], + ): + worldid, jntlimitedid = wp.tid() + jntid = jnt_limited_slide_hinge_adr[jntlimitedid] + jnt_range_id = worldid % jnt_range.shape[0] + jntrange = jnt_range[jnt_range_id, jntid] + + qpos = qpos_in[worldid, jnt_qposadr[jntid]] + jnt_margin_id = worldid % jnt_margin.shape[0] + jntmargin = jnt_margin[jnt_margin_id, jntid] + dist_min, dist_max = qpos - jntrange[0], jntrange[1] - qpos + pos = wp.min(dist_min, dist_max) - jntmargin + active = pos < 0 + + if active: + wp.atomic_add(nl_out, worldid, 1) + efcid = wp.atomic_add(nefc_out, worldid, 1) + + if efcid >= njmax_in: + return + + if wp.static(is_sparse and newton): + jgid = wp.atomic_add(efc_jtdaj_nblock_out, worldid, 1) + efc_jtdaj_adr_out[worldid, jgid] = efcid + efc_jtdaj_nrow_out[worldid, jgid] = 1 + + dofadr = jnt_dofadr[jntid] + + J = float(dist_min < dist_max) * 2.0 - 1.0 + + if wp.static(is_sparse): + efc_J_rownnz_out[worldid, efcid] = 1 + rowadr = wp.atomic_add(efc_nnz_out, worldid, 1) + if rowadr + 1 > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid] = rowadr + efc_J_colind_out[worldid, 0, rowadr] = dofadr + efc_J_out[worldid, 0, rowadr] = J + else: + for i in range(nv): + efc_J_out[worldid, efcid, i] = 0.0 + efc_J_out[worldid, efcid, dofadr] = J + + Jqvel = J * qvel_in[worldid, dofadr] + + dof_invweight0_id = worldid % dof_invweight0.shape[0] + jnt_solref_id = worldid % jnt_solref.shape[0] + jnt_solimp_id = worldid % jnt_solimp.shape[0] + _efc_row( + opt_disableflags, + worldid, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + pos, + pos, + dof_invweight0[dof_invweight0_id, dofadr], + jnt_solref[jnt_solref_id, jntid], + jnt_solimp[jnt_solimp_id, jntid], + jntmargin, + Jqvel, + 0.0, + ConstraintType.LIMIT_JOINT, + jntid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + return kernel + + +@cache_kernel +def _limit_ball(is_sparse: bool, newton: bool): + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + nv: int, + opt_timestep: wp.array[float], + opt_disableflags: int, + jnt_qposadr: wp.array[int], + jnt_dofadr: wp.array[int], + jnt_solref: wp.array2d[wp.vec2], + jnt_solimp: wp.array2d[vec5], + jnt_range: wp.array2d[wp.vec2], + jnt_margin: wp.array2d[float], + dof_invweight0: wp.array2d[float], + jnt_limited_ball_adr: wp.array[int], + # Data in: + qpos_in: wp.array2d[float], + qvel_in: wp.array2d[float], + njmax_in: int, + njmax_nnz_in: int, + # Data out: + nl_out: wp.array[int], + nefc_out: wp.array[int], + efc_type_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_jtdaj_adr_out: wp.array2d[int], + efc_jtdaj_nrow_out: wp.array2d[int], + efc_jtdaj_nblock_out: wp.array[int], + efc_J_rownnz_out: wp.array2d[int], + efc_J_rowadr_out: wp.array2d[int], + efc_J_colind_out: wp.array3d[int], + efc_J_out: wp.array3d[float], + efc_pos_out: wp.array2d[float], + efc_margin_out: wp.array2d[float], + efc_D_out: wp.array2d[float], + efc_vel_out: wp.array2d[float], + efc_aref_out: wp.array2d[float], + efc_frictionloss_out: wp.array2d[float], + # Out: + efc_nnz_out: wp.array[int], + ): + worldid, jntlimitedid = wp.tid() + jntid = jnt_limited_ball_adr[jntlimitedid] + qposadr = jnt_qposadr[jntid] + + qpos = qpos_in[worldid] + jnt_quat = wp.quat(qpos[qposadr + 0], qpos[qposadr + 1], qpos[qposadr + 2], qpos[qposadr + 3]) + jnt_quat = wp.normalize(jnt_quat) + axis_angle = math.quat_to_vel(jnt_quat) + jnt_range_id = worldid % jnt_range.shape[0] + jntrange = jnt_range[jnt_range_id, jntid] + axis, angle = math.normalize_with_norm(axis_angle) + jnt_margin_id = worldid % jnt_margin.shape[0] + jntmargin = jnt_margin[jnt_margin_id, jntid] + + pos = wp.max(jntrange[0], jntrange[1]) - angle - jntmargin + active = pos < 0 + + if active: + wp.atomic_add(nl_out, worldid, 1) + efcid = wp.atomic_add(nefc_out, worldid, 1) + + if efcid >= njmax_in: + return + + if wp.static(is_sparse and newton): + jgid = wp.atomic_add(efc_jtdaj_nblock_out, worldid, 1) + efc_jtdaj_adr_out[worldid, jgid] = efcid + efc_jtdaj_nrow_out[worldid, jgid] = 1 + + dofadr = jnt_dofadr[jntid] + dof0 = dofadr + 0 + dof1 = dofadr + 1 + dof2 = dofadr + 2 + + if wp.static(is_sparse): + efc_J_rownnz_out[worldid, efcid] = 3 + rowadr = wp.atomic_add(efc_nnz_out, worldid, 3) + if rowadr + 3 > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid] = rowadr + + sparseid0 = rowadr + 0 + sparseid1 = rowadr + 1 + sparseid2 = rowadr + 2 + + efc_J_colind_out[worldid, 0, sparseid0] = dof0 + efc_J_colind_out[worldid, 0, sparseid1] = dof1 + efc_J_colind_out[worldid, 0, sparseid2] = dof2 + + efc_J_out[worldid, 0, sparseid0] = -axis[0] + efc_J_out[worldid, 0, sparseid1] = -axis[1] + efc_J_out[worldid, 0, sparseid2] = -axis[2] + else: + for i in range(nv): + efc_J_out[worldid, efcid, i] = 0.0 + efc_J_out[worldid, efcid, dof0] = -axis[0] + efc_J_out[worldid, efcid, dof1] = -axis[1] + efc_J_out[worldid, efcid, dof2] = -axis[2] + + Jqvel = -axis[0] * qvel_in[worldid, dof0] + Jqvel -= axis[1] * qvel_in[worldid, dof1] + Jqvel -= axis[2] * qvel_in[worldid, dof2] + + dof_invweight0_id = worldid % dof_invweight0.shape[0] + jnt_solref_id = worldid % jnt_solref.shape[0] + jnt_solimp_id = worldid % jnt_solimp.shape[0] + _efc_row( + opt_disableflags, + worldid, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + pos, + pos, + dof_invweight0[dof_invweight0_id, dofadr], + jnt_solref[jnt_solref_id, jntid], + jnt_solimp[jnt_solimp_id, jntid], + jntmargin, + Jqvel, + 0.0, + ConstraintType.LIMIT_JOINT, + jntid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + return kernel + + +@cache_kernel +def _limit_tendon(is_sparse: bool, newton: bool): + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + nv: int, + opt_timestep: wp.array[float], + opt_disableflags: int, + ten_J_rownnz: wp.array[int], + ten_J_rowadr: wp.array[int], + ten_J_colind: wp.array[int], + tendon_solref_lim: wp.array2d[wp.vec2], + tendon_solimp_lim: wp.array2d[vec5], + tendon_range: wp.array2d[wp.vec2], + tendon_margin: wp.array2d[float], + tendon_invweight0: wp.array2d[float], + tendon_limited_adr: wp.array[int], + # Data in: + qvel_in: wp.array2d[float], + ten_J_in: wp.array2d[float], + ten_length_in: wp.array2d[float], + njmax_in: int, + njmax_nnz_in: int, + # Data out: + nl_out: wp.array[int], + nefc_out: wp.array[int], + efc_type_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_jtdaj_adr_out: wp.array2d[int], + efc_jtdaj_nrow_out: wp.array2d[int], + efc_jtdaj_nblock_out: wp.array[int], + efc_J_rownnz_out: wp.array2d[int], + efc_J_rowadr_out: wp.array2d[int], + efc_J_colind_out: wp.array3d[int], + efc_J_out: wp.array3d[float], + efc_pos_out: wp.array2d[float], + efc_margin_out: wp.array2d[float], + efc_D_out: wp.array2d[float], + efc_vel_out: wp.array2d[float], + efc_aref_out: wp.array2d[float], + efc_frictionloss_out: wp.array2d[float], + # Out: + efc_nnz_out: wp.array[int], + ): + worldid, tenlimitedid = wp.tid() + tenid = tendon_limited_adr[tenlimitedid] + + tendon_range_id = worldid % tendon_range.shape[0] + tenrange = tendon_range[tendon_range_id, tenid] + length = ten_length_in[worldid, tenid] + dist_min, dist_max = length - tenrange[0], tenrange[1] - length + tendon_margin_id = worldid % tendon_margin.shape[0] + tenmargin = tendon_margin[tendon_margin_id, tenid] + pos = wp.min(dist_min, dist_max) - tenmargin + active = pos < 0 + + if active: + wp.atomic_add(nl_out, worldid, 1) + efcid = wp.atomic_add(nefc_out, worldid, 1) + + if efcid >= njmax_in: + return + + if wp.static(is_sparse and newton): + jgid = wp.atomic_add(efc_jtdaj_nblock_out, worldid, 1) + efc_jtdaj_adr_out[worldid, jgid] = efcid + efc_jtdaj_nrow_out[worldid, jgid] = 1 + + Jqvel = float(0.0) + scl = float(dist_min < dist_max) * 2.0 - 1.0 + + rownnz_tenJ = ten_J_rownnz[tenid] + rowadr_tenJ = ten_J_rowadr[tenid] + if wp.static(is_sparse): + efc_J_rownnz_out[worldid, efcid] = rownnz_tenJ + rowadr_efc = wp.atomic_add(efc_nnz_out, worldid, rownnz_tenJ) + if rowadr_efc + rownnz_tenJ > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid] = rowadr_efc + + for i in range(rownnz_tenJ): + sparseid_ten = rowadr_tenJ + i + sparseid_efc = rowadr_efc + i + colind = ten_J_colind[sparseid_ten] + J = scl * ten_J_in[worldid, sparseid_ten] + efc_J_colind_out[worldid, 0, sparseid_efc] = colind + efc_J_out[worldid, 0, sparseid_efc] = J + Jqvel += J * qvel_in[worldid, colind] + else: + nnz = int(0) + colind = ten_J_colind[rowadr_tenJ] + for i in range(nv): + if nnz < rownnz_tenJ and i == colind: + J = scl * ten_J_in[worldid, rowadr_tenJ + nnz] + efc_J_out[worldid, efcid, i] = J + Jqvel += J * qvel_in[worldid, i] + nnz += 1 + if nnz < rownnz_tenJ: + colind = ten_J_colind[rowadr_tenJ + nnz] + else: + efc_J_out[worldid, efcid, i] = 0.0 + + tendon_invweight0_id = worldid % tendon_invweight0.shape[0] + tendon_solref_lim_id = worldid % tendon_solref_lim.shape[0] + tendon_solimp_lim_id = worldid % tendon_solimp_lim.shape[0] + _efc_row( + opt_disableflags, + worldid, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + pos, + pos, + tendon_invweight0[tendon_invweight0_id, tenid], + tendon_solref_lim[tendon_solref_lim_id, tenid], + tendon_solimp_lim[tendon_solimp_lim_id, tenid], + tenmargin, + Jqvel, + 0.0, + ConstraintType.LIMIT_TENDON, + tenid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + return kernel + + +@wp.func +def _get_contact_bodies_and_weights( + # Model: + geom_bodyid: wp.array[int], + flex_dim: wp.array[int], + flex_vertadr: wp.array[int], + flex_elemdataadr: wp.array[int], + flex_shelldataadr: wp.array[int], + flex_vertbodyid: wp.array[int], + flex_elem: wp.array[int], + flex_shell: wp.array[int], + # Data in: + flexvert_xpos_in: wp.array2d[wp.vec3], + # In: + conid: int, + side: int, + geom: wp.vec2i, + flex: wp.vec2i, + elem: wp.vec2i, + vert: wp.vec2i, + con_pos: wp.vec3, + worldid: int, +) -> Tuple[wp.vec4i, wp.vec4]: + geom_id = geom[side] + flex_id = flex[side] + elem_id = elem[side] + vert_id = vert[side] + + # Rigid Geom Side + if geom_id >= 0: + return wp.vec4i(geom_bodyid[geom_id], -1, -1, -1), wp.vec4(1.0, 0.0, 0.0, 0.0) + + # Plane-Vertex or Vertex-only flex contact + flex_vert_start = flex_vertadr[flex_id] + if vert_id >= 0: + body = flex_vertbodyid[flex_vert_start + vert_id] + return wp.vec4i(body, -1, -1, -1), wp.vec4(1.0, 0.0, 0.0, 0.0) + + # Element contact: Retrieve local vertices + dim = flex_dim[flex_id] + + if dim == 2: + elem_data_start = flex_elemdataadr[flex_id] + elem_id * 3 + v0 = flex_elem[elem_data_start + 0] + v1 = flex_elem[elem_data_start + 1] + v2 = flex_elem[elem_data_start + 2] + + x0 = flexvert_xpos_in[worldid, flex_vert_start + v0] + x1 = flexvert_xpos_in[worldid, flex_vert_start + v1] + x2 = flexvert_xpos_in[worldid, flex_vert_start + v2] + + d0 = wp.length(con_pos - x0) + d1 = wp.length(con_pos - x1) + d2 = wp.length(con_pos - x2) + + w0 = 1.0 / wp.max(types.MJ_MINVAL, d0) + w1 = 1.0 / wp.max(types.MJ_MINVAL, d1) + w2 = 1.0 / wp.max(types.MJ_MINVAL, d2) + + w_sum = w0 + w1 + w2 + w0 = w0 / w_sum + w1 = w1 / w_sum + w2 = w2 / w_sum + + b0 = flex_vertbodyid[flex_vert_start + v0] + b1 = flex_vertbodyid[flex_vert_start + v1] + b2 = flex_vertbodyid[flex_vert_start + v2] + + return wp.vec4i(b0, b1, b2, -1), wp.vec4(w0, w1, w2, 0.0) + + elif dim == 3: + elem_data_start = flex_elemdataadr[flex_id] + elem_id * 4 + v0 = flex_elem[elem_data_start + 0] + v1 = flex_elem[elem_data_start + 1] + v2 = flex_elem[elem_data_start + 2] + v3 = flex_elem[elem_data_start + 3] + + x0 = flexvert_xpos_in[worldid, flex_vert_start + v0] + x1 = flexvert_xpos_in[worldid, flex_vert_start + v1] + x2 = flexvert_xpos_in[worldid, flex_vert_start + v2] + x3 = flexvert_xpos_in[worldid, flex_vert_start + v3] + + d0 = wp.length(con_pos - x0) + d1 = wp.length(con_pos - x1) + d2 = wp.length(con_pos - x2) + d3 = wp.length(con_pos - x3) + + w0 = 1.0 / wp.max(types.MJ_MINVAL, d0) + w1 = 1.0 / wp.max(types.MJ_MINVAL, d1) + w2 = 1.0 / wp.max(types.MJ_MINVAL, d2) + w3 = 1.0 / wp.max(types.MJ_MINVAL, d3) + + w_sum = w0 + w1 + w2 + w3 + w0 = w0 / w_sum + w1 = w1 / w_sum + w2 = w2 / w_sum + w3 = w3 / w_sum + + b0 = flex_vertbodyid[flex_vert_start + v0] + b1 = flex_vertbodyid[flex_vert_start + v1] + b2 = flex_vertbodyid[flex_vert_start + v2] + b3 = flex_vertbodyid[flex_vert_start + v3] + + return wp.vec4i(b0, b1, b2, b3), wp.vec4(w0, w1, w2, w3) + + else: + return wp.vec4i(-1, -1, -1, -1), wp.vec4(0.0, 0.0, 0.0, 0.0) + + +@cache_kernel +def _efc_contact_init(cone_type: types.ConeType, is_sparse: bool, newton: bool): IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC IS_SPARSE = is_sparse @@ -1918,8 +2138,6 @@ def _efc_contact_init(cone_type: types.ConeType, is_sparse: bool): body_dofadr: wp.array[int], dof_parentid: wp.array[int], geom_bodyid: wp.array[int], - flex_vertadr: wp.array[int], - flex_vertbodyid: wp.array[int], # Data in: njmax_in: int, njmax_nnz_in: int, @@ -1930,13 +2148,14 @@ def _efc_contact_init(cone_type: types.ConeType, is_sparse: bool): includemargin_in: wp.array[float], worldid_in: wp.array[int], geom_in: wp.array[wp.vec2i], - flex_in: wp.array[wp.vec2i], - vert_in: wp.array[wp.vec2i], type_in: wp.array[int], # Data out: nefc_out: wp.array[int], contact_efc_address_out: wp.array2d[int], efc_id_out: wp.array2d[int], + efc_jtdaj_adr_out: wp.array2d[int], + efc_jtdaj_nrow_out: wp.array2d[int], + efc_jtdaj_nblock_out: wp.array[int], efc_J_rownnz_out: wp.array2d[int], efc_J_rowadr_out: wp.array2d[int], # Out: @@ -1980,26 +2199,16 @@ def _efc_contact_init(cone_type: types.ConeType, is_sparse: bool): # This is redundant with the _efc_row call later but needed for the jac calculation efc_id_out[worldid, efcid] = conid + if wp.static(is_sparse and newton): + if base_efcid < njmax_in: + jgid = wp.atomic_add(efc_jtdaj_nblock_out, worldid, 1) + efc_jtdaj_adr_out[worldid, jgid] = base_efcid + efc_jtdaj_nrow_out[worldid, jgid] = wp.min(ndim, njmax_in - base_efcid) + if wp.static(IS_SPARSE): geom = geom_in[conid] - - if geom[0] >= 0: - body1 = geom_bodyid[geom[0]] - else: - flex = flex_in[conid] - vert = vert_in[conid] - body1 = flex_vertbodyid[flex_vertadr[flex[0]] + vert[0]] - - if geom[1] >= 0: - body2 = geom_bodyid[geom[1]] - else: - flex = flex_in[conid] - vert = vert_in[conid] - body2 = flex_vertbodyid[flex_vertadr[flex[1]] + vert[1]] - - # skip fixed bodies - body1 = body_weldid[body1] - body2 = body_weldid[body2] + body1 = body_weldid[geom_bodyid[geom[0]]] + body2 = body_weldid[geom_bodyid[geom[1]]] da1 = int(body_dofadr[body1] + body_dofnum[body1] - 1) da2 = int(body_dofadr[body2] + body_dofnum[body2] - 1) @@ -2029,6 +2238,206 @@ def _efc_contact_init(cone_type: types.ConeType, is_sparse: bool): return kernel +@cache_kernel +def _efc_contact_init_flex(cone_type: types.ConeType, is_sparse: bool, newton: bool): + IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC + IS_SPARSE = is_sparse + + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + body_parentid: wp.array[int], + body_weldid: wp.array[int], + body_dofnum: wp.array[int], + body_dofadr: wp.array[int], + dof_parentid: wp.array[int], + geom_bodyid: wp.array[int], + flex_dim: wp.array[int], + flex_vertadr: wp.array[int], + flex_elemdataadr: wp.array[int], + flex_shelldataadr: wp.array[int], + flex_vertbodyid: wp.array[int], + flex_elem: wp.array[int], + flex_shell: wp.array[int], + # Data in: + flexvert_xpos_in: wp.array2d[wp.vec3], + njmax_in: int, + njmax_nnz_in: int, + nacon_in: wp.array[int], + # In: + dist_in: wp.array[float], + pos_in: wp.array[wp.vec3], + condim_in: wp.array[int], + includemargin_in: wp.array[float], + worldid_in: wp.array[int], + geom_in: wp.array[wp.vec2i], + flex_in: wp.array[wp.vec2i], + elem_in: wp.array[wp.vec2i], + vert_in: wp.array[wp.vec2i], + type_in: wp.array[int], + # Data out: + nefc_out: wp.array[int], + contact_efc_address_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_jtdaj_adr_out: wp.array2d[int], + efc_jtdaj_nrow_out: wp.array2d[int], + efc_jtdaj_nblock_out: wp.array[int], + efc_J_rownnz_out: wp.array2d[int], + efc_J_rowadr_out: wp.array2d[int], + # Out: + efc_nnz_out: wp.array[int], + ): + conid = wp.tid() + + if conid >= nacon_in[0]: + return + + if not type_in[conid] & ContactType.CONSTRAINT: + return + + condim = condim_in[conid] + + includemargin = includemargin_in[conid] + pos = dist_in[conid] - includemargin + active = pos < 0 + + if not active: + return + + if wp.static(IS_ELLIPTIC): + ndim = condim + else: + if condim == 1: + ndim = 1 + else: + ndim = 2 * (condim - 1) + + worldid = worldid_in[conid] + + # Allocate contiguous block of efcids for all dimids + base_efcid = wp.atomic_add(nefc_out, worldid, ndim) + for dim in range(ndim): + efcid = base_efcid + dim + if efcid >= njmax_in: + contact_efc_address_out[conid, dim] = -1 + else: + contact_efc_address_out[conid, dim] = efcid + # This is redundant with the _efc_row call later but needed for the jac calculation + efc_id_out[worldid, efcid] = conid + + if wp.static(is_sparse and newton): + if base_efcid < njmax_in: + jgid = wp.atomic_add(efc_jtdaj_nblock_out, worldid, 1) + efc_jtdaj_adr_out[worldid, jgid] = base_efcid + efc_jtdaj_nrow_out[worldid, jgid] = wp.min(ndim, njmax_in - base_efcid) + + if wp.static(IS_SPARSE): + geom = geom_in[conid] + flex = flex_in[conid] + elem = elem_in[conid] + vert = vert_in[conid] + con_pos = pos_in[conid] + + body_ids1, weights1 = _get_contact_bodies_and_weights( + geom_bodyid, + flex_dim, + flex_vertadr, + flex_elemdataadr, + flex_shelldataadr, + flex_vertbodyid, + flex_elem, + flex_shell, + flexvert_xpos_in, + conid, + 0, + geom, + flex, + elem, + vert, + con_pos, + worldid, + ) + body_ids2, weights2 = _get_contact_bodies_and_weights( + geom_bodyid, + flex_dim, + flex_vertadr, + flex_elemdataadr, + flex_shelldataadr, + flex_vertbodyid, + flex_elem, + flex_shell, + flexvert_xpos_in, + conid, + 1, + geom, + flex, + elem, + vert, + con_pos, + worldid, + ) + + b1_0 = body_weldid[body_ids1[0]] + b1_1 = body_weldid[body_ids1[1]] if body_ids1[1] >= 0 else -1 + b1_2 = body_weldid[body_ids1[2]] if body_ids1[2] >= 0 else -1 + b1_3 = body_weldid[body_ids1[3]] if body_ids1[3] >= 0 else -1 + + b2_0 = body_weldid[body_ids2[0]] + b2_1 = body_weldid[body_ids2[1]] if body_ids2[1] >= 0 else -1 + b2_2 = body_weldid[body_ids2[2]] if body_ids2[2] >= 0 else -1 + b2_3 = body_weldid[body_ids2[3]] if body_ids2[3] >= 0 else -1 + + dof1_0 = int(body_dofadr[b1_0] + body_dofnum[b1_0] - 1) if b1_0 >= 0 else -1 + dof1_1 = int(body_dofadr[b1_1] + body_dofnum[b1_1] - 1) if b1_1 >= 0 else -1 + dof1_2 = int(body_dofadr[b1_2] + body_dofnum[b1_2] - 1) if b1_2 >= 0 else -1 + dof1_3 = int(body_dofadr[b1_3] + body_dofnum[b1_3] - 1) if b1_3 >= 0 else -1 + + dof2_0 = int(body_dofadr[b2_0] + body_dofnum[b2_0] - 1) if b2_0 >= 0 else -1 + dof2_1 = int(body_dofadr[b2_1] + body_dofnum[b2_1] - 1) if b2_1 >= 0 else -1 + dof2_2 = int(body_dofadr[b2_2] + body_dofnum[b2_2] - 1) if b2_2 >= 0 else -1 + dof2_3 = int(body_dofadr[b2_3] + body_dofnum[b2_3] - 1) if b2_3 >= 0 else -1 + + # count non-zeros + rownnz = int(0) + while ( + dof1_0 >= 0 or dof1_1 >= 0 or dof1_2 >= 0 or dof1_3 >= 0 or dof2_0 >= 0 or dof2_1 >= 0 or dof2_2 >= 0 or dof2_3 >= 0 + ): + da1_max = wp.max(dof1_0, wp.max(dof1_1, wp.max(dof1_2, dof1_3))) + da2_max = wp.max(dof2_0, wp.max(dof2_1, wp.max(dof2_2, dof2_3))) + da = wp.max(da1_max, da2_max) + + if dof1_0 == da: + dof1_0 = dof_parentid[dof1_0] + if dof1_1 == da: + dof1_1 = dof_parentid[dof1_1] + if dof1_2 == da: + dof1_2 = dof_parentid[dof1_2] + if dof1_3 == da: + dof1_3 = dof_parentid[dof1_3] + + if dof2_0 == da: + dof2_0 = dof_parentid[dof2_0] + if dof2_1 == da: + dof2_1 = dof_parentid[dof2_1] + if dof2_2 == da: + dof2_2 = dof_parentid[dof2_2] + if dof2_3 == da: + dof2_3 = dof_parentid[dof2_3] + + rownnz += 1 + + rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz * ndim) + if rowadr + rownnz * ndim > njmax_nnz_in: + return + for dim in range(ndim): + efcid = base_efcid + dim + if efcid < njmax_in: + efc_J_rowadr_out[worldid, efcid] = rowadr + dim * rownnz + efc_J_rownnz_out[worldid, efcid] = rownnz + + return kernel + + @cache_kernel def _efc_contact_jac_sparse(cone_type: types.ConeType): IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC @@ -2044,8 +2453,6 @@ def _efc_contact_jac_sparse(cone_type: types.ConeType): dof_bodyid: wp.array[int], dof_parentid: wp.array[int], geom_bodyid: wp.array[int], - flex_vertadr: wp.array[int], - flex_vertbodyid: wp.array[int], body_isdofancestor: wp.array2d[int], # Data in: qvel_in: wp.array2d[float], @@ -2058,8 +2465,6 @@ def _efc_contact_jac_sparse(cone_type: types.ConeType): # In: condim_in: wp.array[int], geom_in: wp.array[wp.vec2i], - flex_in: wp.array[wp.vec2i], - vert_in: wp.array[wp.vec2i], pos_in: wp.array[wp.vec3], frame_in: wp.array2d[wp.vec3], friction_in: wp.array2d[float], @@ -2082,23 +2487,8 @@ def _efc_contact_jac_sparse(cone_type: types.ConeType): condim = condim_in[conid] geom = geom_in[conid] - if geom[0] >= 0: - body1 = geom_bodyid[geom[0]] - else: - flex = flex_in[conid] - vert = vert_in[conid] - body1 = flex_vertbodyid[flex_vertadr[flex[0]] + vert[0]] - - if geom[1] >= 0: - body2 = geom_bodyid[geom[1]] - else: - flex = flex_in[conid] - vert = vert_in[conid] - body2 = flex_vertbodyid[flex_vertadr[flex[1]] + vert[1]] - - # skip fixed bodies - body1 = body_weldid[body1] - body2 = body_weldid[body2] + body1 = body_weldid[geom_bodyid[geom[0]]] + body2 = body_weldid[geom_bodyid[geom[1]]] con_pos = pos_in[conid] @@ -2200,6 +2590,274 @@ def _efc_contact_jac_sparse(cone_type: types.ConeType): return kernel +@cache_kernel +def _efc_contact_jac_sparse_flex(cone_type: types.ConeType): + IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC + + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + body_parentid: wp.array[int], + body_rootid: wp.array[int], + body_weldid: wp.array[int], + body_dofnum: wp.array[int], + body_dofadr: wp.array[int], + dof_bodyid: wp.array[int], + dof_parentid: wp.array[int], + geom_bodyid: wp.array[int], + flex_dim: wp.array[int], + flex_vertadr: wp.array[int], + flex_elemdataadr: wp.array[int], + flex_shelldataadr: wp.array[int], + flex_vertbodyid: wp.array[int], + flex_elem: wp.array[int], + flex_shell: wp.array[int], + body_isdofancestor: wp.array2d[int], + # Data in: + qvel_in: wp.array2d[float], + subtree_com_in: wp.array2d[wp.vec3], + cdof_in: wp.array2d[wp.spatial_vector], + flexvert_xpos_in: wp.array2d[wp.vec3], + contact_efc_address_in: wp.array2d[int], + efc_J_rownnz_in: wp.array2d[int], + efc_J_rowadr_in: wp.array2d[int], + nacon_in: wp.array[int], + # In: + condim_in: wp.array[int], + geom_in: wp.array[wp.vec2i], + flex_in: wp.array[wp.vec2i], + elem_in: wp.array[wp.vec2i], + vert_in: wp.array[wp.vec2i], + pos_in: wp.array[wp.vec3], + frame_in: wp.array2d[wp.vec3], + friction_in: wp.array2d[float], + worldid_in: wp.array[int], + # Data out: + efc_J_colind_out: wp.array3d[int], + efc_J_out: wp.array3d[float], + efc_Jqvel_out: wp.array2d[float], + ): + conid, dimid = wp.tid() + + if conid >= nacon_in[0]: + return + + efcid = contact_efc_address_in[conid, dimid] + if efcid < 0: + return + + worldid = worldid_in[conid] + condim = condim_in[conid] + + geom = geom_in[conid] + flex = flex_in[conid] + elem = elem_in[conid] + vert = vert_in[conid] + con_pos = pos_in[conid] + + body_ids1, weights1 = _get_contact_bodies_and_weights( + geom_bodyid, + flex_dim, + flex_vertadr, + flex_elemdataadr, + flex_shelldataadr, + flex_vertbodyid, + flex_elem, + flex_shell, + flexvert_xpos_in, + conid, + 0, + geom, + flex, + elem, + vert, + con_pos, + worldid, + ) + body_ids2, weights2 = _get_contact_bodies_and_weights( + geom_bodyid, + flex_dim, + flex_vertadr, + flex_elemdataadr, + flex_shelldataadr, + flex_vertbodyid, + flex_elem, + flex_shell, + flexvert_xpos_in, + conid, + 1, + geom, + flex, + elem, + vert, + con_pos, + worldid, + ) + + # skip fixed bodies + b1_0 = body_weldid[body_ids1[0]] + b1_1 = body_weldid[body_ids1[1]] if body_ids1[1] >= 0 else -1 + b1_2 = body_weldid[body_ids1[2]] if body_ids1[2] >= 0 else -1 + b1_3 = body_weldid[body_ids1[3]] if body_ids1[3] >= 0 else -1 + + b2_0 = body_weldid[body_ids2[0]] + b2_1 = body_weldid[body_ids2[1]] if body_ids2[1] >= 0 else -1 + b2_2 = body_weldid[body_ids2[2]] if body_ids2[2] >= 0 else -1 + b2_3 = body_weldid[body_ids2[3]] if body_ids2[3] >= 0 else -1 + + if not wp.static(IS_ELLIPTIC): + frame_0 = frame_in[conid, 0] + if condim > 1: + dimid2 = dimid / 2 + 1 + frii = friction_in[conid, dimid2 - 1] + + dof1_0 = int(body_dofadr[b1_0] + body_dofnum[b1_0] - 1) if b1_0 >= 0 else -1 + dof1_1 = int(body_dofadr[b1_1] + body_dofnum[b1_1] - 1) if b1_1 >= 0 else -1 + dof1_2 = int(body_dofadr[b1_2] + body_dofnum[b1_2] - 1) if b1_2 >= 0 else -1 + dof1_3 = int(body_dofadr[b1_3] + body_dofnum[b1_3] - 1) if b1_3 >= 0 else -1 + da1 = wp.max(dof1_0, wp.max(dof1_1, wp.max(dof1_2, dof1_3))) + + dof2_0 = int(body_dofadr[b2_0] + body_dofnum[b2_0] - 1) if b2_0 >= 0 else -1 + dof2_1 = int(body_dofadr[b2_1] + body_dofnum[b2_1] - 1) if b2_1 >= 0 else -1 + dof2_2 = int(body_dofadr[b2_2] + body_dofnum[b2_2] - 1) if b2_2 >= 0 else -1 + dof2_3 = int(body_dofadr[b2_3] + body_dofnum[b2_3] - 1) if b2_3 >= 0 else -1 + da2 = wp.max(dof2_0, wp.max(dof2_1, wp.max(dof2_2, dof2_3))) + + da = wp.max(da1, da2) + + rowadr = efc_J_rowadr_in[worldid, efcid] + rownnz = efc_J_rownnz_in[worldid, efcid] + + Jqvel = float(0.0) + nnz = int(0) + dofid = int(da) + + while True: + if nnz >= rownnz: + break + + if dofid == da: + jac1p = wp.vec3(0.0) + jac1r = wp.vec3(0.0) + if dof1_0 == da: + jp, jr = support.jac_dof( + body_parentid, body_rootid, dof_bodyid, body_isdofancestor, subtree_com_in, cdof_in, con_pos, b1_0, dofid, worldid + ) + jac1p += jp * weights1[0] + jac1r += jr * weights1[0] + if dof1_1 == da: + jp, jr = support.jac_dof( + body_parentid, body_rootid, dof_bodyid, body_isdofancestor, subtree_com_in, cdof_in, con_pos, b1_1, dofid, worldid + ) + jac1p += jp * weights1[1] + jac1r += jr * weights1[1] + if dof1_2 == da: + jp, jr = support.jac_dof( + body_parentid, body_rootid, dof_bodyid, body_isdofancestor, subtree_com_in, cdof_in, con_pos, b1_2, dofid, worldid + ) + jac1p += jp * weights1[2] + jac1r += jr * weights1[2] + if dof1_3 == da: + jp, jr = support.jac_dof( + body_parentid, body_rootid, dof_bodyid, body_isdofancestor, subtree_com_in, cdof_in, con_pos, b1_3, dofid, worldid + ) + jac1p += jp * weights1[3] + jac1r += jr * weights1[3] + + jac2p = wp.vec3(0.0) + jac2r = wp.vec3(0.0) + if dof2_0 == da: + jp, jr = support.jac_dof( + body_parentid, body_rootid, dof_bodyid, body_isdofancestor, subtree_com_in, cdof_in, con_pos, b2_0, dofid, worldid + ) + jac2p += jp * weights2[0] + jac2r += jr * weights2[0] + if dof2_1 == da: + jp, jr = support.jac_dof( + body_parentid, body_rootid, dof_bodyid, body_isdofancestor, subtree_com_in, cdof_in, con_pos, b2_1, dofid, worldid + ) + jac2p += jp * weights2[1] + jac2r += jr * weights2[1] + if dof2_2 == da: + jp, jr = support.jac_dof( + body_parentid, body_rootid, dof_bodyid, body_isdofancestor, subtree_com_in, cdof_in, con_pos, b2_2, dofid, worldid + ) + jac2p += jp * weights2[2] + jac2r += jr * weights2[2] + if dof2_3 == da: + jp, jr = support.jac_dof( + body_parentid, body_rootid, dof_bodyid, body_isdofancestor, subtree_com_in, cdof_in, con_pos, b2_3, dofid, worldid + ) + jac2p += jp * weights2[3] + jac2r += jr * weights2[3] + + jacp_dif = jac2p - jac1p + jacr_dif = jac2r - jac1r + + if wp.static(IS_ELLIPTIC): + J = float(0.0) + if dimid < 3: + frame_row = frame_in[conid, dimid] + for xyz in range(3): + J += frame_row[xyz] * jacp_dif[xyz] + else: + frame_row = frame_in[conid, dimid - 3] + for xyz in range(3): + J += frame_row[xyz] * jacr_dif[xyz] + else: + J = float(0.0) + Ji = float(0.0) + + for xyz in range(3): + J += frame_0[xyz] * jacp_dif[xyz] + + if condim > 1: + if dimid2 < 3: + Ji += frame_in[conid, dimid2][xyz] * jacp_dif[xyz] + else: + Ji += frame_in[conid, dimid2 - 3][xyz] * jacr_dif[xyz] + + if condim > 1: + if dimid % 2 == 0: + J += Ji * frii + else: + J -= Ji * frii + + sparseid = rowadr + nnz + efc_J_colind_out[worldid, 0, sparseid] = dofid + efc_J_out[worldid, 0, sparseid] = J + nnz += 1 + Jqvel += J * qvel_in[worldid, dofid] + + # Advance tree pointers + if dof1_0 == da: + dof1_0 = dof_parentid[dof1_0] + if dof1_1 == da: + dof1_1 = dof_parentid[dof1_1] + if dof1_2 == da: + dof1_2 = dof_parentid[dof1_2] + if dof1_3 == da: + dof1_3 = dof_parentid[dof1_3] + da1 = wp.max(dof1_0, wp.max(dof1_1, wp.max(dof1_2, dof1_3))) + + if dof2_0 == da: + dof2_0 = dof_parentid[dof2_0] + if dof2_1 == da: + dof2_1 = dof_parentid[dof2_1] + if dof2_2 == da: + dof2_2 = dof_parentid[dof2_2] + if dof2_3 == da: + dof2_3 = dof_parentid[dof2_3] + da2 = wp.max(dof2_0, wp.max(dof2_1, wp.max(dof2_2, dof2_3))) + + da = wp.max(da1, da2) + dofid = da + + efc_Jqvel_out[worldid, efcid] = Jqvel + + return kernel + + @cache_kernel def _efc_contact_jac_dense(tile_size: int, cone_type: types.ConeType): TILE_SIZE = tile_size @@ -2210,8 +2868,6 @@ def _efc_contact_jac_dense(tile_size: int, cone_type: types.ConeType): # Model: body_rootid: wp.array[int], geom_bodyid: wp.array[int], - flex_vertadr: wp.array[int], - flex_vertbodyid: wp.array[int], body_isdofancestor: wp.array2d[int], # Data in: ne_in: wp.array[int], @@ -2228,8 +2884,6 @@ def _efc_contact_jac_dense(tile_size: int, cone_type: types.ConeType): nv_padded: int, condim_in: wp.array[int], geom_in: wp.array[wp.vec2i], - flex_in: wp.array[wp.vec2i], - vert_in: wp.array[wp.vec2i], pos_in: wp.array[wp.vec3], frame_in: wp.array2d[wp.vec3], friction_in: wp.array2d[float], @@ -2261,19 +2915,8 @@ def _efc_contact_jac_dense(tile_size: int, cone_type: types.ConeType): condim = condim_in[conid] geom = geom_in[conid] - if geom[0] >= 0: - body1 = geom_bodyid[geom[0]] - else: - flex = flex_in[conid] - vert = vert_in[conid] - body1 = flex_vertbodyid[flex_vertadr[flex[0]] + vert[0]] - - if geom[1] >= 0: - body2 = geom_bodyid[geom[1]] - else: - flex = flex_in[conid] - vert = vert_in[conid] - body2 = flex_vertbodyid[flex_vertadr[flex[1]] + vert[1]] + body1 = geom_bodyid[geom[0]] + body2 = geom_bodyid[geom[1]] con_pos = pos_in[conid] offset1 = con_pos - subtree_com_in[worldid, body_rootid[body1]] @@ -2345,6 +2988,311 @@ def _efc_contact_jac_dense(tile_size: int, cone_type: types.ConeType): return kernel +@cache_kernel +def _efc_contact_jac_dense_flex(tile_size: int, cone_type: types.ConeType): + TILE_SIZE = tile_size + IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC + + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + body_rootid: wp.array[int], + geom_bodyid: wp.array[int], + flex_dim: wp.array[int], + flex_vertadr: wp.array[int], + flex_elemdataadr: wp.array[int], + flex_shelldataadr: wp.array[int], + flex_vertbodyid: wp.array[int], + flex_elem: wp.array[int], + flex_shell: wp.array[int], + body_isdofancestor: wp.array2d[int], + # Data in: + ne_in: wp.array[int], + nf_in: wp.array[int], + nl_in: wp.array[int], + nefc_in: wp.array[int], + qvel_in: wp.array2d[float], + subtree_com_in: wp.array2d[wp.vec3], + cdof_in: wp.array2d[wp.spatial_vector], + flexvert_xpos_in: wp.array2d[wp.vec3], + contact_efc_address_in: wp.array2d[int], + efc_id_in: wp.array2d[int], + njmax_in: int, + # In: + nv_padded: int, + condim_in: wp.array[int], + geom_in: wp.array[wp.vec2i], + flex_in: wp.array[wp.vec2i], + elem_in: wp.array[wp.vec2i], + vert_in: wp.array[wp.vec2i], + pos_in: wp.array[wp.vec3], + frame_in: wp.array2d[wp.vec3], + friction_in: wp.array2d[float], + # Data out: + efc_J_out: wp.array3d[float], + efc_Jqvel_out: wp.array2d[float], + ): + worldid, dof_block_id, tid = wp.tid() + + dof_start = dof_block_id * wp.static(TILE_SIZE) + if dof_start >= nv_padded: + return + + cdof_tile = wp.tile_load(cdof_in[worldid], shape=TILE_SIZE, offset=dof_start, bounds_check=True) + qvel_tile = wp.tile_load(qvel_in[worldid], shape=TILE_SIZE, offset=dof_start, bounds_check=True) + + efcid_start = ne_in[worldid] + nf_in[worldid] + nl_in[worldid] + efcid_end = wp.min(nefc_in[worldid], njmax_in) + + prev_conid = int(-1) + condim = int(0) + + for efcid in range(efcid_start, efcid_end): + conid = efc_id_in[worldid, efcid] + + # Recompute per-contact data only when contact changes + if conid != prev_conid: + prev_conid = conid + condim = condim_in[conid] + + geom = geom_in[conid] + flex = flex_in[conid] + elem = elem_in[conid] + vert = vert_in[conid] + con_pos = pos_in[conid] + + body_ids1, weights1 = _get_contact_bodies_and_weights( + geom_bodyid, + flex_dim, + flex_vertadr, + flex_elemdataadr, + flex_shelldataadr, + flex_vertbodyid, + flex_elem, + flex_shell, + flexvert_xpos_in, + conid, + 0, + geom, + flex, + elem, + vert, + con_pos, + worldid, + ) + body_ids2, weights2 = _get_contact_bodies_and_weights( + geom_bodyid, + flex_dim, + flex_vertadr, + flex_elemdataadr, + flex_shelldataadr, + flex_vertbodyid, + flex_elem, + flex_shell, + flexvert_xpos_in, + conid, + 1, + geom, + flex, + elem, + vert, + con_pos, + worldid, + ) + + b1_0 = body_ids1[0] + b1_1 = body_ids1[1] + b1_2 = body_ids1[2] + b1_3 = body_ids1[3] + + b2_0 = body_ids2[0] + b2_1 = body_ids2[1] + b2_2 = body_ids2[2] + b2_3 = body_ids2[3] + + # Weighted jacp for side 1 + jacp1_tile = wp.tile_map( + support._compute_jacp, + cdof_tile, + con_pos - subtree_com_in[worldid, body_rootid[b1_0]], + wp.tile_load(body_isdofancestor[b1_0], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacp1_tile = wp.tile_map(wp.mul, jacp1_tile, weights1[0]) + if b1_1 >= 0: + t_jacp = wp.tile_map( + support._compute_jacp, + cdof_tile, + con_pos - subtree_com_in[worldid, body_rootid[b1_1]], + wp.tile_load(body_isdofancestor[b1_1], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacp1_tile = wp.tile_map(wp.add, jacp1_tile, wp.tile_map(wp.mul, t_jacp, weights1[1])) + if b1_2 >= 0: + t_jacp = wp.tile_map( + support._compute_jacp, + cdof_tile, + con_pos - subtree_com_in[worldid, body_rootid[b1_2]], + wp.tile_load(body_isdofancestor[b1_2], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacp1_tile = wp.tile_map(wp.add, jacp1_tile, wp.tile_map(wp.mul, t_jacp, weights1[2])) + if b1_3 >= 0: + t_jacp = wp.tile_map( + support._compute_jacp, + cdof_tile, + con_pos - subtree_com_in[worldid, body_rootid[b1_3]], + wp.tile_load(body_isdofancestor[b1_3], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacp1_tile = wp.tile_map(wp.add, jacp1_tile, wp.tile_map(wp.mul, t_jacp, weights1[3])) + + # Weighted jacp for side 2 + jacp2_tile = wp.tile_map( + support._compute_jacp, + cdof_tile, + con_pos - subtree_com_in[worldid, body_rootid[b2_0]], + wp.tile_load(body_isdofancestor[b2_0], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacp2_tile = wp.tile_map(wp.mul, jacp2_tile, weights2[0]) + if b2_1 >= 0: + t_jacp = wp.tile_map( + support._compute_jacp, + cdof_tile, + con_pos - subtree_com_in[worldid, body_rootid[b2_1]], + wp.tile_load(body_isdofancestor[b2_1], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacp2_tile = wp.tile_map(wp.add, jacp2_tile, wp.tile_map(wp.mul, t_jacp, weights2[1])) + if b2_2 >= 0: + t_jacp = wp.tile_map( + support._compute_jacp, + cdof_tile, + con_pos - subtree_com_in[worldid, body_rootid[b2_2]], + wp.tile_load(body_isdofancestor[b2_2], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacp2_tile = wp.tile_map(wp.add, jacp2_tile, wp.tile_map(wp.mul, t_jacp, weights2[2])) + if b2_3 >= 0: + t_jacp = wp.tile_map( + support._compute_jacp, + cdof_tile, + con_pos - subtree_com_in[worldid, body_rootid[b2_3]], + wp.tile_load(body_isdofancestor[b2_3], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacp2_tile = wp.tile_map(wp.add, jacp2_tile, wp.tile_map(wp.mul, t_jacp, weights2[3])) + + jacp_dif_tile = wp.tile_map(wp.sub, jacp2_tile, jacp1_tile) + + # Weighted jacr for side 1 + jacr1_tile = wp.tile_map( + support._compute_jacr, + cdof_tile, + wp.tile_load(body_isdofancestor[b1_0], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacr1_tile = wp.tile_map(wp.mul, jacr1_tile, weights1[0]) + if b1_1 >= 0: + t_jacr = wp.tile_map( + support._compute_jacr, + cdof_tile, + wp.tile_load(body_isdofancestor[b1_1], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacr1_tile = wp.tile_map(wp.add, jacr1_tile, wp.tile_map(wp.mul, t_jacr, weights1[1])) + if b1_2 >= 0: + t_jacr = wp.tile_map( + support._compute_jacr, + cdof_tile, + wp.tile_load(body_isdofancestor[b1_2], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacr1_tile = wp.tile_map(wp.add, jacr1_tile, wp.tile_map(wp.mul, t_jacr, weights1[2])) + if b1_3 >= 0: + t_jacr = wp.tile_map( + support._compute_jacr, + cdof_tile, + wp.tile_load(body_isdofancestor[b1_3], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacr1_tile = wp.tile_map(wp.add, jacr1_tile, wp.tile_map(wp.mul, t_jacr, weights1[3])) + + # Weighted jacr for side 2 + jacr2_tile = wp.tile_map( + support._compute_jacr, + cdof_tile, + wp.tile_load(body_isdofancestor[b2_0], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacr2_tile = wp.tile_map(wp.mul, jacr2_tile, weights2[0]) + if b2_1 >= 0: + t_jacr = wp.tile_map( + support._compute_jacr, + cdof_tile, + wp.tile_load(body_isdofancestor[b2_1], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacr2_tile = wp.tile_map(wp.add, jacr2_tile, wp.tile_map(wp.mul, t_jacr, weights2[1])) + if b2_2 >= 0: + t_jacr = wp.tile_map( + support._compute_jacr, + cdof_tile, + wp.tile_load(body_isdofancestor[b2_2], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacr2_tile = wp.tile_map(wp.add, jacr2_tile, wp.tile_map(wp.mul, t_jacr, weights2[2])) + if b2_3 >= 0: + t_jacr = wp.tile_map( + support._compute_jacr, + cdof_tile, + wp.tile_load(body_isdofancestor[b2_3], shape=TILE_SIZE, offset=dof_start, bounds_check=True), + ) + jacr2_tile = wp.tile_map(wp.add, jacr2_tile, wp.tile_map(wp.mul, t_jacr, weights2[3])) + + jacr_dif_tile = wp.tile_map(wp.sub, jacr2_tile, jacr1_tile) + + if not wp.static(IS_ELLIPTIC): + frame_0 = frame_in[conid, 0] + Ji_0p_tile = wp.tile_map(wp.dot, jacp_dif_tile, frame_0) + + if condim > 1: + Ji_0r_tile = wp.tile_map(wp.dot, jacr_dif_tile, frame_0) + frame_1 = frame_in[conid, 1] + Ji_1p_tile = wp.tile_map(wp.dot, jacp_dif_tile, frame_1) + Ji_1r_tile = wp.tile_map(wp.dot, jacr_dif_tile, frame_1) + frame_2 = frame_in[conid, 2] + Ji_2p_tile = wp.tile_map(wp.dot, jacp_dif_tile, frame_2) + Ji_2r_tile = wp.tile_map(wp.dot, jacr_dif_tile, frame_2) + + if wp.static(IS_ELLIPTIC): + dimid = efcid - contact_efc_address_in[conid, 0] + if dimid < 3: + frame_idx = dimid + else: + frame_idx = dimid - 3 + + frame_row = frame_in[conid, frame_idx] + + if dimid < 3: + J_tile = wp.tile_map(wp.dot, jacp_dif_tile, frame_row) + else: + J_tile = wp.tile_map(wp.dot, jacr_dif_tile, frame_row) + else: + J_tile = Ji_0p_tile + if condim > 1: + dimid = efcid - contact_efc_address_in[conid, 0] + dimid2 = dimid / 2 + 1 + frii = friction_in[conid, dimid2 - 1] + frii_sign = frii * (1.0 - 2.0 * float(dimid & 1)) + + if dimid2 == 1: + J_tile = wp.tile_map(wp.add, J_tile, wp.tile_map(wp.mul, Ji_1p_tile, frii_sign)) + elif dimid2 == 2: + J_tile = wp.tile_map(wp.add, J_tile, wp.tile_map(wp.mul, Ji_2p_tile, frii_sign)) + elif dimid2 == 3: + J_tile = wp.tile_map(wp.add, J_tile, wp.tile_map(wp.mul, Ji_0r_tile, frii_sign)) + elif dimid2 == 4: + J_tile = wp.tile_map(wp.add, J_tile, wp.tile_map(wp.mul, Ji_1r_tile, frii_sign)) + else: + J_tile = wp.tile_map(wp.add, J_tile, wp.tile_map(wp.mul, Ji_2r_tile, frii_sign)) + + wp.tile_store(efc_J_out[worldid, efcid], J_tile, offset=dof_start, bounds_check=True) + + Jqvel_tile = wp.tile_map(wp.mul, J_tile, qvel_tile) + Jqvel_sum = wp.tile_reduce(wp.add, Jqvel_tile) + if tid == 0: + wp.atomic_add(efc_Jqvel_out[worldid], efcid, wp.tile_extract(Jqvel_sum, 0)) + + return kernel + + @cache_kernel def _efc_contact_update(cone_type: types.ConeType): IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC @@ -2357,8 +3305,6 @@ def _efc_contact_update(cone_type: types.ConeType): opt_impratio_invsqrt: wp.array[float], body_invweight0: wp.array2d[wp.vec2], geom_bodyid: wp.array[int], - flex_vertadr: wp.array[int], - flex_vertbodyid: wp.array[int], # Data in: contact_efc_address_in: wp.array2d[int], efc_Jqvel_in: wp.array2d[float], @@ -2369,8 +3315,6 @@ def _efc_contact_update(cone_type: types.ConeType): includemargin_in: wp.array[float], worldid_in: wp.array[int], geom_in: wp.array[wp.vec2i], - flex_in: wp.array[wp.vec2i], - vert_in: wp.array[wp.vec2i], friction_in: wp.array[vec5], solref_in: wp.array[wp.vec2], solreffriction_in: wp.array[wp.vec2], @@ -2419,19 +3363,8 @@ def _efc_contact_update(cone_type: types.ConeType): geom = geom_in[conid] Jqvel = efc_Jqvel_in[worldid, efcid] - if geom[0] >= 0: - body1 = geom_bodyid[geom[0]] - else: - flex = flex_in[conid] - vert = vert_in[conid] - body1 = flex_vertbodyid[flex_vertadr[flex[0]] + vert[0]] - - if geom[1] >= 0: - body2 = geom_bodyid[geom[1]] - else: - flex = flex_in[conid] - vert = vert_in[conid] - body2 = flex_vertbodyid[flex_vertadr[flex[1]] + vert[1]] + body1 = geom_bodyid[geom[0]] + body2 = geom_bodyid[geom[1]] body_invweight0_id = worldid % body_invweight0.shape[0] invweight = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, body2][0] @@ -2499,21 +3432,240 @@ def _efc_contact_update(cone_type: types.ConeType): return kernel +@cache_kernel +def _efc_contact_update_flex(cone_type: types.ConeType): + IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC + + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # Model: + opt_timestep: wp.array[float], + opt_disableflags: int, + opt_impratio_invsqrt: wp.array[float], + body_invweight0: wp.array2d[wp.vec2], + geom_bodyid: wp.array[int], + flex_dim: wp.array[int], + flex_vertadr: wp.array[int], + flex_elemdataadr: wp.array[int], + flex_shelldataadr: wp.array[int], + flex_vertbodyid: wp.array[int], + flex_elem: wp.array[int], + flex_shell: wp.array[int], + # Data in: + flexvert_xpos_in: wp.array2d[wp.vec3], + contact_efc_address_in: wp.array2d[int], + efc_Jqvel_in: wp.array2d[float], + nacon_in: wp.array[int], + # In: + dist_in: wp.array[float], + pos_in: wp.array[wp.vec3], + condim_in: wp.array[int], + includemargin_in: wp.array[float], + worldid_in: wp.array[int], + geom_in: wp.array[wp.vec2i], + flex_in: wp.array[wp.vec2i], + elem_in: wp.array[wp.vec2i], + vert_in: wp.array[wp.vec2i], + friction_in: wp.array[vec5], + solref_in: wp.array[wp.vec2], + solreffriction_in: wp.array[wp.vec2], + solimp_in: wp.array[vec5], + type_in: wp.array[int], + # Data out: + efc_type_out: wp.array2d[int], + efc_id_out: wp.array2d[int], + efc_pos_out: wp.array2d[float], + efc_margin_out: wp.array2d[float], + efc_D_out: wp.array2d[float], + efc_vel_out: wp.array2d[float], + efc_aref_out: wp.array2d[float], + efc_frictionloss_out: wp.array2d[float], + ): + conid, dimid = wp.tid() + + if conid >= nacon_in[0]: + return + + if not type_in[conid] & ContactType.CONSTRAINT: + return + + condim = condim_in[conid] + + if wp.static(IS_ELLIPTIC): + if dimid > condim - 1: + return + else: + if condim == 1 and dimid > 0: + return + elif condim > 1 and dimid >= 2 * (condim - 1): + return + + efcid = contact_efc_address_in[conid, dimid] + if efcid < 0: + return + + worldid = worldid_in[conid] + timestep = opt_timestep[worldid % opt_timestep.shape[0]] + impratio_invsqrt = opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] + + includemargin = includemargin_in[conid] + pos = dist_in[conid] - includemargin + + geom = geom_in[conid] + flex = flex_in[conid] + elem = elem_in[conid] + vert = vert_in[conid] + con_pos = pos_in[conid] + Jqvel = efc_Jqvel_in[worldid, efcid] + + body_ids1, weights1 = _get_contact_bodies_and_weights( + geom_bodyid, + flex_dim, + flex_vertadr, + flex_elemdataadr, + flex_shelldataadr, + flex_vertbodyid, + flex_elem, + flex_shell, + flexvert_xpos_in, + conid, + 0, + geom, + flex, + elem, + vert, + con_pos, + worldid, + ) + body_ids2, weights2 = _get_contact_bodies_and_weights( + geom_bodyid, + flex_dim, + flex_vertadr, + flex_elemdataadr, + flex_shelldataadr, + flex_vertbodyid, + flex_elem, + flex_shell, + flexvert_xpos_in, + conid, + 1, + geom, + flex, + elem, + vert, + con_pos, + worldid, + ) + + b1_0 = body_ids1[0] + b1_1 = body_ids1[1] + b1_2 = body_ids1[2] + b1_3 = body_ids1[3] + + b2_0 = body_ids2[0] + b2_1 = body_ids2[1] + b2_2 = body_ids2[2] + b2_3 = body_ids2[3] + + body_invweight0_id = worldid % body_invweight0.shape[0] + + invweight1 = weights1[0] * body_invweight0[body_invweight0_id, b1_0][0] + if b1_1 >= 0: + invweight1 += weights1[1] * body_invweight0[body_invweight0_id, b1_1][0] + if b1_2 >= 0: + invweight1 += weights1[2] * body_invweight0[body_invweight0_id, b1_2][0] + if b1_3 >= 0: + invweight1 += weights1[3] * body_invweight0[body_invweight0_id, b1_3][0] + + invweight2 = weights2[0] * body_invweight0[body_invweight0_id, b2_0][0] + if b2_1 >= 0: + invweight2 += weights2[1] * body_invweight0[body_invweight0_id, b2_1][0] + if b2_2 >= 0: + invweight2 += weights2[2] * body_invweight0[body_invweight0_id, b2_2][0] + if b2_3 >= 0: + invweight2 += weights2[3] * body_invweight0[body_invweight0_id, b2_3][0] + + invweight = invweight1 + invweight2 + + ref = solref_in[conid] + pos_aref = pos + + if wp.static(IS_ELLIPTIC): + if dimid > 0: + solreffriction = solreffriction_in[conid] + + # non-normal directions use solreffriction (if non-zero) + if solreffriction[0] or solreffriction[1]: + ref = solreffriction + + invweight = invweight * impratio_invsqrt * impratio_invsqrt + friction = friction_in[conid] + + if dimid > 1: + fri0 = friction[0] + frii = friction[dimid - 1] + fri = fri0 * fri0 / (frii * frii) + invweight *= fri + + pos_aref = 0.0 + else: + if condim > 1: + friction = friction_in[conid] + fri0 = friction[0] + invweight = invweight + fri0 * fri0 * invweight + invweight = invweight * 2.0 * fri0 * fri0 * impratio_invsqrt * impratio_invsqrt + + if condim == 1: + efc_type = ConstraintType.CONTACT_FRICTIONLESS + elif wp.static(IS_ELLIPTIC): + efc_type = ConstraintType.CONTACT_ELLIPTIC + else: + efc_type = ConstraintType.CONTACT_PYRAMIDAL + + _efc_row( + opt_disableflags, + worldid, + timestep, + efcid, + pos_aref, + pos, + invweight, + ref, + solimp_in[conid], + includemargin, + Jqvel, + 0.0, + efc_type, + conid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, + ) + + return kernel + + @event_scope def make_constraint(m: types.Model, d: types.Data): """Creates constraint jacobians and other supporting data.""" + newton = m.opt.solver == types.SolverType.NEWTON efc_nnz = wp.empty((d.nworld,), dtype=int) wp.launch( _zero_constraint_counts, dim=d.nworld, - inputs=[d.ne, d.nf, d.nl, d.nefc, efc_nnz], + inputs=[d.ne, d.nf, d.nl, d.nefc, d.efc.jtdaj_nblock, efc_nnz], ) if not (m.opt.disableflags & types.DisableBit.CONSTRAINT): if not (m.opt.disableflags & types.DisableBit.EQUALITY): wp.launch( - _equality_connect, + _equality_connect(m.is_sparse, newton), dim=(d.nworld, m.eq_connect_adr.size), inputs=[ m.nv, @@ -2538,7 +3690,6 @@ def make_constraint(m: types.Model, d: types.Data): m.eq_solref, m.eq_solimp, m.eq_data, - m.is_sparse, m.body_isdofancestor, m.eq_connect_adr, d.qvel, @@ -2559,6 +3710,9 @@ def make_constraint(m: types.Model, d: types.Data): d.nefc, d.efc.type, d.efc.id, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, @@ -2573,7 +3727,7 @@ def make_constraint(m: types.Model, d: types.Data): ], ) wp.launch( - _equality_weld, + _equality_weld(m.is_sparse, newton), dim=(d.nworld, m.eq_wld_adr.size), inputs=[ m.nv, @@ -2599,7 +3753,6 @@ def make_constraint(m: types.Model, d: types.Data): m.eq_solref, m.eq_solimp, m.eq_data, - m.is_sparse, m.body_isdofancestor, m.eq_wld_adr, d.qvel, @@ -2621,6 +3774,9 @@ def make_constraint(m: types.Model, d: types.Data): d.nefc, d.efc.type, d.efc.id, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, @@ -2635,7 +3791,7 @@ def make_constraint(m: types.Model, d: types.Data): ], ) wp.launch( - _equality_joint, + _equality_joint(m.is_sparse, newton), dim=(d.nworld, m.eq_jnt_adr.size), inputs=[ m.nv, @@ -2650,7 +3806,6 @@ def make_constraint(m: types.Model, d: types.Data): m.eq_solref, m.eq_solimp, m.eq_data, - m.is_sparse, m.eq_jnt_adr, d.qpos, d.qvel, @@ -2663,6 +3818,9 @@ def make_constraint(m: types.Model, d: types.Data): d.nefc, d.efc.type, d.efc.id, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, @@ -2677,7 +3835,7 @@ def make_constraint(m: types.Model, d: types.Data): ], ) wp.launch( - _equality_tendon, + _equality_tendon(m.is_sparse, newton), dim=(d.nworld, m.eq_ten_adr.size), inputs=[ m.nv, @@ -2693,7 +3851,6 @@ def make_constraint(m: types.Model, d: types.Data): m.ten_J_colind, m.tendon_length0, m.tendon_invweight0, - m.is_sparse, m.eq_ten_adr, d.qvel, d.eq_active, @@ -2707,6 +3864,9 @@ def make_constraint(m: types.Model, d: types.Data): d.nefc, d.efc.type, d.efc.id, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, @@ -2722,7 +3882,7 @@ def make_constraint(m: types.Model, d: types.Data): ) wp.launch( - _equality_flex(m.is_sparse), + _equality_flex(m.is_sparse, newton), dim=(d.nworld, m.eq_flex_adr.size, m.nflexedge), inputs=[ m.nv, @@ -2751,6 +3911,9 @@ def make_constraint(m: types.Model, d: types.Data): d.nefc, d.efc.type, d.efc.id, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, @@ -2767,7 +3930,7 @@ def make_constraint(m: types.Model, d: types.Data): if not (m.opt.disableflags & types.DisableBit.FRICTIONLOSS): wp.launch( - _friction_dof, + _friction_dof(m.is_sparse, newton), dim=(d.nworld, m.nv), inputs=[ m.nv, @@ -2777,7 +3940,6 @@ def make_constraint(m: types.Model, d: types.Data): m.dof_solimp, m.dof_frictionloss, m.dof_invweight0, - m.is_sparse, d.qvel, d.njmax, d.njmax_nnz, @@ -2787,6 +3949,9 @@ def make_constraint(m: types.Model, d: types.Data): d.nefc, d.efc.type, d.efc.id, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, @@ -2802,7 +3967,7 @@ def make_constraint(m: types.Model, d: types.Data): ) wp.launch( - _friction_tendon, + _friction_tendon(m.is_sparse, newton), dim=(d.nworld, m.ntendon), inputs=[ m.nv, @@ -2815,7 +3980,6 @@ def make_constraint(m: types.Model, d: types.Data): m.tendon_solimp_fri, m.tendon_frictionloss, m.tendon_invweight0, - m.is_sparse, d.qvel, d.ten_J, d.njmax, @@ -2826,6 +3990,9 @@ def make_constraint(m: types.Model, d: types.Data): d.nefc, d.efc.type, d.efc.id, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, @@ -2843,7 +4010,7 @@ def make_constraint(m: types.Model, d: types.Data): # limit if not (m.opt.disableflags & types.DisableBit.LIMIT): wp.launch( - _limit_ball, + _limit_ball(m.is_sparse, newton), dim=(d.nworld, m.jnt_limited_ball_adr.size), inputs=[ m.nv, @@ -2856,7 +4023,6 @@ def make_constraint(m: types.Model, d: types.Data): m.jnt_range, m.jnt_margin, m.dof_invweight0, - m.is_sparse, m.jnt_limited_ball_adr, d.qpos, d.qvel, @@ -2868,6 +4034,9 @@ def make_constraint(m: types.Model, d: types.Data): d.nefc, d.efc.type, d.efc.id, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, @@ -2883,7 +4052,7 @@ def make_constraint(m: types.Model, d: types.Data): ) wp.launch( - _limit_slide_hinge, + _limit_slide_hinge(m.is_sparse, newton), dim=(d.nworld, m.jnt_limited_slide_hinge_adr.size), inputs=[ m.nv, @@ -2896,7 +4065,6 @@ def make_constraint(m: types.Model, d: types.Data): m.jnt_range, m.jnt_margin, m.dof_invweight0, - m.is_sparse, m.jnt_limited_slide_hinge_adr, d.qpos, d.qvel, @@ -2908,6 +4076,9 @@ def make_constraint(m: types.Model, d: types.Data): d.nefc, d.efc.type, d.efc.id, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, @@ -2923,7 +4094,7 @@ def make_constraint(m: types.Model, d: types.Data): ) wp.launch( - _limit_tendon, + _limit_tendon(m.is_sparse, newton), dim=(d.nworld, m.tendon_limited_adr.size), inputs=[ m.nv, @@ -2937,7 +4108,6 @@ def make_constraint(m: types.Model, d: types.Data): m.tendon_range, m.tendon_margin, m.tendon_invweight0, - m.is_sparse, m.tendon_limited_adr, d.qvel, d.ten_J, @@ -2950,6 +4120,9 @@ def make_constraint(m: types.Model, d: types.Data): d.nefc, d.efc.type, d.efc.id, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, @@ -2984,152 +4157,324 @@ def make_constraint(m: types.Model, d: types.Data): copy=False, ) - wp.launch( - _efc_contact_init(m.opt.cone, m.is_sparse), - dim=d.naconmax, - inputs=[ - m.body_weldid, - m.body_dofnum, - m.body_dofadr, - m.dof_parentid, - m.geom_bodyid, - m.flex_vertadr, - m.flex_vertbodyid, - d.njmax, - d.njmax_nnz, - d.nacon, - d.contact.dist, - d.contact.dim, - d.contact.includemargin, - d.contact.worldid, - d.contact.geom, - d.contact.flex, - d.contact.vert, - d.contact.type, - ], - outputs=[ - d.nefc, - d.contact.efc_address, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - efc_nnz, - ], - ) + has_flex = m.nflex > 0 - if m.is_sparse: + if has_flex: wp.launch( - _efc_contact_jac_sparse(m.opt.cone), - dim=(d.naconmax, nmaxdim), + _efc_contact_init_flex(m.opt.cone, m.is_sparse, newton), + dim=d.naconmax, inputs=[ m.body_parentid, - m.body_rootid, m.body_weldid, m.body_dofnum, m.body_dofadr, - m.dof_bodyid, m.dof_parentid, m.geom_bodyid, + m.flex_dim, m.flex_vertadr, + m.flex_elemdataadr, + m.flex_shelldataadr, m.flex_vertbodyid, - m.body_isdofancestor, - d.qvel, - d.subtree_com, - d.cdof, - d.contact.efc_address, - d.efc.J_rownnz, - d.efc.J_rowadr, + m.flex_elem, + m.flex_shell, + d.flexvert_xpos, + d.njmax, + d.njmax_nnz, d.nacon, + d.contact.dist, + d.contact.pos, d.contact.dim, + d.contact.includemargin, + d.contact.worldid, d.contact.geom, d.contact.flex, + d.contact.elem, d.contact.vert, - d.contact.pos, - contact_frame_2d, - contact_friction_2d, - d.contact.worldid, + d.contact.type, ], outputs=[ - d.efc.J_colind, - d.efc.J, - d.efc.Jqvel, + d.nefc, + d.contact.efc_address, + d.efc.id, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, + d.efc.J_rownnz, + d.efc.J_rowadr, + efc_nnz, ], ) + else: + wp.launch( + _efc_contact_init(m.opt.cone, m.is_sparse, newton), + dim=d.naconmax, + inputs=[ + m.body_weldid, + m.body_dofnum, + m.body_dofadr, + m.dof_parentid, + m.geom_bodyid, + d.njmax, + d.njmax_nnz, + d.nacon, + d.contact.dist, + d.contact.dim, + d.contact.includemargin, + d.contact.worldid, + d.contact.geom, + d.contact.type, + ], + outputs=[ + d.nefc, + d.contact.efc_address, + d.efc.id, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, + d.efc.J_rownnz, + d.efc.J_rowadr, + efc_nnz, + ], + ) + + if m.is_sparse: + if has_flex: + wp.launch( + _efc_contact_jac_sparse_flex(m.opt.cone), + dim=(d.naconmax, nmaxdim), + inputs=[ + m.body_parentid, + m.body_rootid, + m.body_weldid, + m.body_dofnum, + m.body_dofadr, + m.dof_bodyid, + m.dof_parentid, + m.geom_bodyid, + m.flex_dim, + m.flex_vertadr, + m.flex_elemdataadr, + m.flex_shelldataadr, + m.flex_vertbodyid, + m.flex_elem, + m.flex_shell, + m.body_isdofancestor, + d.qvel, + d.subtree_com, + d.cdof, + d.flexvert_xpos, + d.contact.efc_address, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.nacon, + d.contact.dim, + d.contact.geom, + d.contact.flex, + d.contact.elem, + d.contact.vert, + d.contact.pos, + contact_frame_2d, + contact_friction_2d, + d.contact.worldid, + ], + outputs=[ + d.efc.J_colind, + d.efc.J, + d.efc.Jqvel, + ], + ) + else: + wp.launch( + _efc_contact_jac_sparse(m.opt.cone), + dim=(d.naconmax, nmaxdim), + inputs=[ + m.body_parentid, + m.body_rootid, + m.body_weldid, + m.body_dofnum, + m.body_dofadr, + m.dof_bodyid, + m.dof_parentid, + m.geom_bodyid, + m.body_isdofancestor, + d.qvel, + d.subtree_com, + d.cdof, + d.contact.efc_address, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.nacon, + d.contact.dim, + d.contact.geom, + d.contact.pos, + contact_frame_2d, + contact_friction_2d, + d.contact.worldid, + ], + outputs=[ + d.efc.J_colind, + d.efc.J, + d.efc.Jqvel, + ], + ) else: d.efc.Jqvel.zero_() tile_size = m.block_dim.contact_jac_tiled n_dof_blocks = (m.nv_pad + tile_size - 1) // tile_size - wp.launch_tiled( - _efc_contact_jac_dense(tile_size, m.opt.cone), - dim=(d.nworld, n_dof_blocks), + if has_flex: + wp.launch_tiled( + _efc_contact_jac_dense_flex(tile_size, m.opt.cone), + dim=(d.nworld, n_dof_blocks), + inputs=[ + m.body_rootid, + m.geom_bodyid, + m.flex_dim, + m.flex_vertadr, + m.flex_elemdataadr, + m.flex_shelldataadr, + m.flex_vertbodyid, + m.flex_elem, + m.flex_shell, + m.body_isdofancestor, + d.ne, + d.nf, + d.nl, + d.nefc, + d.qvel, + d.subtree_com, + d.cdof, + d.flexvert_xpos, + d.contact.efc_address, + d.efc.id, + d.njmax, + m.nv_pad, + d.contact.dim, + d.contact.geom, + d.contact.flex, + d.contact.elem, + d.contact.vert, + d.contact.pos, + contact_frame_2d, + contact_friction_2d, + ], + outputs=[ + d.efc.J, + d.efc.Jqvel, + ], + block_dim=tile_size, + ) + else: + wp.launch_tiled( + _efc_contact_jac_dense(tile_size, m.opt.cone), + dim=(d.nworld, n_dof_blocks), + inputs=[ + m.body_rootid, + m.geom_bodyid, + m.body_isdofancestor, + d.ne, + d.nf, + d.nl, + d.nefc, + d.qvel, + d.subtree_com, + d.cdof, + d.contact.efc_address, + d.efc.id, + d.njmax, + m.nv_pad, + d.contact.dim, + d.contact.geom, + d.contact.pos, + contact_frame_2d, + contact_friction_2d, + ], + outputs=[ + d.efc.J, + d.efc.Jqvel, + ], + block_dim=tile_size, + ) + + if has_flex: + wp.launch( + _efc_contact_update_flex(m.opt.cone), + dim=(d.naconmax, nmaxdim), inputs=[ - m.body_rootid, + m.opt.timestep, + m.opt.disableflags, + m.opt.impratio_invsqrt, + m.body_invweight0, m.geom_bodyid, + m.flex_dim, m.flex_vertadr, + m.flex_elemdataadr, + m.flex_shelldataadr, m.flex_vertbodyid, - m.body_isdofancestor, - d.ne, - d.nf, - d.nl, - d.nefc, - d.qvel, - d.subtree_com, - d.cdof, + m.flex_elem, + m.flex_shell, + d.flexvert_xpos, d.contact.efc_address, - d.efc.id, - d.njmax, - m.nv_pad, + d.efc.Jqvel, + d.nacon, + d.contact.dist, + d.contact.pos, d.contact.dim, + d.contact.includemargin, + d.contact.worldid, d.contact.geom, d.contact.flex, + d.contact.elem, d.contact.vert, - d.contact.pos, - contact_frame_2d, - contact_friction_2d, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + d.contact.type, ], outputs=[ - d.efc.J, - d.efc.Jqvel, + d.efc.type, + d.efc.id, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + ], + ) + else: + wp.launch( + _efc_contact_update(m.opt.cone), + dim=(d.naconmax, nmaxdim), + inputs=[ + m.opt.timestep, + m.opt.disableflags, + m.opt.impratio_invsqrt, + m.body_invweight0, + m.geom_bodyid, + d.contact.efc_address, + d.efc.Jqvel, + d.nacon, + d.contact.dist, + d.contact.dim, + d.contact.includemargin, + d.contact.worldid, + d.contact.geom, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + d.contact.type, + ], + outputs=[ + d.efc.type, + d.efc.id, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, ], - block_dim=tile_size, ) - - wp.launch( - _efc_contact_update(m.opt.cone), - dim=(d.naconmax, nmaxdim), - inputs=[ - m.opt.timestep, - m.opt.disableflags, - m.opt.impratio_invsqrt, - m.body_invweight0, - m.geom_bodyid, - m.flex_vertadr, - m.flex_vertbodyid, - d.contact.efc_address, - d.efc.Jqvel, - d.nacon, - d.contact.dist, - d.contact.dim, - d.contact.includemargin, - d.contact.worldid, - d.contact.geom, - d.contact.flex, - d.contact.vert, - d.contact.friction, - d.contact.solref, - d.contact.solreffriction, - d.contact.solimp, - d.contact.type, - ], - outputs=[ - d.efc.type, - d.efc.id, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], - ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py index 28c8193f..9329302b 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py @@ -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], + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py index 60375d5d..633c2a3a 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py @@ -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) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py index 388ad1aa..c6809d0e 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py @@ -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) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py index c90451ab..9a2fd58e 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py @@ -24,6 +24,7 @@ import warp as wp from mujoco.mjx.third_party.mujoco_warp._src import bvh from mujoco.mjx.third_party.mujoco_warp._src import math as mjmath from mujoco.mjx.third_party.mujoco_warp._src import render_util +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 types from mujoco.mjx.third_party.mujoco_warp._src import warp_util @@ -31,10 +32,6 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType from mujoco.mjx.third_party.mujoco_warp._src.types import vec10 -from mujoco.mjx.third_party.mujoco_warp._src.util_pkg import check_version - -# TODO(team): remove after improving island solver performance -ENABLE_ISLANDS = False def _is_array_spec(typ) -> bool: @@ -42,28 +39,54 @@ def _is_array_spec(typ) -> bool: return isinstance(typ, wp.array) or type(typ).__name__ == "_ArrayAnnotation" -def _create_array(data: Any, spec, sizes: dict[str, int]) -> wp.array | None: +def _mark_batched(obj): + """Recursively set _is_batched = True on all batched warp arrays within obj.""" + if not dataclasses.is_dataclass(obj): + return + for f in dataclasses.fields(obj): + val = getattr(obj, f.name, None) + if val is None: + continue + if dataclasses.is_dataclass(val): + _mark_batched(val) + continue + if not isinstance(val, wp.array): + continue + if not _is_array_spec(f.type): + continue + spec_shape = getattr(f.type, "shape", ()) + if spec_shape and spec_shape[0] in ("*", "nworld"): + val._is_batched = True + + +def _create_array(data: Any, spec, sizes: dict[str, int], batch_size: int = 1) -> wp.array | None: """Creates a warp array and populates it with data. The array shape is determined by a field spec referencing MjModel / MjData array sizes. """ spec_shape = getattr(spec, "shape", (0,)) - shape = None - if spec_shape != (0,): - shape = tuple(sizes[dim] if isinstance(dim, str) else dim for dim in spec_shape) + if spec_shape == (0,): + if data is None: + return None + return wp.array(np.array(data), dtype=spec.dtype) - if data is None and shape is None: - return None # nothing to do - elif data is None: + shape = tuple(batch_size if dim == "*" else (sizes[dim] if isinstance(dim, str) else dim) for dim in spec_shape) + + is_batched = spec_shape[0] in ("*", "nworld") + + if data is None: array = wp.zeros(shape, dtype=spec.dtype) else: - array = wp.array(np.array(data), dtype=spec.dtype, shape=shape) + data = np.array(data) + if is_batched and shape[0] != 1: + target_shape = shape + getattr(spec.dtype, "_shape_", ()) + if data.shape != target_shape: + tail_shape = target_shape[1:] + if data.size == np.prod(tail_shape): + data = data.reshape(tail_shape) + data = np.broadcast_to(data, target_shape).copy() + array = wp.array(data, dtype=spec.dtype, shape=shape) - if spec_shape and spec_shape[0] == "*": - # add private attribute for JAX to determine which fields are batched - array._is_batched = True - # also set stride 0 to 0 which is expected legacy behavior (but is deprecated) - array.strides = (0,) + array.strides[1:] return array @@ -73,55 +96,53 @@ def _create_constraint( njmax: int, njmax_nnz: int, sizes: dict, - island_enabled: bool, mjd=None, ) -> types.Constraint: """Construct a types.Constraint with standard and island local fields allocated properly.""" efc_kwargs = {"J_rownnz": None, "J_rowadr": None, "J_colind": None, "J": None} sparse = is_sparse(mjm) + # The JTDAJ block list is only consumed by the sparse Newton Hessian assembly (_JTDAJ_sparse). + jtdaj_active = sparse and mjm.opt.solver == mujoco.mjtSolver.mjSOL_NEWTON for f in dataclasses.fields(types.Constraint): - if f.name == "itype": - efc_kwargs[f.name] = wp.empty((nworld, njmax if island_enabled else 0), dtype=int) - elif f.name == "iid": - efc_kwargs[f.name] = wp.empty((nworld, njmax if island_enabled else 0), dtype=int) - elif f.name == "iJ_rownnz": - efc_kwargs[f.name] = wp.empty((nworld, njmax if island_enabled else 0) if sparse else (nworld, 0), dtype=int) - elif f.name == "iJ_rowadr": - efc_kwargs[f.name] = wp.empty((nworld, njmax if island_enabled else 0) if sparse else (nworld, 0), dtype=int) - elif f.name == "iJ_colind": - efc_kwargs[f.name] = wp.empty((nworld, 1, njmax_nnz if island_enabled else 0) if sparse else (nworld, 0, 0), dtype=int) - elif f.name == "iJ": - efc_kwargs[f.name] = wp.empty( - (nworld, 1, njmax_nnz if island_enabled else 0) if sparse else (nworld, njmax if island_enabled else 0, mjm.nv), - dtype=float, - ) - elif f.name == "iD": - efc_kwargs[f.name] = wp.empty((nworld, njmax if island_enabled else 0), dtype=float) - elif f.name == "iaref": - efc_kwargs[f.name] = wp.empty((nworld, njmax if island_enabled else 0), dtype=float) - elif f.name == "ifrictionloss": - efc_kwargs[f.name] = wp.empty((nworld, njmax if island_enabled else 0), dtype=float) - elif f.name == "iforce": - efc_kwargs[f.name] = wp.empty((nworld, njmax if island_enabled else 0), dtype=float) - elif f.name == "istate": - efc_kwargs[f.name] = wp.empty((nworld, njmax if island_enabled else 0), dtype=int) + if f.name in ("jtdaj_adr", "jtdaj_nrow"): + efc_kwargs[f.name] = wp.empty((nworld, njmax if jtdaj_active else 0), dtype=int) + elif f.name == "jtdaj_nblock": + efc_kwargs[f.name] = wp.empty((nworld,), dtype=int) else: if f.name in efc_kwargs: continue - if mjd is not None: - shape = tuple(sizes[dim] if isinstance(dim, str) else dim for dim in f.type.shape) - val = np.zeros(shape, dtype=f.type.dtype) - if f.name in ("type", "id", "pos", "margin", "D", "vel", "aref", "frictionloss", "force"): - val[:, : mjd.nefc] = np.tile(getattr(mjd, "efc_" + f.name), (nworld, 1)) - efc_kwargs[f.name] = wp.array(val, dtype=f.type.dtype) - else: - efc_kwargs[f.name] = _create_array(None, f.type, sizes) + if mjd is not None: + shape = tuple(sizes[dim] if isinstance(dim, str) else dim for dim in f.type.shape) + val = np.zeros(shape, dtype=f.type.dtype) + if f.name in ("type", "id", "pos", "margin", "D", "vel", "aref", "frictionloss", "force"): + val[:, : mjd.nefc] = np.tile(getattr(mjd, "efc_" + f.name), (nworld, 1)) + efc_kwargs[f.name] = wp.array(val, dtype=f.type.dtype) + else: + efc_kwargs[f.name] = _create_array(None, f.type, sizes) return types.Constraint(**efc_kwargs) +def _jtdaj_groups(mjd: mujoco.MjData) -> tuple[np.ndarray, np.ndarray]: + """Group loaded efc rows into JTDAJ blocks: maximal runs sharing (efc_type, efc_id). + + MuJoCo lays each constraint's rows out contiguously, so this reproduces the block list + make_constraint builds in-kernel. Returns block start rows (adr) and lengths (nrow). + """ + nefc = mjd.nefc + if nefc == 0: + return np.zeros(0, dtype=int), np.zeros(0, dtype=int) + etype = mjd.efc_type[:nefc] + eid = mjd.efc_id[:nefc] + boundary = np.ones(nefc, dtype=bool) + boundary[1:] = (etype[1:] != etype[:-1]) | (eid[1:] != eid[:-1]) + adr = np.flatnonzero(boundary) + nrow = np.diff(np.append(adr, nefc)) + return adr, nrow + + def is_sparse(mjm: mujoco.MjModel) -> bool: if mjm.opt.jacobian == mujoco.mjtJacobian.mjJAC_AUTO: if mjm.nv > 32: @@ -132,11 +153,128 @@ def is_sparse(mjm: mujoco.MjModel) -> bool: return bool(mujoco.mj_isSparse(mjm)) -def put_model(mjm: mujoco.MjModel) -> types.Model: +def _m_blocks(mjm: mujoco.MjModel): + """The (start, size) diagonal blocks of M: the kinematic trees, each a contiguous dof range. + + M couples a dof only with its tree ancestors, so its diagonal blocks are exactly the trees. + (dof_simplenum is not used to classify blocks: it is a contiguous-suffix run-length, so it + misses interspersed decoupled dofs; the M_rownnz coupling check in m_block_layout catches those.) + """ + return [(int(adr), int(num)) for adr, num in zip(mjm.tree_dofadr, mjm.tree_dofnum) if num > 0] + + +def _m_allow_dense(mjm: mujoco.MjModel) -> bool: + """Whether any block may use the packed dense layout (tendon armature forces all-sparse).""" + # tendon armature accumulates into M in CSR layout, which the packed block layout cannot represent + return not (mjm.ntendon and np.any(mjm.tendon_armature)) + + +def m_block_layout(mjm: mujoco.MjModel) -> dict: + """Per-block dense/sparse layout for M's diagonal blocks. + + Blocks (connected sub-trees, each a contiguous dof range) are classified into three per-block + categories by coupling and size: + - simple: a decoupled block (M is diagonal -- a "simple body" like a point mass on orthogonal + slides) needs no factorization, just D = 1/diag, so it bypasses both factor paths. + - dense: a coupled block small enough for a dense tile-Cholesky (size <= M_BLOCK_DENSE_MAX). + - sparse: a coupled block too large for a tile, via the sparse LDL factor. + Dense block factors are packed back to back (block k's b*b factor at the prefix sum of preceding + dense block areas). Returns: + total: packed length of the dense region (also the offset of the LDL region) + dof_adr: per-dof packed offset within the dense region (0 for non-dense dofs) + blocks: all (start, size) blocks + dense_blocks: (start, size) blocks using the packed dense layout + dof_dense: per-dof flag, 1 if the dof's block is dense + dof_simple: per-dof flag, 1 if the dof's block is simple (diagonal) + has_dense / has_simple / has_sparse: whether any block falls in that category + """ + nv = mjm.nv + blocks = _m_blocks(mjm) + allow_dense = _m_allow_dense(mjm) + rownnz = mjm.M_rownnz + dof_adr = np.zeros(nv, dtype=np.int32) + dof_dense = np.zeros(nv, dtype=np.int32) + dof_simple = np.zeros(nv, dtype=np.int32) + dense_blocks = [] + off = 0 + has_sparse = False + for start, size in blocks: + coupled = bool(np.max(rownnz[start : start + size]) > 1) + if not coupled: + dof_simple[start : start + size] = 1 + elif allow_dense and size <= types.M_BLOCK_DENSE_MAX: + dense_blocks.append((start, size)) + dof_adr[start : start + size] = off + dof_dense[start : start + size] = 1 + off += size * size + else: + has_sparse = True + return { + "total": off, + "dof_adr": dof_adr, + "blocks": blocks, + "dense_blocks": dense_blocks, + "dof_dense": dof_dense, + "dof_simple": dof_simple, + "has_dense": len(dense_blocks) > 0, + "has_simple": bool(dof_simple.any()), + "has_sparse": has_sparse, + } + + +def _filter_tri_geoms( + mjm: mujoco.MjModel, + v0: int, + v1: int, + v2: int, + geomids: np.ndarray, + filterparent: bool, +) -> np.ndarray: + """Vectorized check for a single triangle vs multiple geoms.""" + b0 = mjm.flex_vertbodyid[v0] + b1 = mjm.flex_vertbodyid[v1] + b2 = mjm.flex_vertbodyid[v2] + + w0 = mjm.body_weldid[b0] + w1 = mjm.body_weldid[b1] + w2 = mjm.body_weldid[b2] + + bg = mjm.geom_bodyid[geomids] + wg = mjm.body_weldid[bg] + + is_self = (wg == w0) | (wg == w1) | (wg == w2) + + is_parent = np.zeros_like(is_self, dtype=bool) + if filterparent: + wp0 = mjm.body_weldid[mjm.body_parentid[w0]] + wp1 = mjm.body_weldid[mjm.body_parentid[w1]] + wp2 = mjm.body_weldid[mjm.body_parentid[w2]] + wpg = mjm.body_weldid[mjm.body_parentid[wg]] + + cond0 = (wg != 0) & (w0 != 0) & ((wg == wp0) | (w0 == wpg)) + cond1 = (wg != 0) & (w1 != 0) & ((wg == wp1) | (w1 == wpg)) + cond2 = (wg != 0) & (w2 != 0) & ((wg == wp2) | (w2 == wpg)) + is_parent = cond0 | cond1 | cond2 + + sig0 = (b0 << 16) + geomids + sig1 = (b1 << 16) + geomids + sig2 = (b2 << 16) + geomids + + is_excluded = ( + np.isin(sig0, mjm.exclude_signature) | np.isin(sig1, mjm.exclude_signature) | np.isin(sig2, mjm.exclude_signature) + ) + + return is_self | is_parent | is_excluded + + +def put_model(mjm: mujoco.MjModel, batch_sizes: dict[str, int] | None = None) -> types.Model: """Creates a model on device. Args: mjm: The model containing kinematic and dynamic information (host). + batch_sizes: Optional per-field leading batch sizes for `Model` fields whose + array spec starts with `*`. Fields not listed here keep the default shared + leading dimension of 1. Returns: The model containing kinematic and dynamic information (device). @@ -144,6 +282,16 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: # check for compatible cuda toolkit and driver versions warp_util.check_toolkit_driver() + batch_sizes = batch_sizes or {} + model_fields = {f.name: f.type for f in dataclasses.fields(types.Model) if _is_array_spec(f.type)} + for name, size in batch_sizes.items(): + field_type = model_fields.get(name) + spec_shape = getattr(field_type, "shape", ()) + if not spec_shape or spec_shape[0] != "*": + raise ValueError(f"Model field {name!r} is not a batched array field.") + if size < 1: + raise ValueError(f"batch_sizes[{name!r}] must be positive, got {size}.") + # model: check supported features in array types for field, field_type, mj_type in ( (mjm.actuator_trntype, types.TrnType, mujoco.mjtTrn), @@ -185,12 +333,6 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: if mjm.opt.noslip_iterations > 0: raise NotImplementedError(f"noslip solver not implemented.") - if (mjm.opt.viscosity > 0 or mjm.opt.density > 0) and mjm.opt.integrator in ( - mujoco.mjtIntegrator.mjINT_IMPLICITFAST, - mujoco.mjtIntegrator.mjINT_IMPLICIT, - ): - raise NotImplementedError(f"Implicit integrators and fluid model not implemented.") - if (mjm.body_plugin != -1).any(): raise NotImplementedError("Body plugins not supported.") @@ -209,6 +351,12 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: if mjm.nv > nv_max and mjm.opt.jacobian == mujoco.mjtJacobian.mjJAC_DENSE: raise ValueError(f"Dense is unsupported for nv > {nv_max} (nv = {mjm.nv}).") + # sleeping is supported via a dof-compaction approach. awake dofs are compacted into dense + # nvmax-sized arrays. nvmax is chosen to fit the worst-case active dof set. sleeping is only + # supported for Newton solver and requires nv <= nvmax. + if (mjm.opt.enableflags & mujoco.mjtEnableBit.mjENBL_SLEEP) and mjm.opt.solver != mujoco.mjtSolver.mjSOL_NEWTON: + raise ValueError(f"sleeping requires the Newton solver (got solver={types.SolverType(mjm.opt.solver).name})") + collision_sensors = (mujoco.mjtSensor.mjSENS_GEOMDIST, mujoco.mjtSensor.mjSENS_GEOMNORMAL, mujoco.mjtSensor.mjSENS_GEOMFROMTO) is_collision_sensor = np.isin(mjm.sensor_type, collision_sensors) @@ -245,12 +393,6 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: opt_kwargs["impratio_invsqrt"] = 1.0 / np.sqrt(np.maximum(mjm.opt.impratio, mujoco.mjMINVAL)) opt = types.Option(**opt_kwargs) - # islands are disabled by default while performance is being improved - # override by setting io.ENABLE_ISLANDS = True - # TODO(team): remove after improving island solver performance - if not ENABLE_ISLANDS: - opt.disableflags |= types.DisableBit.ISLAND - # C MuJoCo tolerance was chosen for float64 architecture, but we default to float32 on GPU # adjust the tolerance for lower precision, to avoid the solver spending iterations needlessly # bouncing around the optimal solution @@ -261,6 +403,7 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: opt.broadphase_filter = types.BroadphaseFilter.PLANE | types.BroadphaseFilter.SPHERE | types.BroadphaseFilter.OBB opt.graph_conditional = True opt.run_collision_detection = True + opt.warn_overflow = True contact_sensor_maxmatch_id = mujoco.mj_name2id(mjm, mujoco.mjtObj.mjOBJ_NUMERIC, "contact_sensor_maxmatch") if contact_sensor_maxmatch_id > -1: opt.contact_sensor_maxmatch = mjm.numeric_data[mjm.numeric_adr[contact_sensor_maxmatch_id]] @@ -297,6 +440,8 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: m.nmaxcondim = np.concatenate(condim_arrays).max() m.nmaxpyramid = np.maximum(1, 2 * (m.nmaxcondim - 1)) m.has_sdf_geom = (mjm.geom_type == mujoco.mjtGeom.mjGEOM_SDF).any() + m.has_flex_selfcollide = bool(mjm.nflex > 0 and np.any(mjm.flex_selfcollide != 0)) + m.max_flex_dim = int(np.max(mjm.flex_dim)) if mjm.nflex > 0 else 0 m.block_dim = types.BlockDim() # Derive CG solver block_dim from nv: clamp(round_up_to_32(nv), 32, 256) _nv_block = max(32, min(256, ((mjm.nv + 31) // 32) * 32)) @@ -308,29 +453,8 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: m.block_dim.linesearch_iterative = 512 m.is_sparse = is_sparse(mjm) m.has_fluid = mjm.opt.wind.any() or mjm.opt.density > 0 or mjm.opt.viscosity > 0 - m.max_ten_J_rownnz = int(mjm.ten_J_rownnz.max()) if mjm.ntendon else 0 - # Upper bound on a contact's Jacobian support, to size the elliptic-cone JTCJ launch (one - # thread per (contact, support-pair)). A contact's row spans the dof chains of weld(b1) and - # weld(b2) (see _efc_contact_jac_sparse in constraint.py); take the largest union over all - # geom-carrying bodies -- a safe superset, since over-estimating only adds skipped threads. - # A body's dof chain is exactly the sparsity of its deepest dof's row in the (ancestor- - # structured) mass matrix, so reuse MuJoCo's precomputed M_colind rather than re-walking. - def _dof_chain(body): - if mjm.body_dofnum[body] == 0: - return frozenset() - dof = int(mjm.body_dofadr[body] + mjm.body_dofnum[body] - 1) - adr = int(mjm.M_rowadr[dof]) - return frozenset(int(mjm.M_colind[adr + k]) for k in range(int(mjm.M_rownnz[dof]))) - - chains = list({_dof_chain(int(mjm.body_weldid[b])) for b in mjm.geom_bodyid}) - max_rownnz = 0 - for i, chain_i in enumerate(chains): - for chain_j in chains[i:]: - max_rownnz = max(max_rownnz, len(chain_i | chain_j)) - m.jtcj_max_pairs = max(max_rownnz * (max_rownnz + 1) // 2, 1) - # body ids grouped by tree level (depth-based traversal) bodies, body_depth = {}, np.zeros(mjm.nbody, dtype=int) - 1 for i in range(mjm.nbody): @@ -361,6 +485,12 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: m.mocap_bodyid = m.mocap_bodyid[mjm.body_mocapid[mjm.body_mocapid >= 0].argsort()] m.body_fluid_ellipsoid = np.zeros(mjm.nbody, dtype=bool) m.body_fluid_ellipsoid[mjm.geom_bodyid[mjm.geom_fluid.reshape(mjm.ngeom, mujoco.mjNFLUID)[:, 0] > 0]] = True + m.body_fluid_ellipsoid_adr = np.nonzero(m.body_fluid_ellipsoid)[0] + body_fluid_box = np.zeros(mjm.nbody, dtype=bool) + for b in range(1, mjm.nbody): + if not m.body_fluid_ellipsoid[b] and mjm.body_mass[b] > 0.0: + body_fluid_box[b] = True + m.body_fluid_box_adr = np.nonzero(body_fluid_box)[0] jnt_limited_slide_hinge = mjm.jnt_limited & np.isin(mjm.jnt_type, (mujoco.mjtJoint.mjJNT_SLIDE, mujoco.mjtJoint.mjJNT_HINGE)) m.jnt_limited_slide_hinge_adr = np.nonzero(jnt_limited_slide_hinge)[0] m.jnt_limited_ball_adr = np.nonzero(mjm.jnt_limited & (mjm.jnt_type == mujoco.mjtJoint.mjJNT_BALL))[0] @@ -381,6 +511,13 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: dofid = mjm.dof_parentid[dofid] m.body_isdofancestor = body_isdofancestor + # Upper bound on a contact's Jacobian support-pair count, to size the elliptic-cone JTCJ + # launch. Use body_isdofancestor (the full dof tree), not the mass-matrix sparsity, which the + # simple-dof optimization diagonalizes -- that undercounts the support and NaNs the solve. + support_chains = [set(np.flatnonzero(row).tolist()) for row in np.unique(body_isdofancestor[mjm.geom_bodyid], axis=0)] + max_support = max((len(ci | cj) for i, ci in enumerate(support_chains) for cj in support_chains[i:]), default=0) + m.jtcj_max_pairs = max(max_support * (max_support + 1) // 2, 1) + # precalculated geom pairs filterparent = not (mjm.opt.disableflags & types.DisableBit.FILTERPARENT) @@ -678,16 +815,27 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: for j in range(mjm.mesh_vertnum[mjm.sensor_objid[i]]) ] - # M_tiles records the block diagonal structure of M - tile_corners = [i for i in range(mjm.nv) if mjm.dof_parentid[i] == -1] + # Per-block dense/sparse layout (see m_block_layout). M_tiles holds the dense blocks grouped by + # size; a model may use both paths at once (e.g. one large tree + many small free joints). + _lay = m_block_layout(mjm) + dof_dense = _lay["dof_dense"] + dof_simple = _lay["dof_simple"] + m.qLD_has_dense = _lay["has_dense"] + m.qLD_has_simple = _lay["has_simple"] + m.qLD_has_sparse = _lay["has_sparse"] + m.qLD_block_total = _lay["total"] # packed dense region length / offset of the LDL region + m.qLD_block_adr = _lay["dof_adr"] + m.qLD_dof_dense = dof_dense # per-dof: 1 if the dof's block is dense (packed) + m.qLD_dof_simple = dof_simple # per-dof: 1 if the dof's block is simple (diagonal -> 1/diag) + m.qLD_simple_dofs = np.nonzero(dof_simple)[0].astype(np.int32) # the simple dof indices + tiles = {} - for i in range(len(tile_corners)): - tile_beg = tile_corners[i] - tile_end = mjm.nv if i == len(tile_corners) - 1 else tile_corners[i + 1] - tiles.setdefault(tile_end - tile_beg, []).append(tile_beg) + for start, size in _lay["dense_blocks"]: + tiles.setdefault(size, []).append(start) m.M_tiles = tuple(types.TileSet(adr=wp.array(tiles[sz], dtype=int), size=sz) for sz in sorted(tiles.keys())) - # qLD_updates has dof tree ordering of qLD updates for sparse factor m + # qLD_updates has dof tree ordering of qLD updates for the sparse LDL factor. Only sparse-block + # dofs are included; dense blocks use the packed Cholesky and never touch the LDL region. qLD_updates, dof_depth = {}, np.zeros(mjm.nv, dtype=int) - 1 for k in range(mjm.nv): @@ -695,6 +843,8 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: if mjm.M_rownnz[k] == 1: continue dof_depth[k] = dof_depth[mjm.dof_parentid[k]] + 1 + if dof_dense[k]: + continue # dense block: handled by the packed Cholesky, not the LDL factor i = mjm.dof_parentid[k] diag_k = mjm.M_rowadr[k] + mjm.M_rownnz[k] - 1 Madr_ki = diag_k - 1 @@ -727,6 +877,9 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: j = mjm.dof_parentid[j] # M_elemid maps (row, col) -> madr index in the native CSR M layout M_elemid = np.full((mjm.nv, mjm.nv), -1, dtype=np.int32) + # M_hinit_i: row index of each CSR entry (its madr is the flat index). The dense Newton H-init + # uses (M_hinit_i, M_colind) to scatter M's upper triangle into the dense H tile from CSR. + M_hinit_i = np.zeros(mjm.nC, dtype=np.int32) for i in range(mjm.nv): rowadr = mjm.M_rowadr[i] rownnz = mjm.M_rownnz[i] @@ -734,7 +887,23 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: madr = rowadr + k col = int(mjm.M_colind[madr]) M_elemid[i, col] = madr + M_hinit_i[madr] = i m.M_elemid = M_elemid + m.M_hinit_i = M_hinit_i + + # Precompute per-block gather indices for the dense-block densify (tile_load_indexed). For each + # dense block and flat slot (row, col), store the CSR address of M[max(i,j), min(i,j)], or nC + # (out of bounds -> read as 0) for structurally absent pairs. Laid out [block, slot] so the kernel + # reads slice (block_size^2,) at offset blk * block_size^2. + for tile in m.M_tiles: + sz = tile.size + starts = np.array(tiles[sz], dtype=np.int32) # host block starts; no device round-trip + dofs = starts[:, None] + np.arange(sz)[None, :] # (nblock, sz) global dof per block row + gi = dofs[:, :, None] # (nblock, sz, 1) + gj = dofs[:, None, :] # (nblock, 1, sz) + elemid = M_elemid[np.maximum(gi, gj), np.minimum(gi, gj)] # (nblock, sz, sz), -1 if absent + elemid = np.where(elemid >= 0, elemid, mjm.nC) + tile.elemid = wp.array(elemid.reshape(-1).astype(np.int32), dtype=int) upper_j, upper_i = np.triu_indices(mjm.nv) upper_elemid = M_elemid[upper_i, upper_j] @@ -781,20 +950,136 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: m.flexedge_J_rowadr = mjm.flexedge_J_rowadr m.flexedge_J_colind = mjm.flexedge_J_colind.reshape(-1) - # flex_bendingadr backward compat: flatten old (nflexedge, 17) to 1D - if not check_version("mujoco>=3.8.1.dev909088123"): - m.flex_bendingadr = ( - np.array([mjm.flex_edgeadr[i] * 17 for i in range(mjm.nflex)], dtype=int) if mjm.nflex else np.zeros(0, dtype=int) - ) - m.flex_bending = mjm.flex_bending.ravel() - m.nflexbending = len(m.flex_bending) + # Populate lookup maps and candidate pairs + flexelem_geom_pairs = [] + flexshell_geom_pairs = [] + flexvert_geom_pairs = [] + + flex_elemflexid = np.zeros(mjm.nflexelem, dtype=np.int32) + flex_shellflexid = np.zeros(mjm.nflexshelldata, dtype=np.int32) + flex_evpairflexid = np.zeros(mjm.nflexevpair, dtype=np.int32) + flex_vertflexid = np.zeros(mjm.nflexvert, dtype=np.int32) + flex_shelladr = np.zeros(mjm.nflex, dtype=np.int32) + + if mjm.nflex > 0: + shell_offset = 0 + for fi in range(mjm.nflex): + fct = mjm.flex_contype[fi] + fca = mjm.flex_conaffinity[fi] + fdim = mjm.flex_dim[fi] + + # Mappings loop + elem_start = mjm.flex_elemadr[fi] + elem_num = mjm.flex_elemnum[fi] + flex_elemflexid[elem_start : elem_start + elem_num] = fi + + ev_start = mjm.flex_evpairadr[fi] + ev_num = mjm.flex_evpairnum[fi] + flex_evpairflexid[ev_start : ev_start + ev_num] = fi + + flex_shelladr[fi] = shell_offset + shell_num = mjm.flex_shellnum[fi] + flex_shellflexid[shell_offset : shell_offset + shell_num] = fi + shell_offset += shell_num + + vert_start = mjm.flex_vertadr[fi] + vert_num = mjm.flex_vertnum[fi] + flex_vertflexid[vert_start : vert_start + vert_num] = fi + + # Candidate pairs loop + match = ((mjm.geom_contype & fca) != 0) | ((fct & mjm.geom_conaffinity) != 0) + is_prim = np.isin( + mjm.geom_type, + [ + mujoco.mjtGeom.mjGEOM_SPHERE, + mujoco.mjtGeom.mjGEOM_CAPSULE, + mujoco.mjtGeom.mjGEOM_BOX, + mujoco.mjtGeom.mjGEOM_CYLINDER, + mujoco.mjtGeom.mjGEOM_MESH, + ], + ) + is_pl = mjm.geom_type == mujoco.mjtGeom.mjGEOM_PLANE + + matching_primitive_geoms = np.where(match & is_prim)[0] + matching_plane_geoms = np.where(match & is_pl)[0] + + vert_start = mjm.flex_vertadr[fi] + + if fdim == 2: + elemdata_start = mjm.flex_elemdataadr[fi] + for e in range(elem_num): + elemid = elem_start + e + v0 = vert_start + mjm.flex_elem[elemdata_start + e * 3] + v1 = vert_start + mjm.flex_elem[elemdata_start + e * 3 + 1] + v2 = vert_start + mjm.flex_elem[elemdata_start + e * 3 + 2] + + if len(matching_primitive_geoms) > 0: + filtered = _filter_tri_geoms(mjm, v0, v1, v2, matching_primitive_geoms, filterparent) + for g in matching_primitive_geoms[~filtered]: + flexelem_geom_pairs.append((elemid, g)) + + elif fdim == 3: + shelldata_start = mjm.flex_shelldataadr[fi] + prev_shells_offset = shell_offset - shell_num + for s in range(shell_num): + v0 = vert_start + mjm.flex_shell[shelldata_start + s * 3] + v1 = vert_start + mjm.flex_shell[shelldata_start + s * 3 + 1] + v2 = vert_start + mjm.flex_shell[shelldata_start + s * 3 + 2] + + if len(matching_primitive_geoms) > 0: + filtered = _filter_tri_geoms(mjm, v0, v1, v2, matching_primitive_geoms, filterparent) + for g in matching_primitive_geoms[~filtered]: + flexshell_geom_pairs.append((prev_shells_offset + s, g)) + + # Planes vs Vertices + if len(matching_plane_geoms) > 0: + vert_count = mjm.flex_vertnum[fi] + for v in range(vert_count): + vertid = vert_start + v + bv = mjm.flex_vertbodyid[vertid] + wv = mjm.body_weldid[bv] + + bg = mjm.geom_bodyid[matching_plane_geoms] + wg = mjm.body_weldid[bg] + + mask = wg != wv + + if filterparent: + wpv = mjm.body_weldid[mjm.body_parentid[wv]] + wpg = mjm.body_weldid[mjm.body_parentid[wg]] + mask &= ~((wg != 0) & (wv != 0) & ((wg == wpv) | (wv == wpg))) + + sig = (bv << 16) + matching_plane_geoms + mask &= ~np.isin(sig, mjm.exclude_signature) + + for g in matching_plane_geoms[mask]: + flexvert_geom_pairs.append((vertid, g)) + + if not flexelem_geom_pairs: + flexelem_geom_pairs = np.zeros((0, 2), dtype=np.int32) + if not flexshell_geom_pairs: + flexshell_geom_pairs = np.zeros((0, 2), dtype=np.int32) + if not flexvert_geom_pairs: + flexvert_geom_pairs = np.zeros((0, 2), dtype=np.int32) + + m.flexelem_geom_pair_filtered = np.array(flexelem_geom_pairs, dtype=np.int32) + m.flexshell_geom_pair_filtered = np.array(flexshell_geom_pairs, dtype=np.int32) + m.flexvert_geom_pair_filtered = np.array(flexvert_geom_pairs, dtype=np.int32) + + m.flex_elemflexid = flex_elemflexid + m.flex_shellflexid = flex_shellflexid + m.flex_evpairflexid = flex_evpairflexid + m.flex_vertflexid = flex_vertflexid + m.flex_shelladr = flex_shelladr # place m on device - sizes = dict({"*": 1}, **{f.name: getattr(m, f.name) for f in dataclasses.fields(types.Model) if f.type is int}) + sizes = {f.name: getattr(m, f.name) for f in dataclasses.fields(types.Model) if f.type is int} for f in dataclasses.fields(types.Model): if _is_array_spec(f.type): - setattr(m, f.name, _create_array(getattr(m, f.name), f.type, sizes)) + batch_size = batch_sizes.get(f.name, 1) + setattr(m, f.name, _create_array(getattr(m, f.name), f.type, sizes, batch_size)) + _mark_batched(m) return m @@ -812,6 +1097,12 @@ def _get_padded_sizes(nv: int, njmax: int, is_sparse: bool, tile_size: int): return njmax_padded, nv_padded +def _nvmax_pad(nvmax: int) -> int: + """Round nvmax up to the dense tile size so the blocked Cholesky never overruns its tile.""" + t = types.TILE_SIZE_JTDAJ_DENSE + return ((max(nvmax, 1) + t - 1) // t) * t + + def _default_nconmax(mjm: mujoco.MjModel, mjd: Optional[mujoco.MjData] = None) -> int: """Returns a default guess for an ideal nconmax given a Model and optional Data. @@ -1012,16 +1303,16 @@ def _allocate_island_arrays( d: types.Data, nworld: int, njmax: int, - island_enabled: bool, + enabled: bool, mjd: mujoco.MjData, ): - ntree_size = mjm.ntree if island_enabled else 0 - nv_size = mjm.nv if island_enabled else 0 - njmax_size = njmax if island_enabled else 0 + ntree_size = mjm.ntree if enabled else 0 + nv_size = mjm.nv if enabled else 0 + njmax_size = njmax if enabled else 0 d.nisland = wp.array(np.full(nworld, mjd.nisland), dtype=int) - d.tree_island = wp.array(np.tile(mjd.tree_island, (nworld, 1 if island_enabled else 0)), dtype=int) - d.dof_island = wp.array(np.tile(mjd.dof_island, (nworld, 1 if island_enabled else 0)), dtype=int) + d.tree_island = wp.array(np.tile(mjd.tree_island, (nworld, 1 if enabled else 0)), dtype=int) + d.dof_island = wp.array(np.tile(mjd.dof_island, (nworld, 1 if enabled else 0)), dtype=int) d.island_dofadr = wp.empty((nworld, ntree_size), dtype=int) d.island_idofadr = wp.empty((nworld, ntree_size), dtype=int) @@ -1029,8 +1320,8 @@ def _allocate_island_arrays( d.island_nefc = wp.empty((nworld, ntree_size), dtype=int) d.island_ne = wp.empty((nworld, ntree_size), dtype=int) d.island_nf = wp.empty((nworld, ntree_size), dtype=int) - d.island_efcadr = wp.empty((nworld, ntree_size), dtype=int) - d.nidof = wp.empty((nworld if island_enabled else 0,), dtype=int) + d.island_iefcadr = wp.empty((nworld, ntree_size), dtype=int) + d.nidof = wp.empty((nworld if enabled else 0,), dtype=int) d.map_dof2idof = wp.empty((nworld, nv_size), dtype=int) d.map_idof2dof = wp.empty((nworld, nv_size), dtype=int) d.map_efc2iefc = wp.empty((nworld, njmax_size), dtype=int) @@ -1038,10 +1329,79 @@ def _allocate_island_arrays( d.dof_islandid = wp.empty((nworld, nv_size), dtype=int) d.efc_islandid = wp.empty((nworld, njmax_size), dtype=int) - d.iqacc = wp.empty((nworld, nv_size), dtype=float) - d.iqacc_smooth = wp.empty((nworld, nv_size), dtype=float) - d.iqfrc_smooth = wp.empty((nworld, nv_size), dtype=float) - d.iqfrc_constraint = wp.empty((nworld, nv_size), dtype=float) + + +def _allocate_compact_arrays( + mjm: mujoco.MjModel, + d: types.Data, + nworld: int, + nvmax_pad: int, + njmax_pad: int, + compact: bool, +): + """Allocate workspace for the compacted dense factor/solve (when nvmax is requested). + + Mirrors the island-local ``i*`` Data fields with a ``c*`` (compact) prefix. The + constant model-shadows (tolerances, dof-pair indices) are derived on the host since + they depend only on nvmax_pad; the workspace shadows are sized by nvmax_pad so the + blocked Cholesky never reads out of bounds on its partial tile. When the user does not + request compaction (nvmax is None) everything is allocated empty. + + TODO(team): once the compact path replaces the island solver, the whole forward + pipeline can run in compacted space and ``d.M`` / ``d.qacc`` etc. become nvmax-sized + directly, collapsing these ``c*`` shadows into the primary Data fields. + """ + nw = nworld if compact else 0 + nvp = nvmax_pad if compact else 0 + njp = njmax_pad if compact else 0 + + if compact: + # match the float32 tolerance clamp applied in put_model; rescale by nv/nvmax_pad so + # the solver's nv-normalized convergence test matches the full-model baseline. + scale = float(mjm.nv) / float(nvmax_pad) + tol = max(float(mjm.opt.tolerance), 1e-6) + ls_tol = float(mjm.opt.ls_tolerance) + d.ctol = wp.array([tol * scale], dtype=float) + d.cls_tol = wp.array([ls_tol * scale], dtype=float) + # all (i, j) DOF pairs of the nvmax_pad-wide compacted Hessian (the global dof_tri, + # triu over full nv, would index out of bounds). + idx = np.arange(nvmax_pad, dtype=np.int32) + d.cdof_tri_row = wp.array(np.repeat(idx, nvmax_pad), dtype=int) + d.cdof_tri_col = wp.array(np.tile(idx, nvmax_pad), dtype=int) + else: + d.ctol = wp.empty(0, dtype=float) + d.cls_tol = wp.empty(0, dtype=float) + d.cdof_tri_row = wp.empty(0, dtype=int) + d.cdof_tri_col = wp.empty(0, dtype=int) + + d.cM = wp.empty((nw, nvp, nvp), dtype=float) + d.cqLD = wp.empty((nw, nvp, nvp), dtype=float) + d.crhs = wp.empty((nw, nvp, 1), dtype=float) + d.cx = wp.empty((nw, nvp, 1), dtype=float) + d.cJ = wp.empty((nw, njp, nvp), dtype=float) + d.cMa = wp.empty((nw, nvp), dtype=float) + d.cqfrc_smooth = wp.empty((nw, nvp), dtype=float) + d.cqacc_smooth = wp.empty((nw, nvp), dtype=float) + d.cqacc_warmstart = wp.empty((nw, nvp), dtype=float) + d.cqacc = wp.empty((nw, nvp), dtype=float) + d.cqfrc_constraint = wp.empty((nw, nvp), dtype=float) + + +def _initial_body_awake(mjm: mujoco.MjModel, nworld: int, init_asleep: bool) -> np.ndarray: + """Returns the initial body awake array.""" + body_awake_np = np.zeros((nworld, mjm.nbody), dtype=np.int32) + for b in range(mjm.nbody): + tree = mjm.body_treeid[b] + if tree < 0: + root = mjm.body_rootid[b] + mocap = mjm.body_mocapid[root] + if mocap >= 0: + body_awake_np[:, b] = int(types.SleepState.AWAKE) + else: + body_awake_np[:, b] = int(types.SleepState.STATIC) + else: + body_awake_np[:, b] = int(types.SleepState.ASLEEP) if init_asleep else int(types.SleepState.AWAKE) + return body_awake_np def make_data( @@ -1053,6 +1413,7 @@ def make_data( njmax_nnz: Optional[int] = None, naconmax: Optional[int] = None, naccdmax: Optional[int] = None, + nvmax: Optional[int] = None, ) -> types.Data: """Creates a data object on device. @@ -1067,6 +1428,7 @@ def make_data( njmax_nnz: Number of non-zeros in constraint Jacobian (sparse). Defaults to njmax * nv. naconmax: Number of contacts to allocate for all worlds. Overrides nconmax. naccdmax: Maximum number of CCD contacts. Defaults to naconmax. + nvmax: Capacity for compacted active DOFs per world. Defaults to nv. Returns: The data object containing the current state and output arrays (device). @@ -1078,12 +1440,24 @@ def make_data( if njmax is None: njmax = _default_njmax(mjm) + island_alloc = True + sleep_enabled = bool(mjm.opt.enableflags & mujoco.mjtEnableBit.mjENBL_SLEEP) and not bool( + mjm.opt.disableflags & mujoco.mjtDisableBit.mjDSBL_ISLAND + ) + compact_alloc = sleep_enabled or (nvmax is not None) + + if nvmax is None: + nvmax = mjm.nv + if nconmax < 0: raise ValueError("nconmax must be >= 0") if njmax < 0: raise ValueError("njmax must be >= 0") + if nvmax < 0 or nvmax > mjm.nv: + raise ValueError(f"nvmax ({nvmax}) must be in [0, nv ({mjm.nv})]") + if nworld < 1: raise ValueError(f"nworld must be >= 1") @@ -1105,6 +1479,8 @@ def make_data( elif nccdmax > nconmax: raise ValueError(f"nccdmax ({nccdmax}) must be <= nconmax ({nconmax})") + nv_compact = nvmax < mjm.nv + sizes = dict({"*": 1}, **{f.name: getattr(mjm, f.name, None) for f in dataclasses.fields(types.Model) if f.type is int}) condim_arrays = [np.array([0]), mjm.geom_condim, mjm.pair_dim] if mjm.nflex > 0: @@ -1116,6 +1492,8 @@ def make_data( sizes["nworld"] = nworld sizes["naconmax"] = naconmax sizes["njmax"] = njmax + sizes["nvmax"] = nvmax + sizes["nvmax_pad"] = _nvmax_pad(nvmax) if njmax_nnz is None: if is_sparse(mjm): @@ -1123,10 +1501,16 @@ def make_data( else: njmax_nnz = njmax * mjm.nv - contact = types.Contact(**{f.name: _create_array(None, f.type, sizes) for f in dataclasses.fields(types.Contact)}) + contact_kwargs = {} + for f in dataclasses.fields(types.Contact): + if f.name in ["flex", "elem", "vert"] and mjm.nflex == 0: + contact_kwargs[f.name] = wp.empty(0, dtype=wp.vec2i) + else: + contact_kwargs[f.name] = _create_array(None, f.type, sizes) + contact = types.Contact(**contact_kwargs) contact.efc_address = wp.array(np.full((naconmax, sizes["nmaxpyramid"]), -1, dtype=int), dtype=int) - efc = _create_constraint(mjm, nworld, njmax, njmax_nnz, sizes, ENABLE_ISLANDS) + efc = _create_constraint(mjm, nworld, njmax, njmax_nnz, sizes) if is_sparse(mjm): efc.J_rownnz = wp.zeros((nworld, njmax), dtype=int) @@ -1139,14 +1523,9 @@ def make_data( efc.J_colind = wp.zeros((nworld, 0, 0), dtype=int) efc.J = wp.zeros((nworld, sizes["njmax_pad"], sizes["nv_pad"]), dtype=float) - contact_kwargs = {} - for f in dataclasses.fields(types.Contact): - contact_kwargs[f.name] = _create_array(None, f.type, sizes) - contact = types.Contact(**contact_kwargs) - - # world body and static geom (attached to the world) poses are precomputed - # this speeds up scenes with many static geoms (e.g. terrains) - # TODO(team): remove this when we introduce dof islands + sleeping + # Compute initial kinematic state. Static geom positions (geom_xpos, geom_xmat) are set here + # and never updated by the physics loop (see smooth.py geom_kinematics), so this call is the + # only place they are initialized. Also seeds body poses (xquat, xmat, ximat) at qpos0. mjd = mujoco.MjData(mjm) mujoco.mj_kinematics(mjm, mjd) @@ -1162,6 +1541,8 @@ def make_data( "naconmax": naconmax, "naccdmax": naccdmax, "njmax": njmax, + "nvmax": nvmax, + "nvmax_pad": sizes["nvmax_pad"], "njmax_pad": sizes["njmax_pad"], "njmax_nnz": njmax_nnz, "M": None, @@ -1190,7 +1571,7 @@ def make_data( "island_nefc": None, "island_ne": None, "island_nf": None, - "island_efcadr": None, + "island_iefcadr": None, "nidof": None, "map_dof2idof": None, "map_idof2dof": None, @@ -1198,14 +1579,9 @@ def make_data( "map_iefc2efc": None, "dof_islandid": None, "efc_islandid": None, - "iqacc": None, - "iqacc_smooth": None, - "iqfrc_smooth": None, - "iqfrc_constraint": None, - # sleep state: all trees start fully awake - "tree_asleep": wp.array(np.full((nworld, mjm.ntree), -(1 + types.MJ_MINAWAKE)), dtype=int), - "tree_awake": wp.array(np.ones((nworld, mjm.ntree)), dtype=int), - "body_awake": wp.array(np.ones((nworld, mjm.nbody)), dtype=int), + "tree_asleep": wp.array(np.full((nworld, mjm.ntree), -(1 + types.MJ_MINAWAKE), dtype=np.int32), dtype=int), + "tree_awake": wp.array(np.ones((nworld, mjm.ntree), dtype=np.int32), dtype=int), + "body_awake": wp.array(_initial_body_awake(mjm, nworld, False), dtype=int), } for f in dataclasses.fields(types.Data): if f.name in d_kwargs: @@ -1214,15 +1590,21 @@ def make_data( d = types.Data(**d_kwargs) - if is_sparse(mjm): - d.M = wp.zeros((nworld, 1, mjm.nC), dtype=float) - d.qLD = wp.zeros((nworld, 1, mjm.nC), dtype=float) - else: - d.M = wp.zeros((nworld, sizes["nv_pad"], sizes["nv_pad"]), dtype=float) - d.qLD = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float) + # qLD holds the factor: a packed dense region for dense blocks followed + # by an nC-length LDL region for sparse blocks (present only when some block is sparse). Either + # region may be empty (pure dense / pure sparse). + d.M = wp.zeros((nworld, mjm.nC), dtype=float) + _lay = m_block_layout(mjm) + qld_total = _lay["total"] + (mjm.nC if _lay["has_sparse"] else 0) + d.qLD = wp.zeros((nworld, qld_total), dtype=float) - _allocate_island_arrays(mjm, d, nworld, njmax, ENABLE_ISLANDS, mjd) + _allocate_island_arrays(mjm, d, nworld, njmax, island_alloc, mjd) + _allocate_compact_arrays(mjm, d, nworld, sizes["nvmax_pad"], sizes["njmax_pad"], compact_alloc) + d.ncdof.zero_() + d.dof_cdof.fill_(-1) + d.cdof_dof.fill_(-1) + _mark_batched(d) return d @@ -1236,6 +1618,7 @@ def put_data( njmax_nnz: Optional[int] = None, naconmax: Optional[int] = None, naccdmax: Optional[int] = None, + nvmax: Optional[int] = None, ) -> types.Data: """Moves data from host to a device. @@ -1251,6 +1634,7 @@ def put_data( njmax_nnz: Number of non-zeros in constraint Jacobian (sparse). Defaults to njmax * nv. naconmax: Number of contacts to allocate for all worlds. Overrides nconmax. naccdmax: Maximum number of CCD contacts. Defaults to naconmax. + nvmax: Capacity for compacted active DOFs per world. Defaults to nv. Returns: The data object containing the current state and output arrays (device). @@ -1265,12 +1649,24 @@ def put_data( if njmax is None: njmax = _default_njmax(mjm, mjd) + island_alloc = True + sleep_enabled = bool(mjm.opt.enableflags & mujoco.mjtEnableBit.mjENBL_SLEEP) and not bool( + mjm.opt.disableflags & mujoco.mjtDisableBit.mjDSBL_ISLAND + ) + compact_alloc = sleep_enabled or (nvmax is not None) + + if nvmax is None: + nvmax = mjm.nv + if nconmax < 0: raise ValueError("nconmax must be >= 0") if njmax < 0: raise ValueError("njmax must be >= 0") + if nvmax < 0 or nvmax > mjm.nv: + raise ValueError(f"nvmax ({nvmax}) must be in [0, nv ({mjm.nv})]") + if nworld < 1: raise ValueError(f"nworld must be >= 1") @@ -1302,6 +1698,8 @@ def put_data( if mjd.nefc > njmax: raise ValueError(f"njmax overflow (njmax must be >= {mjd.nefc})") + nv_compact = nvmax < mjm.nv + sizes = dict({"*": 1}, **{f.name: getattr(mjm, f.name, None) for f in dataclasses.fields(types.Model) if f.type is int}) condim_arrays = [np.array([0]), mjm.geom_condim, mjm.pair_dim] if mjm.nflex > 0: @@ -1313,6 +1711,8 @@ def put_data( sizes["nworld"] = nworld sizes["naconmax"] = naconmax sizes["njmax"] = njmax + sizes["nvmax"] = nvmax + sizes["nvmax_pad"] = _nvmax_pad(nvmax) if njmax_nnz is None: if is_sparse(mjm): @@ -1320,8 +1720,13 @@ def put_data( else: njmax_nnz = njmax * mjm.nv - # ensure static geom positions are computed - # TODO: remove once MjData creation semantics are fixed + # Capture sleep state before mj_kinematics, which resets tree_asleep as a side effect. + tree_asleep_init = mjd.tree_asleep.copy() + body_awake_init = mjd.body_awake.copy() + + # Ensure kinematic state is populated. mujoco.MjData() does not call mj_kinematics, so a freshly + # created mjd has zero geom positions. Static geoms are never updated by the physics loop + # (see smooth.py geom_kinematics), so without this call they would remain at (0,0,0). mujoco.mj_kinematics(mjm, mjd) # create contact @@ -1329,8 +1734,11 @@ def put_data( for f in dataclasses.fields(types.Contact): if f.name in contact_kwargs: continue + if f.name in ["flex", "elem", "vert"] and mjm.nflex == 0: + contact_kwargs[f.name] = wp.empty(0, dtype=wp.vec2i) + continue val = getattr(mjd.contact, f.name) - val = np.repeat(val, nworld, axis=0) + val = np.tile(val, (nworld,) + (1,) * (val.ndim - 1)) width = ((0, naconmax - val.shape[0]),) + ((0, 0),) * (val.ndim - 1) val = np.pad(val, width) contact_kwargs[f.name] = _create_array(val, f.type, sizes) @@ -1356,7 +1764,20 @@ def put_data( # create efc efc_kwargs = {"J_rownnz": None, "J_rowadr": None, "J_colind": None, "J": None} - efc = _create_constraint(mjm, nworld, njmax, njmax_nnz, sizes, ENABLE_ISLANDS, mjd) + efc = _create_constraint(mjm, nworld, njmax, njmax_nnz, sizes, mjd) + + # make_constraint builds the block list in-kernel; put_data does not run it, so build it here + # -- otherwise solving a put_data state would assemble an empty J^T D J. + if is_sparse(mjm) and mjm.opt.solver == mujoco.mjtSolver.mjSOL_NEWTON: + jtdaj_adr, jtdaj_nrow = _jtdaj_groups(mjd) + nblock = jtdaj_adr.shape[0] + adr_row = np.zeros(njmax, dtype=int) + nrow_row = np.zeros(njmax, dtype=int) + adr_row[:nblock] = jtdaj_adr + nrow_row[:nblock] = jtdaj_nrow + efc.jtdaj_adr = wp.array(np.tile(adr_row, (nworld, 1)), dtype=int) + efc.jtdaj_nrow = wp.array(np.tile(nrow_row, (nworld, 1)), dtype=int) + efc.jtdaj_nblock = wp.array(np.full(nworld, nblock, dtype=int), dtype=int) if is_sparse(mjm): J_rownnz = np.zeros(njmax, dtype=np.int32) @@ -1403,6 +1824,8 @@ def put_data( "naconmax": naconmax, "naccdmax": naccdmax, "njmax": njmax, + "nvmax": nvmax, + "nvmax_pad": sizes["nvmax_pad"], "njmax_pad": sizes["njmax_pad"], "njmax_nnz": njmax_nnz, # fields set after initialization: @@ -1420,7 +1843,7 @@ def put_data( "island_nefc": None, "island_ne": None, "island_nf": None, - "island_efcadr": None, + "island_iefcadr": None, "nidof": None, "map_dof2idof": None, "map_idof2dof": None, @@ -1428,46 +1851,47 @@ def put_data( "map_iefc2efc": None, "dof_islandid": None, "efc_islandid": None, - "iqacc": None, - "iqacc_smooth": None, - "iqfrc_smooth": None, - "iqfrc_constraint": None, + "tree_asleep": wp.array(np.tile(tree_asleep_init, (nworld, 1)), dtype=int), + "tree_awake": wp.array(np.tile((tree_asleep_init < 0).astype(np.int32), (nworld, 1)), dtype=int), + "body_awake": wp.array(np.tile(body_awake_init.astype(np.int32), (nworld, 1)), dtype=int), } for f in dataclasses.fields(types.Data): if f.name in d_kwargs: continue val = getattr(mjd, f.name, None) - if val is not None: - shape = val.shape if hasattr(val, "shape") else () - val = np.full((nworld,) + shape, val) d_kwargs[f.name] = _create_array(val, f.type, sizes) d = types.Data(**d_kwargs) d.solver_niter = wp.full((nworld,), mjd.solver_niter[0], dtype=int) - if is_sparse(mjm): - if check_version("mujoco>=3.8.1.dev910242375"): - d.M = wp.array(np.full((nworld, 1, mjm.nC), mjd.M), dtype=float) - else: - d.M = wp.array(np.full((nworld, 1, mjm.nC), mjd.qM[mjm.mapM2M]), dtype=float) - d.qLD = wp.array(np.full((nworld, 1, mjm.nC), mjd.qLD), dtype=float) - else: - M = np.zeros((mjm.nv, mjm.nv)) - if check_version("mujoco>=3.8.1.dev910242375"): - mujoco.mju_sym2dense(M, mjd.M, mjm.M_rownnz, mjm.M_rowadr, mjm.M_colind) - qLD = np.linalg.cholesky(M).T if (mjd.M != 0.0).any() and (mjd.qLD != 0.0).any() else np.zeros((mjm.nv, mjm.nv)) - else: - mujoco.mj_fullM(mjm, M, mjd.qM) - qLD = np.linalg.cholesky(M).T if (mjd.qM != 0.0).any() and (mjd.qLD != 0.0).any() else np.zeros((mjm.nv, mjm.nv)) - padding = sizes["nv_pad"] - mjm.nv - M_padded = np.pad(M, ((0, padding), (0, padding)), mode="constant", constant_values=0.0) - d.M = wp.array(np.full((nworld, sizes["nv_pad"], sizes["nv_pad"]), M_padded), dtype=float) - d.qLD = wp.array(np.full((nworld, mjm.nv, mjm.nv), qLD), dtype=float) + d.M = wp.array(np.full((nworld, mjm.nC), mjd.M), dtype=float) + # qLD = [packed dense-block Cholesky | nC LDL region]. Dense blocks store their upper Cholesky + # packed; the LDL region (present iff some block is sparse) holds MuJoCo's full L'DL factor (only + # its sparse-block entries are read by the solve). + lay = m_block_layout(mjm) + qld_total = lay["total"] + (mjm.nC if lay["has_sparse"] else 0) + qLD = np.zeros(qld_total, dtype=np.float32) + if lay["has_dense"]: + Mfull = np.zeros((mjm.nv, mjm.nv)) + mujoco.mju_sym2dense(Mfull, mjd.M, mjm.M_rownnz, mjm.M_rowadr, mjm.M_colind) + for start, size in lay["dense_blocks"]: + off = lay["dof_adr"][start] + blk = Mfull[start : start + size, start : start + size] + if blk.any(): + qLD[off : off + size * size] = np.linalg.cholesky(blk).T.reshape(-1) + if lay["has_sparse"]: + qLD[lay["total"] :] = mjd.qLD + d.qLD = wp.array(np.full((nworld, qld_total), qLD), dtype=float) - _allocate_island_arrays(mjm, d, nworld, njmax, ENABLE_ISLANDS, mjd) + _allocate_island_arrays(mjm, d, nworld, njmax, island_alloc, mjd) + _allocate_compact_arrays(mjm, d, nworld, sizes["nvmax_pad"], sizes["njmax_pad"], compact_alloc) + d.ncdof.zero_() + d.dof_cdof.fill_(-1) + d.cdof_dof.fill_(-1) d.nacon = wp.array([mjd.ncon * nworld], dtype=int) + _mark_batched(d) return d @@ -1601,6 +2025,9 @@ def get_data_into( if mjm.nhistory > 0: result.history[:] = d.history.numpy()[world_id] + if mjm.nuserdata > 0: + result.userdata[:] = d.userdata.numpy()[world_id] + # contact result.contact.dist[:ncon] = d.contact.dist.numpy()[ncon_filter] result.contact.pos[:ncon] = d.contact.pos.numpy()[ncon_filter] @@ -1612,31 +2039,21 @@ def get_data_into( result.contact.solimp[:ncon] = d.contact.solimp.numpy()[ncon_filter] result.contact.dim[:ncon] = d.contact.dim.numpy()[ncon_filter] result.contact.geom[:ncon] = d.contact.geom.numpy()[ncon_filter] + if mjm.nflex > 0: + result.contact.flex[:ncon] = d.contact.flex.numpy()[ncon_filter] + result.contact.elem[:ncon] = d.contact.elem.numpy()[ncon_filter] + result.contact.vert[:ncon] = d.contact.vert.numpy()[ncon_filter] result.contact.efc_address[:ncon] = contact_efc_address_ordered[:ncon] - if is_sparse(mjm): - if check_version("mujoco>=3.8.1.dev910242375"): - result.M[:] = d.M.numpy()[world_id, 0] - else: - result.qM[mjm.mapM2M] = d.M.numpy()[world_id, 0] - result.qLD[:] = d.qLD.numpy()[world_id, 0] - else: - M = d.M.numpy()[world_id] - if check_version("mujoco>=3.8.1.dev910242375"): - for i in range(mjm.nv): - adr = mjm.M_rowadr[i] - for k in range(mjm.M_rownnz[i]): - col = mjm.M_colind[adr + k] - result.M[adr + k] = M[i, col] - else: - adr = 0 - for i in range(mjm.nv): - j = i - while j >= 0: - result.qM[adr] = M[i, j] - j = mjm.dof_parentid[j] - adr += 1 + result.M[:] = d.M.numpy()[world_id] + _lay = m_block_layout(mjm) + if _lay["has_dense"] or _lay["has_simple"]: + # d.qLD is not MuJoCo's LDL: dense blocks are a packed Cholesky and simple blocks are factored + # into qLDiagInv (their LDL slots are never written). Recompute the LDL factor from M. mujoco.mj_factorM(mjm, result) + else: + # Pure sparse: qLD is exactly MuJoCo's nC LDL factor. + result.qLD[:] = d.qLD.numpy()[world_id] if nefc > 0: if is_sparse(mjm): @@ -1710,44 +2127,13 @@ def get_data_into( result.island_nefc[:nisland] = d.island_nefc.numpy()[world_id, :nisland] result.island_ne[:nisland] = d.island_ne.numpy()[world_id, :nisland] result.island_nf[:nisland] = d.island_nf.numpy()[world_id, :nisland] - result.island_iefcadr[:nisland] = d.island_efcadr.numpy()[world_id, :nisland] + result.island_iefcadr[:nisland] = d.island_iefcadr.numpy()[world_id, :nisland] nv = mjm.nv result.map_dof2idof[:nv] = d.map_dof2idof.numpy()[world_id, :nv] result.map_idof2dof[:nv] = d.map_idof2dof.numpy()[world_id, :nv] result.map_efc2iefc[:nefc] = d.map_efc2iefc.numpy()[world_id, :nefc] result.map_iefc2efc[:nefc] = d.map_iefc2efc.numpy()[world_id, :nefc] - result.iefc_type[:nefc] = d.efc.itype.numpy()[world_id, :nefc] - result.iefc_id[:nefc] = d.efc.iid.numpy()[world_id, :nefc] - result.iefc_D[:nefc] = d.efc.iD.numpy()[world_id, :nefc] - result.iefc_aref[:nefc] = d.efc.iaref.numpy()[world_id, :nefc] - result.iefc_frictionloss[:nefc] = d.efc.ifrictionloss.numpy()[world_id, :nefc] - result.iefc_state[:nefc] = d.efc.istate.numpy()[world_id, :nefc] - result.iefc_force[:nefc] = d.efc.iforce.numpy()[world_id, :nefc] - - if is_sparse(mjm): - iefc_J = np.zeros((nefc, mjm.nv)) - mujoco.mju_sparse2dense( - iefc_J, - d.efc.iJ.numpy()[world_id, 0], - d.efc.iJ_rownnz.numpy()[world_id, :nefc], - d.efc.iJ_rowadr.numpy()[world_id, :nefc], - d.efc.iJ_colind.numpy()[world_id, 0], - ) - else: - iefc_J = d.efc.iJ.numpy()[world_id, :nefc, : mjm.nv] - - if mujoco.mj_isSparse(mjm): - mujoco.mju_dense2sparse( - result.iefc_J, - iefc_J, - result.iefc_J_rownnz, - result.iefc_J_rowadr, - result.iefc_J_colind, - ) - else: - result.iefc_J[: nefc * mjm.nv] = iefc_J.flatten() - def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): """Clear data, set defaults; optionally by world. @@ -1757,6 +2143,7 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): d: The data object containing the current state and output arrays (device). reset: Per-world bitmask. Reset if True. """ + sleep_enabled = bool(m.opt.enableflags & types.EnableBit.SLEEP) @wp.kernel(module="unique", enable_backward=False) def reset_xfrc_applied(reset_in: wp.array[bool], xfrc_applied_out: wp.array2d[wp.spatial_vector]): @@ -1769,14 +2156,14 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): xfrc_applied_out[worldid, bodyid][elemid] = 0.0 @wp.kernel(module="unique", enable_backward=False) - def reset_M(reset_in: wp.array[bool], M_out: wp.array3d[float]): - worldid, elemid1, elemid2 = wp.tid() + def reset_M(reset_in: wp.array[bool], M_out: wp.array2d[float]): + worldid, elemid = wp.tid() if wp.static(reset is not None): if not reset_in[worldid]: return - M_out[worldid, elemid1, elemid2] = 0.0 + M_out[worldid, elemid] = 0.0 @wp.kernel(module="unique", enable_backward=False) def reset_nworld( @@ -1788,6 +2175,7 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): nbody: int, ntree: int, neq: int, + nuserdata: int, nsensordata: int, qpos0: wp.array2d[float], eq_active0: wp.array[bool], @@ -1815,8 +2203,10 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): eq_active_out: wp.array2d[bool], qacc_out: wp.array2d[float], act_dot_out: wp.array2d[float], + userdata_out: wp.array2d[float], sensordata_out: wp.array2d[float], nacon_out: wp.array[int], + overflow_out: wp.array[int], ): worldid = wp.tid() @@ -1853,6 +2243,9 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): eq_active_out[worldid, i] = eq_active0[i] for i in range(nsensordata): sensordata_out[worldid, i] = 0.0 + for i in range(nuserdata): + userdata_out[worldid, i] = 0.0 + overflow_out[worldid] = 0 @wp.kernel(module="unique", enable_backward=False) def reset_mocap( @@ -1897,6 +2290,7 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): contact_dim_out: wp.array[int], contact_geom_out: wp.array[wp.vec2i], contact_flex_out: wp.array[wp.vec2i], + contact_elem_out: wp.array[wp.vec2i], contact_vert_out: wp.array[wp.vec2i], contact_efc_address_out: wp.array2d[int], contact_worldid_out: wp.array[int], @@ -1924,8 +2318,12 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): contact_solimp_out[conid] = types.vec5(0.0, 0.0, 0.0, 0.0, 0.0) contact_dim_out[conid] = 0 contact_geom_out[conid] = wp.vec2i(0, 0) - contact_flex_out[conid] = wp.vec2i(0, 0) - contact_vert_out[conid] = wp.vec2i(0, 0) + if contact_flex_out.shape[0] > 0: + contact_flex_out[conid] = wp.vec2i(0, 0) + if contact_elem_out.shape[0] > 0: + contact_elem_out[conid] = wp.vec2i(0, 0) + if contact_vert_out.shape[0] > 0: + contact_vert_out[conid] = wp.vec2i(0, 0) for i in range(nefcaddress): contact_efc_address_out[conid, i] = -1 contact_worldid_out[conid] = 0 @@ -1978,7 +2376,7 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): wp.launch(reset_xfrc_applied, dim=(d.nworld, m.nbody, 6), inputs=[reset_input], outputs=[d.xfrc_applied]) wp.launch( reset_M, - dim=(d.nworld, d.M.shape[1], d.M.shape[2]), + dim=(d.nworld, d.M.shape[1]), inputs=[reset_input], outputs=[d.M], ) @@ -2008,6 +2406,7 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): d.contact.dim, d.contact.geom, d.contact.flex, + d.contact.elem, d.contact.vert, d.contact.efc_address, d.contact.worldid, @@ -2032,7 +2431,21 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): wp.launch( reset_nworld, dim=d.nworld, - inputs=[m.nq, m.nv, m.nu, m.na, m.nbody, m.ntree, m.neq, m.nsensordata, m.qpos0, m.eq_active0, d.nworld, reset_input], + inputs=[ + m.nq, + m.nv, + m.nu, + m.na, + m.nbody, + m.ntree, + m.neq, + m.nuserdata, + m.nsensordata, + m.qpos0, + m.eq_active0, + d.nworld, + reset_input, + ], outputs=[ d.solver_niter, d.ne, @@ -2053,11 +2466,16 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): d.eq_active, d.qacc, d.act_dot, + d.userdata, d.sensordata, d.nacon, + d.overflow, ], ) + if sleep_enabled: + sleep.update_sleep(m, d) + # kernel_analyzer: off @wp.kernel @@ -2105,13 +2523,25 @@ def _copy_tendon_length0( tendon_length0_out[tendon_length0_id, tenid] = ten_length_in[worldid, tenid] +@wp.kernel +def _resolve_tendon_lengthspring( + ten_length_in: wp.array2d[float], + tendon_lengthspring_out: wp.array2d[wp.vec2], +): + worldid, tenid = wp.tid() + tendon_lengthspring_id = worldid % tendon_lengthspring_out.shape[0] + val = tendon_lengthspring_out[tendon_lengthspring_id, tenid] + if val[0] == -1.0 and val[1] == -1.0: + l = ten_length_in[worldid, tenid] + tendon_lengthspring_out[tendon_lengthspring_id, tenid] = wp.vec2(l, l) + + @wp.kernel def _compute_meaninertia( nv: int, - is_sparse: bool, M_rownnz_in: wp.array[int], M_rowadr_in: wp.array[int], - M_in: wp.array3d[float], + M_in: wp.array2d[float], meaninertia_out: wp.array[float], ): """Compute mean diagonal inertia from M at qpos0.""" @@ -2123,13 +2553,9 @@ def _compute_meaninertia( total = float(0.0) for i in range(nv): - if is_sparse: - # Sparse: M is in CSR format, diagonal at M_rowadr_in[i] + M_rownnz_in[i] - 1 - madr = M_rowadr_in[i] + M_rownnz_in[i] - 1 - total += M_in[worldid, 0, madr] - else: - # Dense: M is 2D matrix, diagonal at [i,i] - total += M_in[worldid, i, i] + # CSR row diagonal is the last entry: M_rowadr_in[i] + M_rownnz_in[i] - 1 + madr = M_rowadr_in[i] + M_rownnz_in[i] - 1 + total += M_in[worldid, madr] meaninertia_out[worldid % meaninertia_out.shape[0]] = total / float(nv) @@ -2559,7 +2985,6 @@ def set_const_fixed(m: types.Model, d: types.Data): Computes: - body_subtreemass: mass of body and all descendants (depends on body_mass) - - ngravcomp: count of bodies with gravity compensation (depends on body_gravcomp) Args: m: The model containing kinematic and dynamic information (device). @@ -2574,12 +2999,8 @@ def set_const_fixed(m: types.Model, d: types.Data): inputs=[m.body_parentid, m.body_subtreemass, body_tree], ) - # TODO(team): refactor for graph capture compatibility - body_gravcomp_np = m.body_gravcomp.numpy() - m.ngravcomp = int((body_gravcomp_np > 0.0).any(axis=0).sum()) - -def set_const_0(m: types.Model, d: types.Data): +def set_const_0(m: types.Model, d: types.Data, restore: bool = True): """Compute quantities that depend on qpos0. Computes: @@ -2597,6 +3018,7 @@ def set_const_0(m: types.Model, d: types.Data): Args: m: The model containing kinematic and dynamic information (device). d: The data object containing the current state and output arrays (device). + restore: Whether to restore state fields to correspond to d.qpos. """ qpos_saved = wp.clone(d.qpos) @@ -2616,7 +3038,7 @@ def set_const_0(m: types.Model, d: types.Data): wp.launch( _compute_meaninertia, dim=d.nworld, - inputs=[m.nv, m.is_sparse, m.M_rownnz, m.M_rowadr, d.M], + inputs=[m.nv, m.M_rownnz, m.M_rowadr, d.M], outputs=[m.stat.meaninertia], ) @@ -2766,8 +3188,53 @@ def set_const_0(m: types.Model, d: types.Data): wp.copy(d.qpos, qpos_saved) + if restore: + smooth.kinematics(m, d) + smooth.com_pos(m, d) + smooth.camlight(m, d) + smooth.flex(m, d) + smooth.tendon(m, d) + smooth.crb(m, d) + smooth.tendon_armature(m, d) + smooth.factor_m(m, d) + smooth.transmission(m, d) -def set_const(m: types.Model, d: types.Data): + +def set_const_spring(m: types.Model, d: types.Data, restore: bool = True): + """Compute quantities that depend on qpos_spring. + + Computes: + - tendon_lengthspring: spring resting length range + """ + if m.ntendon == 0: + return + + qpos_saved = wp.clone(d.qpos) + + wp.launch(_copy_qpos0_to_qpos, dim=(d.nworld, m.nq), inputs=[m.qpos_spring], outputs=[d.qpos]) + + smooth.kinematics(m, d) + smooth.com_pos(m, d) + smooth.tendon(m, d) + smooth.transmission(m, d) + + wp.launch( + _resolve_tendon_lengthspring, + dim=(d.nworld, m.ntendon), + inputs=[d.ten_length], + outputs=[m.tendon_lengthspring], + ) + + wp.copy(d.qpos, qpos_saved) + + if restore: + smooth.kinematics(m, d) + smooth.com_pos(m, d) + smooth.tendon(m, d) + smooth.transmission(m, d) + + +def set_const(m: types.Model, d: types.Data, restore: bool = True): """Recomputes qpos0-dependent constant model fields. This function propagates changes from some model fields to derived fields, @@ -2805,7 +3272,6 @@ def set_const(m: types.Model, d: types.Data): Computes: - Fixed quantities (via set_const_fixed): - body_subtreemass: mass of body and all descendants - - ngravcomp: count of bodies with gravity compensation - qpos0-dependent quantities (via set_const_0): - tendon_length0: tendon resting lengths - dof_invweight0: inverse inertia for DOFs @@ -2821,9 +3287,22 @@ def set_const(m: types.Model, d: types.Data): Args: m: The model containing kinematic and dynamic information (device). d: The data object containing the current state and output arrays (device). + restore: Whether to restore state fields to correspond to d.qpos. """ set_const_fixed(m, d) - set_const_0(m, d) + set_const_0(m, d, restore=False) + set_const_spring(m, d, restore=False) + + if restore: + smooth.kinematics(m, d) + smooth.com_pos(m, d) + smooth.camlight(m, d) + smooth.flex(m, d) + smooth.tendon(m, d) + smooth.crb(m, d) + smooth.tendon_armature(m, d) + smooth.factor_m(m, d) + smooth.transmission(m, d) def set_length_range(m: types.Model, d: types.Data, index: int = -1): @@ -2952,9 +3431,6 @@ def override_model(model: types.Model | mujoco.MjModel, overrides: dict[str, Any else: val = typ(val) - if attr == "disableflags" and isinstance(obj, types.Option) and not ENABLE_ISLANDS: - val = int(val) | types.DisableBit.ISLAND - setattr(obj, attr, val) @@ -3187,12 +3663,12 @@ def create_render_context( # Locate skybox texture skybox_tex_ids = np.nonzero(mjm.tex_type == mujoco.mjtTexture.mjTEXTURE_SKYBOX)[0] if mjm.ntex else np.array([], dtype=int) if render_skybox and skybox_tex_ids.size > 0: - skybox_tex_id = int(skybox_tex_ids[0]) - skybox_face_width = int(mjm.tex_width[skybox_tex_id]) + skybox_tex_id_np = np.array([skybox_tex_ids[0]], dtype=int) + skybox_face_width_np = np.array([mjm.tex_width[skybox_tex_ids[0]]], dtype=int) else: render_skybox = False - skybox_tex_id = -1 - skybox_face_width = 1 + skybox_tex_id_np = np.array([-1], dtype=int) + skybox_face_width_np = np.array([1], dtype=int) # Filter active cameras if cam_active is not None: @@ -3311,8 +3787,8 @@ def create_render_context( ), use_precomputed_rays=use_precomputed_rays, render_skybox=render_skybox, - skybox_tex_id=skybox_tex_id, - skybox_face_width=skybox_face_width, + skybox_tex_id=wp.array(skybox_tex_id_np, dtype=int), + skybox_face_width=wp.array(skybox_face_width_np, dtype=int), headlight_active=bool(mjm.vis.headlight.active), headlight_ambient=wp.vec3(mjm.vis.headlight.ambient), headlight_diffuse=wp.vec3(mjm.vis.headlight.diffuse), @@ -3368,4 +3844,5 @@ def create_render_context( bvh.build_scene_bvh(mjm, mjd, rc, nworld) + _mark_batched(rc) return rc diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/island.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/island.py index 806f71d3..6c863355 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/island.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/island.py @@ -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], ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py index a1893ff6..059756ce 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py @@ -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, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py index e33dd623..1e294952 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py @@ -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=[ diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py index 619e55cf..dfcc8f95 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py @@ -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 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sleep.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sleep.py index 724212ad..3e9521cf 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sleep.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sleep.py @@ -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], ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py index 4f35eab0..4c2939e8 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py @@ -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. diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py index a8459410..f6ae3f7a 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py @@ -13,6 +13,7 @@ # limitations under the License. # ============================================================================== +import dataclasses from math import ceil import warp as wp @@ -25,7 +26,6 @@ from mujoco.mjx.third_party.mujoco_warp._src import types from mujoco.mjx.third_party.mujoco_warp._src.block_cholesky import create_blocked_cholesky_factorize_solve_func from mujoco.mjx.third_party.mujoco_warp._src.block_cholesky import create_blocked_cholesky_solve_func from mujoco.mjx.third_party.mujoco_warp._src.types import InverseContext -from mujoco.mjx.third_party.mujoco_warp._src.types import IslandSolverContext from mujoco.mjx.third_party.mujoco_warp._src.types import SolverContext from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope @@ -57,50 +57,6 @@ def create_inverse_context(m: types.Model, d: types.Data) -> InverseContext: ) -def _create_island_solver_context(m: types.Model, d: types.Data) -> IslandSolverContext: - """Create an IslandSolverContext with allocated workspace arrays. - - Args: - m: Model. - d: Data. - - Returns: - IslandSolverContext with allocated arrays. - """ - nworld = d.nworld - nv = m.nv - nv_pad = m.nv_pad - njmax = d.njmax - ntree = m.ntree - - alloc_h = m.opt.solver == types.SolverType.NEWTON - alloc_island_cg = m.opt.solver == types.SolverType.CG - - return IslandSolverContext( - Jaref=wp.empty((nworld, njmax), dtype=float), - jv=wp.empty((nworld, njmax), dtype=float), - search=wp.empty((nworld, nv), dtype=float), - mv=wp.empty((nworld, nv), dtype=float), - grad=wp.zeros((nworld, nv_pad), dtype=float), - Mgrad=wp.zeros((nworld, nv_pad), dtype=float), - prev_grad=wp.empty((nworld, nv), dtype=float) if alloc_island_cg else wp.empty((nworld, 0), dtype=float), - prev_Mgrad=wp.empty((nworld, nv), dtype=float) if alloc_island_cg else wp.empty((nworld, 0), dtype=float), - h=wp.zeros((nworld, nv_pad, nv_pad), dtype=float) if alloc_h else wp.empty((nworld, 0, 0), dtype=float), - # Per-island solver scalars - cost=wp.empty((nworld, ntree), dtype=float), - prev_cost=wp.empty((nworld, ntree), dtype=float), - gauss=wp.empty((nworld, ntree), dtype=float), - search_dot=wp.empty((nworld, ntree), dtype=float), - grad_dot=wp.empty((nworld, ntree), dtype=float), - done=wp.empty((nworld, ntree), dtype=bool), - solver_niter=wp.empty((nworld, ntree), dtype=int), - beta=wp.empty((nworld, ntree), dtype=float) if alloc_island_cg else wp.empty((nworld, 0), dtype=float), - beta_den=wp.empty((nworld, ntree), dtype=float) if alloc_island_cg else wp.empty((nworld, 0), dtype=float), - alpha=wp.empty((nworld, ntree), dtype=float), - Ma=wp.empty((nworld, nv), dtype=float), - ) - - def _create_solver_context(m: types.Model, d: types.Data) -> SolverContext: """Create a SolverContext with allocated workspace arrays. @@ -1950,7 +1906,7 @@ def _update_gradient_init_h_sparse( nv: int, M_elemid: wp.array2d[int], # Data in: - M_in: wp.array3d[float], + M_in: wp.array2d[float], # In: ctx_done_in: wp.array[bool], # Out: @@ -1969,10 +1925,10 @@ def _update_gradient_init_h_sparse( ctx_h_out[worldid, i, j] = 0.0 return - # M is stored in the lower triangle, so transpose the lookup for the upper + # sparse M is stored in the lower triangle, so transpose the lookup for the upper elemid = M_elemid[j, i] if elemid >= 0: - ctx_h_out[worldid, i, j] = M_in[worldid, 0, elemid] + ctx_h_out[worldid, i, j] = M_in[worldid, elemid] else: ctx_h_out[worldid, i, j] = 0.0 @@ -1994,7 +1950,14 @@ def _active_check(tid: int, threshold: int) -> float: @cache_kernel -def _update_gradient_JTDAJ_dense_tiled(nv_pad: int, tile_size: int, njmax: int): +def _update_gradient_JTDAJ_dense_tiled_compact(nv_pad: int, tile_size: int, njmax: int): + """Compact-path variant of _update_gradient_JTDAJ_dense_tiled. + + Takes M_in as a dense 3D array (nworld, nv_pad, nv_pad) -- the compacted active-DOF + inertia block cM -- instead of a 2D CSR array. Cholesky reads fill_mode="upper"; + cM is full-symmetric so the tile_load covers both triangles and the upper triangle + of the result is correct. + """ if njmax < tile_size: tile_size = njmax @@ -2004,7 +1967,7 @@ def _update_gradient_JTDAJ_dense_tiled(nv_pad: int, tile_size: int, njmax: int): def kernel( # Data in: nefc_in: wp.array[int], - M_in: wp.array3d[float], + M_in: wp.array3d[float], # kernel_analyzer: ignore; compact dense inertia (nworld, nv_pad, nv_pad) efc_J_in: wp.array3d[float], efc_D_in: wp.array2d[float], efc_state_in: wp.array2d[int], @@ -2013,14 +1976,87 @@ def _update_gradient_JTDAJ_dense_tiled(nv_pad: int, tile_size: int, njmax: int): # Out: ctx_h_out: wp.array3d[float], ): - worldid = wp.tid() + worldid, rank = wp.tid() if ctx_done_in[worldid]: return nefc = nefc_in[worldid] - sum_val = wp.tile_load(M_in[worldid], shape=(nv_pad, nv_pad), bounds_check=True) + # Load full symmetric compact inertia tile (Cholesky only reads upper triangle). + sum_val = wp.tile_load(M_in[worldid], shape=(wp.static(nv_pad), wp.static(nv_pad)), bounds_check=True) + + for k in range(0, njmax, TILE_SIZE_K): + if k >= nefc: + break + + J_kj = wp.tile_load(efc_J_in[worldid], shape=(TILE_SIZE_K, nv_pad), offset=(k, 0), bounds_check=False) + + D_k = wp.tile_load(efc_D_in[worldid], shape=TILE_SIZE_K, offset=k, bounds_check=False) + state = wp.tile_load(efc_state_in[worldid], shape=TILE_SIZE_K, offset=k, bounds_check=False) + + D_k = wp.tile_map(_state_check, D_k, state) + + tid_tile = wp.tile_arange(TILE_SIZE_K, dtype=int) + threshold_tile = wp.tile_ones(shape=TILE_SIZE_K, dtype=int) * (nefc - k) + + active_tile = wp.tile_map(_active_check, tid_tile, threshold_tile) + D_k = wp.tile_map(wp.mul, active_tile, D_k) + + J_ki = wp.tile_map(wp.mul, wp.tile_transpose(J_kj), wp.tile_broadcast(D_k, shape=(nv_pad, TILE_SIZE_K))) + + sum_val += wp.tile_matmul(J_ki, J_kj) + + wp.tile_store(ctx_h_out[worldid], sum_val, bounds_check=False) + + return kernel + + +@cache_kernel +def _update_gradient_JTDAJ_dense_tiled(nv_pad: int, tile_size: int, njmax: int, nC: int): + if njmax < tile_size: + tile_size = njmax + + TILE_SIZE_K = tile_size + + @wp.kernel(module="unique", enable_backward=False, module_options={"enable_mathdx_gemm": False}) + def kernel( + # Model: + M_colind: wp.array[int], # column index of each CSR entry + M_hinit_i: wp.array[int], # row index of each CSR entry + # Data in: + nefc_in: wp.array[int], + M_in: wp.array2d[float], # CSR M (nworld, nC) + efc_J_in: wp.array3d[float], + efc_D_in: wp.array2d[float], + efc_state_in: wp.array2d[int], + # In: + ctx_done_in: wp.array[bool], + # Out: + ctx_h_out: wp.array3d[float], + ): + worldid, rank = wp.tid() + + if ctx_done_in[worldid]: + return + + nefc = nefc_in[worldid] + + # Densify M's upper triangle from CSR into the shared H tile (Cholesky reads fill_mode="upper"). + # Entry (i, col<=i) of M goes to the upper position (col, i); the lower triangle stays zero. + m_tile = wp.tile_zeros(shape=(wp.static(nv_pad * nv_pad),), dtype=float, storage="shared") + # Uniform trip count from the RUNTIME lane count: every lane iterates iters times (so the + # collective tile_scatter_add is always called by all lanes -- a divergent range(rank, nC, bd) + # deadlocks). wp.block_dim() is 1 on the CPU backend, so lane 0 then covers every entry. + lanes = wp.block_dim() + iters = (nC + lanes - 1) // lanes + for it in range(iters): + e = it * lanes + rank + enable = e < nC + ec = wp.where(enable, e, 0) + pos = M_colind[ec] * nv_pad + M_hinit_i[ec] + wp.tile_scatter_add(m_tile, pos, wp.where(enable, M_in[worldid, ec], 0.0), enable) + sum_val = wp.tile_reshape(m_tile, (nv_pad, nv_pad)) # Each tile processes one output tile by looping over all constraints for k in range(0, njmax, TILE_SIZE_K): @@ -2104,6 +2140,8 @@ def _update_gradient_JTCJ_sparse( continue efcid0 = contact_efc_address_in[conid, 0] + if efcid0 < 0: + continue if efc_state_in[worldid, efcid0] != types.ConstraintState.CONE: continue @@ -2141,7 +2179,10 @@ def _update_gradient_JTCJ_sparse( tt = float(0.0) for j in range(1, condim): efcidj = contact_efc_address_in[conid, j] - uj = ctx_Jaref_in[worldid, efcidj] * fri[j - 1] + if efcidj >= 0: + uj = ctx_Jaref_in[worldid, efcidj] * fri[j - 1] + else: + uj = 0.0 tt += uj * uj u[j] = uj @@ -2165,6 +2206,8 @@ def _update_gradient_JTCJ_sparse( dm_fri1 = dm * mu else: efcid1 = contact_efc_address_in[conid, dim1id] + if efcid1 < 0: + continue rowadr1 = efc_J_rowadr_in[worldid, efcid1] dm_fri1 = dm * fri[dim1id - 1] @@ -2180,6 +2223,8 @@ def _update_gradient_JTCJ_sparse( dm_fri12 = dm_fri1 * mu else: efcid2 = contact_efc_address_in[conid, dim2id] + if efcid2 < 0: + continue rowadr2 = efc_J_rowadr_in[worldid, efcid2] dm_fri12 = dm_fri1 * fri[dim2id - 1] @@ -2215,6 +2260,181 @@ def _update_gradient_JTCJ_sparse( wp.atomic_add(ctx_h_out[worldid, dof1id], dof2id, h) +@wp.kernel +def _update_gradient_JTCJ_compact( + # Model: + opt_impratio_invsqrt: wp.array[float], + # Data in: + contact_dist_in: wp.array[float], + contact_includemargin_in: wp.array[float], + contact_friction_in: wp.array[types.vec5], + contact_dim_in: wp.array[int], + contact_efc_address_in: wp.array2d[int], + contact_worldid_in: wp.array[int], + efc_J_rownnz_in: wp.array2d[int], + efc_J_rowadr_in: wp.array2d[int], + efc_J_colind_in: wp.array3d[int], + efc_J_in: wp.array3d[float], + efc_D_in: wp.array2d[float], + efc_state_in: wp.array2d[int], + dof_cdof_in: wp.array2d[int], + naconmax_in: int, + nacon_in: wp.array[int], + # In: + ctx_Jaref_in: wp.array2d[float], + ctx_done_in: wp.array[bool], + nblocks_perblock: int, + dim_block: int, + # Out: + ctx_h_out: wp.array3d[float], +): + conid_start, pairid = wp.tid() + + for i in range(nblocks_perblock): + conid = conid_start + i * dim_block + + if conid >= min(nacon_in[0], naconmax_in): + return + + worldid = contact_worldid_in[conid] + if ctx_done_in[worldid]: + continue + + condim = contact_dim_in[conid] + + if condim == 1: + continue + + # check contact status + if contact_dist_in[conid] - contact_includemargin_in[conid] >= 0.0: + continue + + efcid0 = contact_efc_address_in[conid, 0] + if efcid0 < 0: + continue + if efc_state_in[worldid, efcid0] != types.ConstraintState.CONE: + continue + + rownnz = efc_J_rownnz_in[worldid, efcid0] + npairs = rownnz * (rownnz + 1) // 2 + if pairid >= npairs: + continue + + rowadr0 = efc_J_rowadr_in[worldid, efcid0] + pos1 = int(0) + rem = pairid + while rem >= rownnz - pos1: + rem -= rownnz - pos1 + pos1 += 1 + pos2 = pos1 + rem + + dofa = efc_J_colind_in[worldid, 0, rowadr0 + pos1] + dofb = efc_J_colind_in[worldid, 0, rowadr0 + pos2] + + # Map to compacted DOFs + dof1id = dof_cdof_in[worldid, dofa] + dof2id = dof_cdof_in[worldid, dofb] + + if dof1id < 0 or dof2id < 0: + continue + + c_dof1 = wp.min(dof1id, dof2id) + c_dof2 = wp.max(dof1id, dof2id) + + fri = contact_friction_in[conid] + mu = fri[0] * opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] + + mu2 = mu * mu + dm = math.safe_div(efc_D_in[worldid, efcid0], mu2 * (1.0 + mu2)) + + if dm == 0.0: + continue + + n = ctx_Jaref_in[worldid, efcid0] * mu + u = types.vec6(n, 0.0, 0.0, 0.0, 0.0, 0.0) + + tt = float(0.0) + for j in range(1, condim): + efcidj = contact_efc_address_in[conid, j] + if efcidj >= 0: + uj = ctx_Jaref_in[worldid, efcidj] * fri[j - 1] + else: + uj = 0.0 + tt += uj * uj + u[j] = uj + + if tt <= 0.0: + t = 0.0 + else: + t = wp.sqrt(tt) + t = wp.max(t, types.MJ_MINVAL) + ttt = wp.max(t * t * t, types.MJ_MINVAL) + + # Precompute common subexpressions. + mu_over_t = math.safe_div(mu, t) + mu_n_over_ttt = mu * math.safe_div(n, ttt) + mu2_minus_mu_n_over_t = mu2 - mu * math.safe_div(n, t) + + h = float(0.0) + + for dim1id in range(condim): + if dim1id == 0: + efcid1 = efcid0 + dm_fri1 = dm * mu + else: + efcid1 = contact_efc_address_in[conid, dim1id] + if efcid1 < 0: + continue + dm_fri1 = dm * fri[dim1id - 1] + + # Read from the compacted dense Jacobian (efc_J_in) using the mapped compacted DOFs + efc_J11 = efc_J_in[worldid, efcid1, c_dof1] + efc_J12 = efc_J_in[worldid, efcid1, c_dof2] + + ui = u[dim1id] + + for dim2id in range(0, dim1id + 1): + if dim2id == 0: + efcid2 = efcid0 + dm_fri12 = dm_fri1 * mu + else: + efcid2 = contact_efc_address_in[conid, dim2id] + if efcid2 < 0: + continue + dm_fri12 = dm_fri1 * fri[dim2id - 1] + + # Read from the compacted dense Jacobian using the mapped compacted DOFs + efc_J21 = efc_J_in[worldid, efcid2, c_dof1] + efc_J22 = efc_J_in[worldid, efcid2, c_dof2] + + uj = u[dim2id] + + # set first row/column: (1, -mu/t * u) + if dim1id == 0 and dim2id == 0: + hcone = 1.0 + elif dim1id == 0: + hcone = -mu_over_t * uj + elif dim2id == 0: + hcone = -mu_over_t * ui + else: + hcone = mu_n_over_ttt * ui * uj + + # add to diagonal: mu^2 - mu * n / t + if dim1id == dim2id: + hcone += mu2_minus_mu_n_over_t + + hcone *= dm_fri12 + + if hcone != 0.0: + h += hcone * efc_J11 * efc_J22 + + if dim1id != dim2id: + h += hcone * efc_J12 * efc_J21 + + # multiple contacts can contribute to the same (c_dof1, c_dof2); atomic_add is exact + wp.atomic_add(ctx_h_out[worldid, c_dof1], c_dof2, h) + + @wp.kernel def _update_gradient_JTCJ_dense( # Model: @@ -2266,6 +2486,8 @@ def _update_gradient_JTCJ_dense( continue efcid0 = contact_efc_address_in[conid, 0] + if efcid0 < 0: + continue if efc_state_in[worldid, efcid0] != types.ConstraintState.CONE: continue @@ -2284,7 +2506,10 @@ def _update_gradient_JTCJ_dense( tt = float(0.0) for j in range(1, condim): efcidj = contact_efc_address_in[conid, j] - uj = ctx_Jaref_in[worldid, efcidj] * fri[j - 1] + if efcidj >= 0: + uj = ctx_Jaref_in[worldid, efcidj] * fri[j - 1] + else: + uj = 0.0 tt += uj * uj u[j] = uj @@ -2302,6 +2527,8 @@ def _update_gradient_JTCJ_dense( efcid1 = efcid0 else: efcid1 = contact_efc_address_in[conid, dim1id] + if efcid1 < 0: + continue efc_J11 = efc_J_in[worldid, efcid1, dof1id] efc_J12 = efc_J_in[worldid, efcid1, dof2id] @@ -2313,6 +2540,8 @@ def _update_gradient_JTCJ_dense( efcid2 = efcid0 else: efcid2 = contact_efc_address_in[conid, dim2id] + if efcid2 < 0: + continue efc_J21 = efc_J_in[worldid, efcid2, dof1id] efc_J22 = efc_J_in[worldid, efcid2, dof2id] @@ -2373,16 +2602,16 @@ def _update_gradient_cholesky(tile_size: int): return mat_tile = wp.tile_load(h_in[worldid], shape=(TILE_SIZE, TILE_SIZE)) - fact_tile = wp.tile_cholesky(mat_tile, fill_mode="upper") + wp.tile_cholesky_inplace(mat_tile, fill_mode="upper") input_tile = wp.tile_load(ctx_grad_in[worldid], shape=TILE_SIZE) - output_tile = wp.tile_cholesky_solve(fact_tile, input_tile, fill_mode="upper") + output_tile = wp.tile_cholesky_solve(mat_tile, input_tile, fill_mode="upper") wp.tile_store(ctx_Mgrad_out[worldid], output_tile) return kernel @cache_kernel -def _update_gradient_cholesky_blocked(tile_size: int, matrix_size: int): +def _update_gradient_cholesky_blocked(tile_size: int, matrix_size: int, check_skip: bool = True): @wp.kernel(module="unique", enable_backward=False, module_options={"enable_mathdx_gemm": False}) def kernel( # In: @@ -2396,8 +2625,9 @@ def _update_gradient_cholesky_blocked(tile_size: int, matrix_size: int): worldid = wp.tid() TILE_SIZE = wp.static(tile_size) - if ctx_done_in[worldid]: - return + if wp.static(check_skip): + if ctx_done_in[worldid]: + return # We need matrix size both as a runtime input as well as a static input: # static input is needed to specify the tile sizes for the compiler @@ -2495,10 +2725,25 @@ def _cholesky_factorize_solve(m: types.Model, d: types.Data, ctx: SolverContext, ) +# --------------------------------------------------------------------------- +# H += J^T D J. D diagonal, so each efc row adds one rank-1 outer product. make_constraint +# groups a constraint's contiguous efc rows (shared colind = dof support S) into one |S|x|S| +# block, stored densely per world in efc.jtdaj_{adr,nrow,nblock}. The launch fills the +# GPU once (groups_per_world slots/world) then grid-strides the rest, so no thread lands on a +# non-head efc row. A block's upper-triangular entries split across THREADS_PER_GROUP threads +# (one warp -> coalesced J reads); entry -> (block_row, block_col) is the triangular-number +# inverse, exact in float32 since column boundaries are perfect squares (8*entry+1 = (2c+1)^2). +# --------------------------------------------------------------------------- +_JTDAJ_THREADS_PER_GROUP = 32 # one warp per group, so its J reads coalesce +_JTDAJ_OVERSUBSCRIBE_WAVES = 6 # grid-stride depth; short per-warp chains load-balance groups + + @wp.kernel def _JTDAJ_sparse( # Data in: - nefc_in: wp.array[int], + efc_jtdaj_adr_in: wp.array2d[int], + efc_jtdaj_nrow_in: wp.array2d[int], + efc_jtdaj_nblock_in: wp.array[int], efc_J_rownnz_in: wp.array2d[int], efc_J_rowadr_in: wp.array2d[int], efc_J_colind_in: wp.array3d[int], @@ -2507,48 +2752,50 @@ def _JTDAJ_sparse( efc_state_in: wp.array2d[int], # In: ctx_done_in: wp.array[bool], + groups_per_world: int, # Out: h_out: wp.array3d[float], ): - worldid, efcid = wp.tid() - + worldid, slot, lane = wp.tid() if ctx_done_in[worldid]: return - - if efcid >= nefc_in[worldid]: - return - - efc_D = efc_D_in[worldid, efcid] - efc_state = efc_state_in[worldid, efcid] - - if _state_check(efc_D, efc_state) == 0.0: - return - - rownnz = efc_J_rownnz_in[worldid, efcid] - rowadr = efc_J_rowadr_in[worldid, efcid] - - for i in range(rownnz): - sparseidi = rowadr + i - Ji = efc_J_in[worldid, 0, sparseidi] - colindi = efc_J_colind_in[worldid, 0, sparseidi] - for j in range(i, rownnz): - if j == i: - sparseidj = sparseidi - Jj = Ji - colindj = colindi - else: - sparseidj = rowadr + j - Jj = efc_J_in[worldid, 0, sparseidj] - colindj = efc_J_colind_in[worldid, 0, sparseidj] - - h = Ji * Jj * efc_D - # Store in upper triangle only: ensure row <= col. - row = wp.min(colindi, colindj) - col = wp.max(colindi, colindj) - wp.atomic_add(h_out[worldid, row], col, h) + count = efc_jtdaj_nblock_in[worldid] + for groupid in range(slot, count, groups_per_world): # grid-stride this world's group list + head_row = efc_jtdaj_adr_in[worldid, groupid] + block_rows = efc_jtdaj_nrow_in[worldid, groupid] + head_adr = efc_J_rowadr_in[worldid, head_row] + support = efc_J_rownnz_in[worldid, head_row] # dofs the constraint touches = block dimension + n_entries = support * (support + 1) // 2 # upper-triangular entries of the |S|x|S| block + for entry in range(lane, n_entries, wp.static(_JTDAJ_THREADS_PER_GROUP)): + block_col = int((wp.sqrt(float(8 * entry + 1)) - 1.0) * 0.5) + block_row = entry - block_col * (block_col + 1) // 2 + dof_row = efc_J_colind_in[worldid, 0, head_adr + block_row] + dof_col = efc_J_colind_in[worldid, 0, head_adr + block_col] + hval = float(0.0) + for member in range(block_rows): + member_row = head_row + member + if efc_state_in[worldid, member_row] == types.ConstraintState.QUADRATIC.value: + member_adr = efc_J_rowadr_in[worldid, member_row] + j_row = efc_J_in[worldid, 0, member_adr + block_row] + j_col = efc_J_in[worldid, 0, member_adr + block_col] + hval += j_row * efc_D_in[worldid, member_row] * j_col + if hval != 0.0: # skip the atomic when no member row is active + wp.atomic_add(h_out[worldid, wp.min(dof_row, dof_col)], wp.max(dof_row, dof_col), hval) -def _update_gradient(m: types.Model, d: types.Data, ctx: SolverContext): +def _jtdaj_groups_per_world(nworld: int, njmax: int) -> int: + # Per-world width of the grid stride. Target one warp per group-slot (njmax), but cap the grid at + # _JTDAJ_OVERSUBSCRIBE_WAVES device waves -- else high-njmax worlds dispatch many idle tail warps + # (njmax >> actual groups). A few waves of oversubscription keep each warp's serial chain short, + # load-balancing the variable group sizes (measured plateau: ~4-8 waves). + block_size, min_grid_size = wp.get_suggested_block_size(_JTDAJ_sparse) + # block_size * min_grid_size = full-device thread count (block_size cancels): the kernel's max + # resident threads (one wave), a device property independent of nworld and our launch block_dim. + device_warps = max(1, block_size * min_grid_size // _JTDAJ_THREADS_PER_GROUP) + return max(1, min(njmax, _JTDAJ_OVERSUBSCRIBE_WAVES * device_warps // nworld)) + + +def _update_gradient(m: types.Model, d: types.Data, ctx: SolverContext, compact: bool = False): # grad = Ma - qfrc_smooth - qfrc_constraint if m.opt.solver == types.SolverType.CG: wp.launch_tiled( @@ -2579,27 +2826,60 @@ def _update_gradient(m: types.Model, d: types.Data, ctx: SolverContext): outputs=[ctx.h], ) + groups_per_world = _jtdaj_groups_per_world(d.nworld, d.njmax) wp.launch( _JTDAJ_sparse, - dim=(d.nworld, d.njmax), - inputs=[d.nefc, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, d.efc.J, d.efc.D, d.efc.state, ctx.done], - outputs=[ctx.h], - ) - else: - wp.launch_tiled( - _update_gradient_JTDAJ_dense_tiled(m.nv_pad, types.TILE_SIZE_JTDAJ_DENSE, d.njmax), - dim=d.nworld, + dim=(d.nworld, groups_per_world, _JTDAJ_THREADS_PER_GROUP), inputs=[ - d.nefc, - d.M, + d.efc.jtdaj_adr, + d.efc.jtdaj_nrow, + d.efc.jtdaj_nblock, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, d.efc.J, d.efc.D, d.efc.state, ctx.done, + groups_per_world, ], outputs=[ctx.h], - block_dim=m.block_dim.update_gradient_JTDAJ_dense, + block_dim=m.block_dim.update_gradient_JTDAJ_sparse, ) + else: + if compact: + # compact path: d.M is the dense 3D compact inertia block (nworld, nv_pad, nv_pad) + wp.launch_tiled( + _update_gradient_JTDAJ_dense_tiled_compact(m.nv_pad, types.TILE_SIZE_JTDAJ_DENSE, d.njmax), + dim=d.nworld, + inputs=[ + d.nefc, + d.M, + d.efc.J, + d.efc.D, + d.efc.state, + ctx.done, + ], + outputs=[ctx.h], + block_dim=m.block_dim.update_gradient_JTDAJ_dense, + ) + else: + wp.launch_tiled( + _update_gradient_JTDAJ_dense_tiled(m.nv_pad, types.TILE_SIZE_JTDAJ_DENSE, d.njmax, m.M_colind.shape[0]), + dim=d.nworld, + inputs=[ + m.M_colind, + m.M_hinit_i, + d.nefc, + d.M, + d.efc.J, + d.efc.D, + d.efc.state, + ctx.done, + ], + outputs=[ctx.h], + block_dim=m.block_dim.update_gradient_JTDAJ_dense, + ) if m.opt.cone == types.ConeType.ELLIPTIC: # Optimization: launching update_gradient_JTCJ with limited number of blocks on a GPU. @@ -2613,7 +2893,13 @@ def _update_gradient(m: types.Model, d: types.Data, ctx: SolverContext): # we don't over-launch naconmax (capacity) threads when active contacts are far fewer. The # sparse kernel uses one thread per (contact, support-pair) (jtcj_max_pairs), the dense one # per (contact, dof-pair) (dof_tri_row.size). - jtcj_second_dim = m.jtcj_max_pairs if m.is_sparse else m.dof_tri_row.size + # `compact` is set by solve_compact's inner solve, which runs the dense factor/solve on + # the nvmax_pad block but maps sparse contact support-pairs to compacted DOFs via dof_cdof. + # (Don't infer it from `d.nvmax < m.nv`: after solve_compact's shallow m2/d2 replace that + # reduces to `nvmax < nvmax_pad`, which is false whenever nvmax is a tile multiple and + # silently falls back to the O(nvmax_pad^2) dense cone scan.) + is_sparse_compact = compact and (d.efc.J_colind.shape[1] > 0) + jtcj_second_dim = m.jtcj_max_pairs if (m.is_sparse or is_sparse_compact) else m.dof_tri_row.size if wp.get_device().is_cuda: sm_count = wp.get_device().sm_count @@ -2656,31 +2942,60 @@ def _update_gradient(m: types.Model, d: types.Data, ctx: SolverContext): outputs=[ctx.h], ) else: - wp.launch( - _update_gradient_JTCJ_dense, - dim=(dim_block, m.dof_tri_row.size), - inputs=[ - m.opt.impratio_invsqrt, - m.dof_tri_row, - m.dof_tri_col, - d.contact.dist, - d.contact.includemargin, - d.contact.friction, - d.contact.dim, - d.contact.efc_address, - d.contact.worldid, - d.efc.J, - d.efc.D, - d.efc.state, - d.naconmax, - d.nacon, - ctx.Jaref, - ctx.done, - nblocks_perblock, - dim_block, - ], - outputs=[ctx.h], - ) + if is_sparse_compact: + wp.launch( + _update_gradient_JTCJ_compact, + dim=(dim_block, m.jtcj_max_pairs), + inputs=[ + m.opt.impratio_invsqrt, + d.contact.dist, + d.contact.includemargin, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.contact.worldid, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.D, + d.efc.state, + d.dof_cdof, + d.naconmax, + d.nacon, + ctx.Jaref, + ctx.done, + nblocks_perblock, + dim_block, + ], + outputs=[ctx.h], + ) + else: + wp.launch( + _update_gradient_JTCJ_dense, + dim=(dim_block, m.dof_tri_row.size), + inputs=[ + m.opt.impratio_invsqrt, + m.dof_tri_row, + m.dof_tri_col, + d.contact.dist, + d.contact.includemargin, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.contact.worldid, + d.efc.J, + d.efc.D, + d.efc.state, + d.naconmax, + d.nacon, + ctx.Jaref, + ctx.done, + nblocks_perblock, + dim_block, + ], + outputs=[ctx.h], + ) _cholesky_factorize_solve(m, d, ctx) else: @@ -2982,6 +3297,7 @@ def _solver_iteration( d: types.Data, ctx: SolverContext, nsolving: wp.array[int], + compact: bool = False, ): _linesearch(m, d, ctx) @@ -3000,7 +3316,7 @@ def _solver_iteration( if incremental: _update_gradient_incremental(m, d, ctx) else: - _update_gradient(m, d, ctx) + _update_gradient(m, d, ctx, compact=compact) # polak-ribiere if m.opt.solver == types.SolverType.CG: @@ -3071,7 +3387,7 @@ def _solver_iteration( ) -def init_context(m: types.Model, d: types.Data, ctx: SolverContext | InverseContext, grad: bool = True): +def init_context(m: types.Model, d: types.Data, ctx: SolverContext | InverseContext, grad: bool = True, compact: bool = False): # initialize some efc arrays wp.launch( _solve_init_efc, @@ -3107,29 +3423,29 @@ def init_context(m: types.Model, d: types.Data, ctx: SolverContext | InverseCont _update_constraint(m, d, ctx) if grad: - _update_gradient(m, d, ctx) + _update_gradient(m, d, ctx, compact=compact) @event_scope def solve(m: types.Model, d: types.Data): + if m.opt.enableflags & types.EnableBit.SLEEP: + # Self-contained like the island branch below: rebuild the active-DOF mapping from + # tree_awake so solve() works when called directly (not only via fwd_acceleration). + island.update_active_dofs(m, d) + solve_compact(m, d) + if m.ntree > 1: + island.compute_island_mapping(m, d) + return + if d.njmax == 0 or m.nv == 0: wp.copy(d.qacc, d.qacc_smooth) d.solver_niter.fill_(0) else: - if m.ntree > 1 and not (m.opt.disableflags & types.DisableBit.ISLAND): - ctx = _create_island_solver_context(m, d) - island.compute_island_mapping(m, d, ctx) - island.gather_island_inputs(m, d, ctx) - _solve_island(m, d, ctx) - # Ma is needed by Euler/implicit integrators for implicit damping - scatter_Ma = m.opt.integrator != types.IntegratorType.RK4 - island.scatter_island_results(m, d, ctx, scatter_Ma=scatter_Ma) - else: - ctx = _create_solver_context(m, d) - _solve(m, d, ctx) + ctx = _create_solver_context(m, d) + _solve(m, d, ctx) -def _solve(m: types.Model, d: types.Data, ctx: SolverContext): +def _solve(m: types.Model, d: types.Data, ctx: SolverContext, compact: bool = False): """Finds forces that satisfy constraints.""" if not (m.opt.disableflags & types.DisableBit.WARMSTART): wp.copy(d.qacc, d.qacc_warmstart) @@ -3137,7 +3453,7 @@ def _solve(m: types.Model, d: types.Data, ctx: SolverContext): wp.copy(d.qacc, d.qacc_smooth) # context - init_context(m, d, ctx, grad=True) + init_context(m, d, ctx, grad=True, compact=compact) # search = -Mgrad if m.opt.solver == types.SolverType.CG: @@ -3166,2176 +3482,322 @@ def _solve(m: types.Model, d: types.Data, ctx: SolverContext): # When the number of iterations reaches m.opt.iterations, solver_niter # becomes zero and all worlds are marked as converged to avoid an infinite loop. # note: we only launch the iteration kernel if everything is not done - wp.capture_while(nsolving, while_body=_solver_iteration, m=m, d=d, ctx=ctx, nsolving=nsolving) + wp.capture_while(nsolving, while_body=_solver_iteration, m=m, d=d, ctx=ctx, nsolving=nsolving, compact=compact) else: # This branch is mostly for when JAX is used as it is currently not compatible # with CUDA graph conditional. # It should be removed when JAX becomes compatible. for _ in range(m.opt.iterations): - _solver_iteration(m, d, ctx, nsolving) + _solver_iteration(m, d, ctx, nsolving, compact=compact) -# TODO(team): Consolidate monolithic and island solver code where possible -@event_scope -def _solve_island(m: types.Model, d: types.Data, ctx: IslandSolverContext): - """Solve constraints for all islands in parallel. - - All islands are processed simultaneously. Island-local arrays in ctx - (iacc, iefc_J, iefc_D, etc.) are indexed by idof/iefc, and each thread - determines its island via idof_islandid/iefc_islandid lookup tables. - """ - # Initialize iacc from warmstart or smooth - if not (m.opt.disableflags & types.DisableBit.WARMSTART): - wp.launch( - _gather_warmstart_island, - dim=(d.nworld, m.nv), - inputs=[d.nidof, d.qacc_warmstart, d.map_idof2dof], - outputs=[d.iqacc], - ) - else: - wp.copy(d.iqacc, d.iqacc_smooth) - - # nsolving tracks how many active islands still have unconverged globally - nsolving = wp.zeros((1,), dtype=int) - - # Initialize island context - _init_context_island(m, d, ctx, nsolving) - - # search = -Mgrad - wp.launch( - _solve_init_search_island, - dim=(d.nworld, m.nv), - inputs=[d.nidof, ctx.Mgrad, d.dof_islandid, ctx.done], - outputs=[ctx.search, ctx.search_dot], - ) - - if m.opt.iterations != 0 and m.opt.graph_conditional: - wp.capture_while( - nsolving, - while_body=_solver_iteration_island, - m=m, - d=d, - ctx=ctx, - nsolving=nsolving, - ) - else: - for _ in range(m.opt.iterations): - _solver_iteration_island(m, d, ctx, nsolving) +# Active-DOF compaction solve (nvmax < nv). +# +# When fewer than nv DOFs are active, the active set is packed into a contiguous +# [0, ncdof) range (see island.update_active_dofs) and the dense factor/solve runs at +# size nvmax_pad instead of nv. The compacted workspace lives on Data as c* fields +# (mirroring the island-local i* fields). The constrained solve gathers into compacted +# arrays, shallow-replaces (m, d) so the stock dense Newton solver runs at nvmax_pad, +# then scatters qacc/qfrc_constraint back, freezing inactive DOFs to 0. @wp.kernel -def _gather_warmstart_island( +def _init_compact_inertia( # Data in: - nidof_in: wp.array[int], - qacc_warmstart_in: wp.array2d[float], - map_idof2dof_in: wp.array2d[int], + ncdof_in: wp.array[int], # Out: - iacc_out: wp.array2d[float], + M_c_out: wp.array3d[float], ): - """Gather qacc_warmstart into island-local order.""" - worldid, idofid = wp.tid() - - if idofid >= nidof_in[worldid]: - return - - dof = map_idof2dof_in[worldid, idofid] - iacc_out[worldid, idofid] = qacc_warmstart_in[worldid, dof] - - -@cache_kernel -def _solve_init_efc_island(enable_sleep: bool): - @wp.kernel(module="unique", enable_backward=False) - def kernel( - # Model: - ntree: int, - # Data in: - nisland_in: wp.array[int], - tree_awake_in: wp.array2d[int], - tree_island_in: wp.array2d[int], - # Out: - island_cost_out: wp.array2d[float], - island_search_dot_out: wp.array2d[float], - island_done_out: wp.array2d[bool], - island_solver_niter_out: wp.array2d[int], - nsolving_out: wp.array[int], - ): - """Initialize per-island solver scalars.""" - worldid, islandid = wp.tid() - - if islandid >= nisland_in[worldid]: - return - - is_asleep_flag = int(0) - if wp.static(enable_sleep): - has_awake_tree = int(0) - for t in range(ntree): - if tree_island_in[worldid, t] == islandid: - if tree_awake_in[worldid, t] == 1: - has_awake_tree = int(1) - break - if has_awake_tree == 0: - is_asleep_flag = int(1) - - island_cost_out[worldid, islandid] = 0.0 - island_search_dot_out[worldid, islandid] = 0.0 - island_done_out[worldid, islandid] = is_asleep_flag == 1 - island_solver_niter_out[worldid, islandid] = 0 - if is_asleep_flag == 0: - wp.atomic_add(nsolving_out, 0, 1) - - return kernel + worldid, i, j = wp.tid() + val = 0.0 + if i == j and i >= ncdof_in[worldid]: + val = 1.0 + M_c_out[worldid, i, j] = val @wp.kernel -def _solve_init_jaref_island( +def _gather_M_sparse( # Model: - is_sparse: bool, + M_rownnz: wp.array[int], + M_rowadr: wp.array[int], + M_colind: wp.array[int], # Data in: - nefc_in: wp.array[int], - island_idofadr_in: wp.array2d[int], - island_nv_in: wp.array2d[int], - njmax_in: int, - # In: - iefc_J_rownnz_in: wp.array2d[int], - iefc_J_rowadr_in: wp.array2d[int], - iefc_J_colind_in: wp.array3d[int], - iefc_J_in: wp.array3d[float], - iacc_in: wp.array2d[float], - iefc_aref_in: wp.array2d[float], - iefc_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], + M_in: wp.array2d[float], + dof_cdof_in: wp.array2d[int], # Out: - Jaref_out: wp.array2d[float], + M_c_out: wp.array3d[float], ): - """Jaref[iefcid] = iefc_J[iefcid] @ iacc - iefc_aref[iefcid] for all island EFCs.""" - worldid, iefcid = wp.tid() - - if iefcid >= wp.min(njmax_in, nefc_in[worldid]): + worldid, i = wp.tid() + ci = dof_cdof_in[worldid, i] + if ci < 0: return - - islandid = iefc_islandid_in[worldid, iefcid] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - acc = float(0.0) - if is_sparse: - rownnz = iefc_J_rownnz_in[worldid, iefcid] - rowadr = iefc_J_rowadr_in[worldid, iefcid] - for k in range(rownnz): - adr = rowadr + k - Ji = iefc_J_in[worldid, 0, adr] - idof = iefc_J_colind_in[worldid, 0, adr] - acc += Ji * iacc_in[worldid, idof] - else: - idofadr = island_idofadr_in[worldid, islandid] - inv = island_nv_in[worldid, islandid] - for i in range(inv): - idof = idofadr + i - acc += iefc_J_in[worldid, iefcid, idof] * iacc_in[worldid, idof] - - Jaref_out[worldid, iefcid] = acc - iefc_aref_in[worldid, iefcid] - - -@wp.kernel -def _solve_init_search_island( - # Data in: - nidof_in: wp.array[int], - # In: - Mgrad_in: wp.array2d[float], - idof_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - search_out: wp.array2d[float], - island_search_dot_out: wp.array2d[float], -): - """Search = -Mgrad for all island DOFs.""" - worldid, idofid = wp.tid() - - if idofid >= nidof_in[worldid]: - return - - islandid = idof_islandid_in[worldid, idofid] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - s = -Mgrad_in[worldid, idofid] - search_out[worldid, idofid] = s - wp.atomic_add(island_search_dot_out, worldid, islandid, s * s) - - -# TODO(team): remove after updating island solver done criteria to use delta cost -@wp.kernel -def _update_constraint_init_cost( - # In: - cost_in: wp.array[float], - done_in: wp.array[bool], - # Out: - gauss_out: wp.array[float], - cost_out: wp.array[float], - prev_cost_out: wp.array[float], -): - tid = wp.tid() - if done_in[tid]: - return - - prev_cost_out[tid] = cost_in[tid] - cost_out[tid] = 0.0 - gauss_out[tid] = 0.0 - - -@wp.kernel -def _update_constraint_efc_island( - # Model: - opt_impratio_invsqrt: wp.array[float], - # Data in: - nefc_in: wp.array[int], - contact_friction_in: wp.array[types.vec5], - contact_dim_in: wp.array[int], - contact_efc_address_in: wp.array2d[int], - island_nefc_in: wp.array2d[int], - island_ne_in: wp.array2d[int], - island_nf_in: wp.array2d[int], - island_efcadr_in: wp.array2d[int], - map_efc2iefc_in: wp.array2d[int], - njmax_in: int, - nacon_in: wp.array[int], - # In: - iefc_type_in: wp.array2d[int], - iefc_id_in: wp.array2d[int], - iefc_D_in: wp.array2d[float], - iefc_frictionloss_in: wp.array2d[float], - Jaref_in: wp.array2d[float], - iefc_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - iefc_force_out: wp.array2d[float], - iefc_state_out: wp.array2d[int], - island_cost_out: wp.array2d[float], -): - """Compute force, state, and cost for each island constraint.""" - worldid, iefcid = wp.tid() - - if iefcid >= wp.min(njmax_in, nefc_in[worldid]): - return - - islandid = iefc_islandid_in[worldid, iefcid] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - # Local position within island - iefcadr = island_efcadr_in[worldid, islandid] - local_iefcid = iefcid - iefcadr - ine = island_ne_in[worldid, islandid] - inf = island_nf_in[worldid, islandid] - - jaref = Jaref_in[worldid, iefcid] - D = iefc_D_in[worldid, iefcid] - - is_equality = local_iefcid < ine - is_friction = (not is_equality) and (local_iefcid < ine + inf) - is_elliptic = iefc_type_in[worldid, iefcid] == types.ConstraintType.CONTACT_ELLIPTIC - - frictionloss = iefc_frictionloss_in[worldid, iefcid] if is_friction else 0.0 - - ic0 = int(-1) - jaref0 = float(0.0) - D0 = float(0.0) - mu = float(0.0) - ufrictionj = float(0.0) - TT = float(0.0) - - if is_elliptic: - conid = iefc_id_in[worldid, iefcid] - if conid >= nacon_in[0]: - return - efcid0_global = contact_efc_address_in[conid, 0] - if efcid0_global < 0: - return - ic0 = map_efc2iefc_in[worldid, efcid0_global] - - dim = contact_dim_in[conid] - friction = contact_friction_in[conid] - mu = friction[0] * opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] - jaref0 = Jaref_in[worldid, ic0] - D0 = iefc_D_in[worldid, ic0] - - for j in range(1, dim): - efcidj_global = contact_efc_address_in[conid, j] - if efcidj_global < 0: - return - icj = map_efc2iefc_in[worldid, efcidj_global] - frictionj = friction[j - 1] - uj = Jaref_in[worldid, icj] * frictionj - TT += uj * uj - if iefcid == icj: - ufrictionj = uj * frictionj - - res = _eval_constraint( - is_equality, - is_friction, - is_elliptic, - jaref, - D, - frictionloss, - iefcid, - ic0, - jaref0, - D0, - mu, - ufrictionj, - TT, - ) - - iefc_force_out[worldid, iefcid] = res[0] - iefc_state_out[worldid, iefcid] = int(res[1]) - cost = res[2] - if cost != 0.0: - wp.atomic_add(island_cost_out, worldid, islandid, cost) - - -@wp.kernel -def _update_constraint_init_qfrc_constraint_dense_island( - # Data in: - nefc_in: wp.array[int], - nidof_in: wp.array[int], - island_nefc_in: wp.array2d[int], - island_efcadr_in: wp.array2d[int], - njmax_in: int, - # In: - iefc_J_in: wp.array3d[float], - iefc_force_in: wp.array2d[float], - idof_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - ifrc_constraint_out: wp.array2d[float], -): - """ifrc_constraint = iefc_J.T @ iefc_force for all island DOFs.""" - worldid, idofid = wp.tid() - - if idofid >= nidof_in[worldid]: - ifrc_constraint_out[worldid, idofid] = 0.0 - return - - islandid = idof_islandid_in[worldid, idofid] - if islandid < 0: - ifrc_constraint_out[worldid, idofid] = 0.0 - return - if island_done_in[worldid, islandid]: - return - - iefcadr = island_efcadr_in[worldid, islandid] - inefc = island_nefc_in[worldid, islandid] - acc = float(0.0) - for iefcid in range(iefcadr, iefcadr + inefc): - acc += iefc_J_in[worldid, iefcid, idofid] * iefc_force_in[worldid, iefcid] - - ifrc_constraint_out[worldid, idofid] = acc - - -@wp.kernel -def _update_constraint_init_qfrc_constraint_sparse_island( - # Data in: - nefc_in: wp.array[int], - njmax_in: int, - # In: - iefc_J_rownnz_in: wp.array2d[int], - iefc_J_rowadr_in: wp.array2d[int], - iefc_J_colind_in: wp.array3d[int], - iefc_J_in: wp.array3d[float], - iefc_force_in: wp.array2d[float], - iefc_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - ifrc_constraint_out: wp.array2d[float], -): - """ifrc_constraint += iefc_J.T @ iefc_force for all island EFCs (sparse parallel per EFC).""" - worldid, iefcid = wp.tid() - - if iefcid >= wp.min(njmax_in, nefc_in[worldid]): - return - - islandid = iefc_islandid_in[worldid, iefcid] - if islandid < 0: - return - - rownnz = iefc_J_rownnz_in[worldid, iefcid] - rowadr = iefc_J_rowadr_in[worldid, iefcid] - force = iefc_force_in[worldid, iefcid] - - for k in range(rownnz): + rowadr = M_rowadr[i] + for k in range(M_rownnz[i]): adr = rowadr + k - Ji = iefc_J_in[worldid, 0, adr] - idof = iefc_J_colind_in[worldid, 0, adr] - wp.atomic_add(ifrc_constraint_out, worldid, idof, Ji * force) + cj = dof_cdof_in[worldid, M_colind[adr]] + if cj >= 0: + val = M_in[worldid, adr] + M_c_out[worldid, ci, cj] = val + M_c_out[worldid, cj, ci] = val @wp.kernel -def _update_constraint_gauss_cost_island( +def _gather_rhs_compact( # Data in: - nidof_in: wp.array[int], + cdof_dof_in: wp.array2d[int], # In: - iacc_in: wp.array2d[float], - ifrc_smooth_in: wp.array2d[float], - iacc_smooth_in: wp.array2d[float], - iMa_in: wp.array2d[float], - idof_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], + vec_in: wp.array2d[float], # Out: - island_gauss_out: wp.array2d[float], - island_cost_out: wp.array2d[float], + rhs_out: wp.array3d[float], ): - """Gauss cost: 0.5 * (Ma - qfrc_smooth).T @ (qacc - qacc_smooth) per island.""" - worldid, idofid = wp.tid() - - if idofid >= nidof_in[worldid]: - return - - islandid = idof_islandid_in[worldid, idofid] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - dq = iacc_in[worldid, idofid] - iacc_smooth_in[worldid, idofid] - df = iMa_in[worldid, idofid] - ifrc_smooth_in[worldid, idofid] - gauss = 0.5 * df * dq - - wp.atomic_add(island_gauss_out, worldid, islandid, gauss) - wp.atomic_add(island_cost_out, worldid, islandid, gauss) - - -@wp.kernel -def _update_gradient_grad_island( - # Data in: - nidof_in: wp.array[int], - # In: - ifrc_smooth_in: wp.array2d[float], - ifrc_constraint_in: wp.array2d[float], - iMa_in: wp.array2d[float], - idof_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - grad_out: wp.array2d[float], - island_grad_dot_out: wp.array2d[float], -): - """Grad = Ma - qfrc_smooth - qfrc_constraint, grad_dot per island.""" - worldid, idofid = wp.tid() - - if idofid >= nidof_in[worldid]: - return - - islandid = idof_islandid_in[worldid, idofid] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - g = iMa_in[worldid, idofid] - ifrc_smooth_in[worldid, idofid] - ifrc_constraint_in[worldid, idofid] - grad_out[worldid, idofid] = g - wp.atomic_add(island_grad_dot_out, worldid, islandid, g * g) - - -@wp.kernel -def _linesearch_jv_island( - # Model: - is_sparse: bool, - # Data in: - nefc_in: wp.array[int], - nidof_in: wp.array[int], - island_idofadr_in: wp.array2d[int], - island_nv_in: wp.array2d[int], - njmax_in: int, - # In: - iefc_J_rownnz_in: wp.array2d[int], - iefc_J_rowadr_in: wp.array2d[int], - iefc_J_colind_in: wp.array3d[int], - iefc_J_in: wp.array3d[float], - search_in: wp.array2d[float], - iefc_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - jv_out: wp.array2d[float], -): - """Jv = iefc_J @ search for all island EFCs.""" - worldid, iefcid = wp.tid() - - if iefcid >= wp.min(njmax_in, nefc_in[worldid]): - return - - islandid = iefc_islandid_in[worldid, iefcid] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - acc = float(0.0) - if is_sparse: - rownnz = iefc_J_rownnz_in[worldid, iefcid] - rowadr = iefc_J_rowadr_in[worldid, iefcid] - for k in range(rownnz): - adr = rowadr + k - Ji = iefc_J_in[worldid, 0, adr] - idof = iefc_J_colind_in[worldid, 0, adr] - acc += Ji * search_in[worldid, idof] + worldid, ci = wp.tid() + dof = cdof_dof_in[worldid, ci] + if dof >= 0: + rhs_out[worldid, ci, 0] = vec_in[worldid, dof] else: - idofadr = island_idofadr_in[worldid, islandid] - inv = island_nv_in[worldid, islandid] - for i in range(inv): - idof = idofadr + i - acc += iefc_J_in[worldid, iefcid, idof] * search_in[worldid, idof] - - jv_out[worldid, iefcid] = acc + rhs_out[worldid, ci, 0] = 0.0 -@wp.func -def _eval_elliptic_cost_island( - # Model: - opt_impratio_invsqrt: float, # kernel_analyzer: off - # Data in: - contact_friction_in: wp.array[types.vec5], - contact_dim_in: wp.array[int], - contact_efc_address_in: wp.array2d[int], - map_efc2iefc_in: wp.array2d[int], - # In: - alpha: float, - conid: int, - iefc_D_in: wp.array2d[float], - Jaref_in: wp.array2d[float], - jv_in: wp.array2d[float], - worldid: int, -) -> wp.vec3: - dim = contact_dim_in[conid] - friction = contact_friction_in[conid] - mu = friction[0] * opt_impratio_invsqrt - - ic0 = map_efc2iefc_in[worldid, contact_efc_address_in[conid, 0]] - D0 = iefc_D_in[worldid, ic0] - ja0 = Jaref_in[worldid, ic0] - jv0 = jv_in[worldid, ic0] - - # Bottom-zone quad for the full contact (scalar quadratic over all rows) - quad = wp.vec3(0.5 * ja0 * ja0 * D0, jv0 * ja0 * D0, 0.5 * jv0 * jv0 * D0) - - u0 = ja0 * mu - v0 = jv0 * mu - uu = float(0.0) - uv = float(0.0) - vv = float(0.0) - - for j in range(1, dim): - icj = map_efc2iefc_in[worldid, contact_efc_address_in[conid, j]] - jaj = Jaref_in[worldid, icj] - jvj = jv_in[worldid, icj] - dj = iefc_D_in[worldid, icj] - DJj = dj * jaj - - quad += wp.vec3(0.5 * jaj * DJj, jvj * DJj, 0.5 * jvj * dj * jvj) - - frictionj = friction[j - 1] - uj = jaj * frictionj - vj = jvj * frictionj - uu += uj * uj - uv += uj * vj - vv += vj * vj - - mu2 = mu * mu - dm = math.safe_div(D0, mu2 * (1.0 + mu2)) - - quad1 = wp.vec3(u0, v0, uu) - quad2 = wp.vec3(uv, vv, dm) - - return _eval_elliptic(mu, quad, quad1, quad2, alpha) - - -# TODO(team): refactor _linesearch_kernel_island @wp.kernel -def _linesearch_kernel_island( - # Model: - opt_tolerance: wp.array[float], - opt_ls_tolerance: wp.array[float], - opt_ls_iterations: int, - opt_impratio_invsqrt: wp.array[float], - stat_meaninertia: wp.array[float], +def _scatter_solution( # Data in: - nefc_in: wp.array[int], - nisland_in: wp.array[int], - contact_friction_in: wp.array[types.vec5], - contact_dim_in: wp.array[int], - contact_efc_address_in: wp.array2d[int], - nidof_in: wp.array[int], - island_nv_in: wp.array2d[int], - island_ne_in: wp.array2d[int], - island_nf_in: wp.array2d[int], - island_efcadr_in: wp.array2d[int], - island_nefc_in: wp.array2d[int], - map_efc2iefc_in: wp.array2d[int], - njmax_in: int, - nacon_in: wp.array[int], - island_idofadr_in: wp.array2d[int], + dof_cdof_in: wp.array2d[int], # In: - iefc_type_in: wp.array2d[int], - iefc_id_in: wp.array2d[int], - iefc_D_in: wp.array2d[float], - iefc_frictionloss_in: wp.array2d[float], - Jaref_in: wp.array2d[float], - jv_in: wp.array2d[float], - mv_in: wp.array2d[float], - search_in: wp.array2d[float], - ifrc_smooth_in: wp.array2d[float], - iMa_in: wp.array2d[float], - island_search_dot_in: wp.array2d[float], - island_gauss_in: wp.array2d[float], - island_done_in: wp.array2d[bool], + x_in: wp.array3d[float], # Out: - island_alpha_out: wp.array2d[float], + vec_out: wp.array2d[float], ): - """Linesearch per island.""" - worldid, islandid = wp.tid() - nisland = nisland_in[worldid] - if islandid >= nisland: - island_alpha_out[worldid, islandid] = 0.0 - return - nefc = wp.min(njmax_in, nefc_in[worldid]) - tolerance = opt_tolerance[worldid % opt_tolerance.shape[0]] - ls_tolerance = opt_ls_tolerance[worldid % opt_ls_tolerance.shape[0]] - meaninertia = stat_meaninertia[worldid % stat_meaninertia.shape[0]] - impratio_invsqrt = opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] - - if island_done_in[worldid, islandid]: - island_alpha_out[worldid, islandid] = 0.0 - return - - iefcadr = island_efcadr_in[worldid, islandid] - ine = island_ne_in[worldid, islandid] - inf = island_nf_in[worldid, islandid] - inv = island_nv_in[worldid, islandid] - idofadr = island_idofadr_in[worldid, islandid] - - # Get island nefc - isle_nefc_end = iefcadr + island_nefc_in[worldid, islandid] - - # Compute gauss quad: [gauss, s.T @ (Ma - frc_smooth), 0.5 * s.T @ mv] - quad_gauss_1 = float(0.0) - quad_gauss_2 = float(0.0) - for i in range(inv): - idof = idofadr + i - s = search_in[worldid, idof] - quad_gauss_1 += s * (iMa_in[worldid, idof] - ifrc_smooth_in[worldid, idof]) - quad_gauss_2 += 0.5 * s * mv_in[worldid, idof] - # Costs are evaluated as deltas from alpha=0 to keep float32 precision on large - # absolute costs, so the constant gauss cost (island_gauss_in) is dropped here. - quad_gauss = wp.vec3(0.0, quad_gauss_1, quad_gauss_2) - - # gtol - snorm = wp.sqrt(island_search_dot_in[worldid, islandid]) - scale = meaninertia * float(inv) - gtol = wp.max(tolerance * ls_tolerance * snorm * scale, 1e-6) - - # p0: cost/grad/hessian at alpha=0 - p0 = wp.vec3(quad_gauss[0], quad_gauss[1], 2.0 * quad_gauss[2]) - for iefcid in range(iefcadr, isle_nefc_end): - if iefcid >= nefc: - break - local_iefcid = iefcid - iefcadr - D = iefc_D_in[worldid, iefcid] - ja = Jaref_in[worldid, iefcid] - jv_val = jv_in[worldid, iefcid] - if local_iefcid < ine: - # Equality: always active - jvD = jv_val * D - p0 += wp.vec3(0.5 * D * ja * ja, jvD * ja, jv_val * jvD) - elif local_iefcid < ine + inf: - # Friction - f = iefc_frictionloss_in[worldid, iefcid] - rf = math.safe_div(f, D) - p0 += _eval_frictionloss_pt(ja, f, rf, jv_val, D) - elif iefc_type_in[worldid, iefcid] == types.ConstraintType.CONTACT_ELLIPTIC: - conid = iefc_id_in[worldid, iefcid] - if conid < nacon_in[0]: - ic0 = map_efc2iefc_in[worldid, contact_efc_address_in[conid, 0]] - if iefcid == ic0: - p0 += _eval_elliptic_cost_island( - impratio_invsqrt, - contact_friction_in, - contact_dim_in, - contact_efc_address_in, - map_efc2iefc_in, - 0.0, - conid, - iefc_D_in, - Jaref_in, - jv_in, - worldid, - ) - else: - # Inequality - if ja < 0.0: - jvD = jv_val * D - p0 += wp.vec3(0.5 * D * ja * ja, jvD * ja, jv_val * jvD) - - # Zero the cost component: alpha=0 is the delta reference (grad/hessian are kept). - p0 = wp.vec3(0.0, p0[1], p0[2]) - - # Newton step: lo_alpha_in = -p0[1] / p0[2] - lo_alpha_in = -math.safe_div(p0[1], p0[2]) - - # Evaluate at Newton step - lo_in = _eval_pt(quad_gauss, lo_alpha_in) - for iefcid in range(iefcadr, isle_nefc_end): - if iefcid >= nefc: - break - local_iefcid = iefcid - iefcadr - D = iefc_D_in[worldid, iefcid] - ja = Jaref_in[worldid, iefcid] - jv_val = jv_in[worldid, iefcid] - if local_iefcid < ine: - lo_in += _eval_pt_direct_shifted(ja, jv_val, D, lo_alpha_in, 0.0) - elif local_iefcid < ine + inf: - f = iefc_frictionloss_in[worldid, iefcid] - rf = math.safe_div(f, D) - x_a = ja + lo_alpha_in * jv_val - lo_in += _shift_cost(_eval_frictionloss_pt(x_a, f, rf, jv_val, D), _eval_frictionloss_cost(ja, f, rf, D)) - elif iefc_type_in[worldid, iefcid] == types.ConstraintType.CONTACT_ELLIPTIC: - conid = iefc_id_in[worldid, iefcid] - if conid < nacon_in[0]: - ic0 = map_efc2iefc_in[worldid, contact_efc_address_in[conid, 0]] - if iefcid == ic0: - cost0 = _eval_elliptic_cost_island( - impratio_invsqrt, - contact_friction_in, - contact_dim_in, - contact_efc_address_in, - map_efc2iefc_in, - 0.0, - conid, - iefc_D_in, - Jaref_in, - jv_in, - worldid, - )[0] - lo_in += _shift_cost( - _eval_elliptic_cost_island( - impratio_invsqrt, - contact_friction_in, - contact_dim_in, - contact_efc_address_in, - map_efc2iefc_in, - lo_alpha_in, - conid, - iefc_D_in, - Jaref_in, - jv_in, - worldid, - ), - cost0, - ) - else: - # Inequality - x_a = ja + lo_alpha_in * jv_val - quad0 = _eval_pt_direct_cost_alpha_zero(ja, D) - cost0 = wp.where(ja < 0.0, quad0, 0.0) - if x_a < 0.0: - lo_in += _eval_pt_direct_shifted(ja, jv_val, D, lo_alpha_in, quad0 - cost0) - else: - lo_in += wp.vec3(-cost0, 0.0, 0.0) - - # Accept Newton step if derivative is small and cost improved - initial_converged = wp.abs(lo_in[1]) < gtol and lo_in[0] < 0.0 - - if initial_converged: - alpha = lo_alpha_in + worldid, i = wp.tid() + ci = dof_cdof_in[worldid, i] + if ci >= 0: + vec_out[worldid, i] = x_in[worldid, ci, 0] else: - alpha = float(0.0) + vec_out[worldid, i] = 0.0 # frozen inactive DOF - # Initialize brackets - lo_less = int(0) - if lo_in[1] < p0[1]: - lo_less = int(1) - if lo_less == 1: - lo = lo_in - lo_alpha = lo_alpha_in - hi = p0 - hi_alpha = float(0.0) - else: - lo = p0 - lo_alpha = float(0.0) - hi = lo_in - hi_alpha = lo_alpha_in - for _iter in range(opt_ls_iterations): - lo_next_alpha = lo_alpha - math.safe_div(lo[1], lo[2]) - hi_next_alpha = hi_alpha - math.safe_div(hi[1], hi[2]) - mid_alpha = 0.5 * (lo_alpha + hi_alpha) +@event_scope +def smooth_solve_compact(m: types.Model, d: types.Data): + """Compacted equivalent of solve_m: qacc_smooth[active] = cM^-1 qfrc_smooth[active]. - # Evaluate at 3 candidate alphas - lo_next = _eval_pt(quad_gauss, lo_next_alpha) - hi_next = _eval_pt(quad_gauss, hi_next_alpha) - mid = _eval_pt(quad_gauss, mid_alpha) - - for iefcid in range(iefcadr, isle_nefc_end): - if iefcid >= nefc: - break - local_iefcid = iefcid - iefcadr - D = iefc_D_in[worldid, iefcid] - ja = Jaref_in[worldid, iefcid] - jv_val = jv_in[worldid, iefcid] - if local_iefcid < ine: - r_lo, r_hi, r_mid = _eval_pt_direct_shifted_3alphas(ja, jv_val, D, lo_next_alpha, hi_next_alpha, mid_alpha, 0.0) - elif local_iefcid < ine + inf: - f = iefc_frictionloss_in[worldid, iefcid] - rf = math.safe_div(f, D) - cost0 = _eval_frictionloss_cost(ja, f, rf, D) - x_lo = ja + lo_next_alpha * jv_val - x_hi = ja + hi_next_alpha * jv_val - x_mid = ja + mid_alpha * jv_val - r_lo = _shift_cost(_eval_frictionloss_pt(x_lo, f, rf, jv_val, D), cost0) - r_hi = _shift_cost(_eval_frictionloss_pt(x_hi, f, rf, jv_val, D), cost0) - r_mid = _shift_cost(_eval_frictionloss_pt(x_mid, f, rf, jv_val, D), cost0) - elif iefc_type_in[worldid, iefcid] == types.ConstraintType.CONTACT_ELLIPTIC: - conid = iefc_id_in[worldid, iefcid] - r_lo = wp.vec3(0.0) - r_hi = wp.vec3(0.0) - r_mid = wp.vec3(0.0) - if conid < nacon_in[0]: - ic0 = map_efc2iefc_in[worldid, contact_efc_address_in[conid, 0]] - if iefcid == ic0: - cost0 = _eval_elliptic_cost_island( - impratio_invsqrt, - contact_friction_in, - contact_dim_in, - contact_efc_address_in, - map_efc2iefc_in, - 0.0, - conid, - iefc_D_in, - Jaref_in, - jv_in, - worldid, - )[0] - r_lo = _shift_cost( - _eval_elliptic_cost_island( - impratio_invsqrt, - contact_friction_in, - contact_dim_in, - contact_efc_address_in, - map_efc2iefc_in, - lo_next_alpha, - conid, - iefc_D_in, - Jaref_in, - jv_in, - worldid, - ), - cost0, - ) - r_hi = _shift_cost( - _eval_elliptic_cost_island( - impratio_invsqrt, - contact_friction_in, - contact_dim_in, - contact_efc_address_in, - map_efc2iefc_in, - hi_next_alpha, - conid, - iefc_D_in, - Jaref_in, - jv_in, - worldid, - ), - cost0, - ) - r_mid = _shift_cost( - _eval_elliptic_cost_island( - impratio_invsqrt, - contact_friction_in, - contact_dim_in, - contact_efc_address_in, - map_efc2iefc_in, - mid_alpha, - conid, - iefc_D_in, - Jaref_in, - jv_in, - worldid, - ), - cost0, - ) - else: - # Inequality - x_lo = ja + lo_next_alpha * jv_val - x_hi = ja + hi_next_alpha * jv_val - x_mid = ja + mid_alpha * jv_val - quad0 = _eval_pt_direct_cost_alpha_zero(ja, D) - cost0 = wp.where(ja < 0.0, quad0, 0.0) - offset = quad0 - cost0 - neg_cost0 = wp.vec3(-cost0, 0.0, 0.0) - r_lo = neg_cost0 - r_hi = neg_cost0 - r_mid = neg_cost0 - if x_lo < 0.0: - r_lo = _eval_pt_direct_shifted(ja, jv_val, D, lo_next_alpha, offset) - if x_hi < 0.0: - r_hi = _eval_pt_direct_shifted(ja, jv_val, D, hi_next_alpha, offset) - if x_mid < 0.0: - r_mid = _eval_pt_direct_shifted(ja, jv_val, D, mid_alpha, offset) - lo_next += r_lo - hi_next += r_hi - mid += r_mid - - # Bracket swapping - swap_lo = int(0) - if _in_bracket(lo, lo_next): - lo = lo_next - lo_alpha = lo_next_alpha - swap_lo = int(1) - if _in_bracket(lo, mid): - lo = mid - lo_alpha = mid_alpha - swap_lo = int(1) - if _in_bracket(lo, hi_next): - lo = hi_next - lo_alpha = hi_next_alpha - swap_lo = int(1) - - swap_hi = int(0) - if _in_bracket(hi, hi_next): - hi = hi_next - hi_alpha = hi_next_alpha - swap_hi = int(1) - if _in_bracket(hi, mid): - hi = mid - hi_alpha = mid_alpha - swap_hi = int(1) - if _in_bracket(hi, lo_next): - hi = lo_next - hi_alpha = lo_next_alpha - swap_hi = int(1) - - # Done check - ls_done = (swap_lo == 0 and swap_hi == 0) or (lo[1] < 0.0 and lo[1] > -gtol) or (hi[1] > 0.0 and hi[1] < gtol) - - # Update alpha if improved - if lo[0] < 0.0 or hi[0] < 0.0: - if lo[0] < hi[0]: - alpha = lo_alpha - else: - alpha = hi_alpha - - if ls_done: - break - - island_alpha_out[worldid, islandid] = alpha - - -@wp.kernel -def _linesearch_qacc_ma_island( - # Data in: - nidof_in: wp.array[int], - # In: - search_in: wp.array2d[float], - mv_in: wp.array2d[float], - island_alpha_in: wp.array2d[float], - idof_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - iacc_out: wp.array2d[float], - iMa_out: wp.array2d[float], -): - """Update iacc and iMa after linesearch.""" - worldid, tid = wp.tid() - - # Process DOFs - if tid < nidof_in[worldid]: - idof = tid - islandid = idof_islandid_in[worldid, idof] - if islandid >= 0 and not island_done_in[worldid, islandid]: - alpha = island_alpha_in[worldid, islandid] - iacc_out[worldid, idof] += alpha * search_in[worldid, idof] - iMa_out[worldid, idof] += alpha * mv_in[worldid, idof] - - -@wp.kernel -def _linesearch_jaref_island( - # Data in: - nefc_in: wp.array[int], - njmax_in: int, - # In: - jv_in: wp.array2d[float], - island_alpha_in: wp.array2d[float], - iefc_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - Jaref_out: wp.array2d[float], -): - """Update Jaref after linesearch.""" - worldid, iefcid = wp.tid() - - if iefcid >= wp.min(njmax_in, nefc_in[worldid]): - return - - islandid = iefc_islandid_in[worldid, iefcid] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - alpha = island_alpha_in[worldid, islandid] - Jaref_out[worldid, iefcid] += alpha * jv_in[worldid, iefcid] - - -@wp.kernel -def _solve_prev_grad_Mgrad_island( - # Data in: - nidof_in: wp.array[int], - # In: - grad_in: wp.array2d[float], - Mgrad_in: wp.array2d[float], - idof_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - prev_grad_out: wp.array2d[float], - prev_Mgrad_out: wp.array2d[float], -): - """Save prev_grad and prev_Mgrad per island DOF.""" - worldid, idofid = wp.tid() - - if idofid >= nidof_in[worldid]: - return - - islandid = idof_islandid_in[worldid, idofid] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - prev_grad_out[worldid, idofid] = grad_in[worldid, idofid] - prev_Mgrad_out[worldid, idofid] = Mgrad_in[worldid, idofid] - - -@wp.kernel -def _solve_beta_island_zero( - # In: - nisland_in: wp.array[int], - # Out: - island_beta_num_out: wp.array2d[float], - island_beta_den_out: wp.array2d[float], -): - """Zero Polak-Ribière numerator and denominator per island.""" - worldid, islandid = wp.tid() - - if islandid >= nisland_in[worldid]: - return - - island_beta_num_out[worldid, islandid] = 0.0 - island_beta_den_out[worldid, islandid] = 0.0 - - -@wp.kernel -def _solve_beta_island_accumulate( - # Data in: - nidof_in: wp.array[int], - # In: - idof_islandid_in: wp.array2d[int], - grad_in: wp.array2d[float], - Mgrad_in: wp.array2d[float], - prev_grad_in: wp.array2d[float], - prev_Mgrad_in: wp.array2d[float], - island_done_in: wp.array2d[bool], - # Out: - island_beta_num_out: wp.array2d[float], - island_beta_den_out: wp.array2d[float], -): - """Parallel Polak-Ribière beta accumulation per island DOF.""" - worldid, idofid = wp.tid() - - if idofid >= nidof_in[worldid]: - return - - islandid = idof_islandid_in[worldid, idofid] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - pMg = prev_Mgrad_in[worldid, idofid] - num = grad_in[worldid, idofid] * (Mgrad_in[worldid, idofid] - pMg) - den = prev_grad_in[worldid, idofid] * pMg - wp.atomic_add(island_beta_num_out, worldid, islandid, num) - wp.atomic_add(island_beta_den_out, worldid, islandid, den) - - -@wp.kernel -def _solve_beta_island_finalize( - # Data in: - nisland_in: wp.array[int], - # In: - island_beta_num_in: wp.array2d[float], - island_beta_den_in: wp.array2d[float], - island_done_in: wp.array2d[bool], - # Out: - island_beta_out: wp.array2d[float], -): - """Finalize Polak-Ribière beta per island.""" - worldid, islandid = wp.tid() - - if islandid >= nisland_in[worldid]: - return - - if island_done_in[worldid, islandid]: - island_beta_out[worldid, islandid] = 0.0 - return - - island_beta_out[worldid, islandid] = wp.max( - 0.0, island_beta_num_in[worldid, islandid] / wp.max(types.MJ_MINVAL, island_beta_den_in[worldid, islandid]) - ) - - -@wp.kernel -def _solve_search_update_island( - # Model: - opt_solver: int, - # Data in: - nidof_in: wp.array[int], - # In: - Mgrad_in: wp.array2d[float], - search_in: wp.array2d[float], - island_beta_in: wp.array2d[float], - idof_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - search_out: wp.array2d[float], - island_search_dot_out: wp.array2d[float], -): - """Update search direction per island DOF.""" - worldid, idofid = wp.tid() - - if idofid >= nidof_in[worldid]: - return - - islandid = idof_islandid_in[worldid, idofid] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - s = -Mgrad_in[worldid, idofid] - if opt_solver == types.SolverType.CG: - s += island_beta_in[worldid, islandid] * search_in[worldid, idofid] - - search_out[worldid, idofid] = s - wp.atomic_add(island_search_dot_out, worldid, islandid, s * s) - - -@wp.kernel -def _solve_done_island( - # Model: - opt_tolerance: wp.array[float], - opt_iterations: int, - stat_meaninertia: wp.array[float], - # Data in: - nisland_in: wp.array[int], - island_nv_in: wp.array2d[int], - # In: - island_grad_dot_in: wp.array2d[float], - island_cost_in: wp.array2d[float], - island_prev_cost_in: wp.array2d[float], - island_done_in: wp.array2d[bool], - # Data out: - solver_niter_out: wp.array[int], - # Out: - island_done_out: wp.array2d[bool], - island_solver_niter_out: wp.array2d[int], - nsolving_out: wp.array[int], -): - """Check convergence per island.""" - worldid, islandid = wp.tid() - - if islandid >= nisland_in[worldid]: - return - - tolerance = opt_tolerance[worldid % opt_tolerance.shape[0]] - meaninertia = stat_meaninertia[worldid % stat_meaninertia.shape[0]] - - if island_done_in[worldid, islandid]: - niter = island_solver_niter_out[worldid, islandid] - wp.atomic_max(solver_niter_out, worldid, niter) - return - - island_solver_niter_out[worldid, islandid] += 1 - niter = island_solver_niter_out[worldid, islandid] - wp.atomic_max(solver_niter_out, worldid, niter) - - inv = island_nv_in[worldid, islandid] - improvement = _rescale(inv, meaninertia, island_prev_cost_in[worldid, islandid] - island_cost_in[worldid, islandid]) - gradient = _rescale(inv, meaninertia, wp.sqrt(island_grad_dot_in[worldid, islandid])) - done = (improvement < tolerance) or (gradient < tolerance) - if done or niter >= opt_iterations: - island_done_out[worldid, islandid] = True - wp.atomic_sub(nsolving_out, 0, 1) - - -@wp.kernel -def _update_gradient_JTDAJ_island( - # Model: - is_sparse: bool, - # Data in: - nefc_in: wp.array[int], - njmax_in: int, - island_idofadr_in: wp.array2d[int], - island_nv_in: wp.array2d[int], - # In: - iefc_J_rownnz_in: wp.array2d[int], - iefc_J_rowadr_in: wp.array2d[int], - iefc_J_colind_in: wp.array3d[int], - iefc_J_in: wp.array3d[float], - iefc_D_in: wp.array2d[float], - iefc_state_in: wp.array2d[int], - iefc_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - ih_out: wp.array3d[float], -): - """Build island Hessian: ih += Jᵀ·D·J for active constraints.""" - worldid, iefcid = wp.tid() - - if iefcid >= wp.min(njmax_in, nefc_in[worldid]): - return - - islandid = iefc_islandid_in[worldid, iefcid] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - state = iefc_state_in[worldid, iefcid] - if state != types.ConstraintState.QUADRATIC.value: - return - D = iefc_D_in[worldid, iefcid] - - idofadr = island_idofadr_in[worldid, islandid] - inv = island_nv_in[worldid, islandid] - if is_sparse: - rownnz = iefc_J_rownnz_in[worldid, iefcid] - rowadr = iefc_J_rowadr_in[worldid, iefcid] - for k1 in range(rownnz): - adr1 = rowadr + k1 - Ji = iefc_J_in[worldid, 0, adr1] - i = iefc_J_colind_in[worldid, 0, adr1] - - for k2 in range(k1 + 1): - adr2 = rowadr + k2 - Jj = iefc_J_in[worldid, 0, adr2] - j = iefc_J_colind_in[worldid, 0, adr2] - - h = Ji * Jj * D - wp.atomic_add(ih_out[worldid, i], j, h) - if i != j: - wp.atomic_add(ih_out[worldid, j], i, h) - else: - for ii in range(inv): - i = idofadr + ii - Ji = iefc_J_in[worldid, iefcid, i] - if Ji == 0.0: - continue - for jj in range(ii + 1): - j = idofadr + jj - Jj = iefc_J_in[worldid, iefcid, j] - if Jj == 0.0: - continue - h = Ji * Jj * D - wp.atomic_add(ih_out[worldid, i], j, h) - if i != j: - wp.atomic_add(ih_out[worldid, j], i, h) - - -@wp.kernel -def _update_gradient_set_h_M_sparse_island( - # Model: - M_fullm_i: wp.array[int], - M_fullm_j: wp.array[int], - M_elemid: wp.array2d[int], - # Data in: - nidof_in: wp.array[int], - M_in: wp.array3d[float], - dof_island_in: wp.array2d[int], - map_dof2idof_in: wp.array2d[int], - # In: - island_done_in: wp.array2d[bool], - # Out: - ih_out: wp.array3d[float], -): - """Add sparse mass matrix to island Hessian using global-to-island DOF mapping.""" - worldid, elementid = wp.tid() - - i_global = M_fullm_i[elementid] - j_global = M_fullm_j[elementid] - - madr = M_elemid[i_global, j_global] - if madr < 0: - return - - # Check both DOFs belong to an island - island_i = dof_island_in[worldid, i_global] - if island_i < 0: - return - if island_done_in[worldid, island_i]: - return - - island_j = dof_island_in[worldid, j_global] - if island_j < 0: - return - - # Both DOFs must be in the same island - if island_i != island_j: - return - - idof_i = map_dof2idof_in[worldid, i_global] - idof_j = map_dof2idof_in[worldid, j_global] - - val = M_in[worldid, 0, madr] - ih_out[worldid, idof_i, idof_j] += val - if idof_i != idof_j: - ih_out[worldid, idof_j, idof_i] += val - - -@wp.kernel -def _update_gradient_set_h_M_dense_island( - # Model: - nv: int, - # Data in: - nidof_in: wp.array[int], - M_in: wp.array3d[float], - map_idof2dof_in: wp.array2d[int], - # In: - idof_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - ih_out: wp.array3d[float], -): - """Add dense mass matrix to island Hessian.""" - worldid, idofid = wp.tid() - - if idofid >= nidof_in[worldid]: - return - - islandid = idof_islandid_in[worldid, idofid] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - dof_i = map_idof2dof_in[worldid, idofid] - - # Copy row from M to ih, mapping columns - nid = nidof_in[worldid] - for jdof in range(nid): - dof_j = map_idof2dof_in[worldid, jdof] - ih_out[worldid, idofid, jdof] += M_in[worldid, dof_i, dof_j] - - -@wp.kernel -def _update_gradient_JTCJ_island( - # Model: - opt_impratio_invsqrt: wp.array[float], - is_sparse: bool, - # Data in: - nacon_in: wp.array[int], - contact_friction_in: wp.array[types.vec5], - contact_dim_in: wp.array[int], - contact_efc_address_in: wp.array2d[int], - contact_worldid_in: wp.array[int], - island_idofadr_in: wp.array2d[int], - naconmax_in: int, - nidof_in: wp.array[int], - map_efc2iefc_in: wp.array2d[int], - island_nv_in: wp.array2d[int], - # In: - iefc_J_rownnz_in: wp.array2d[int], - iefc_J_rowadr_in: wp.array2d[int], - iefc_J_colind_in: wp.array3d[int], - iefc_J_in: wp.array3d[float], - iefc_D_in: wp.array2d[float], - iefc_state_in: wp.array2d[int], - Jaref_in: wp.array2d[float], - iefc_islandid_in: wp.array2d[int], - island_done_in: wp.array2d[bool], - # Out: - ih_out: wp.array3d[float], -): - """Add elliptic cone Hessian correction: Jᵀ·C·J for contacts in CONE state.""" - conid = wp.tid() - - if conid >= wp.min(naconmax_in, nacon_in[0]): - return - - worldid = contact_worldid_in[conid] - condim = contact_dim_in[conid] - - if condim == 1: - return - - efcid0_global = contact_efc_address_in[conid, 0] - if efcid0_global < 0: - return - - ic0 = map_efc2iefc_in[worldid, efcid0_global] - if iefc_state_in[worldid, ic0] != types.ConstraintState.CONE.value: - return - - islandid = iefc_islandid_in[worldid, ic0] - if islandid < 0: - return - if island_done_in[worldid, islandid]: - return - - inv = island_nv_in[worldid, islandid] - idofadr = island_idofadr_in[worldid, islandid] - - fri = contact_friction_in[conid] - mu = fri[0] * opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] - mu2 = mu * mu - dm = math.safe_div(iefc_D_in[worldid, ic0], mu2 * (1.0 + mu2)) - - if dm == 0.0: - return - - # Compute n and u vector - n = Jaref_in[worldid, ic0] * mu - u = types.vec6(n, 0.0, 0.0, 0.0, 0.0, 0.0) - - tt = float(0.0) - for j in range(1, condim): - efcidj_global = contact_efc_address_in[conid, j] - if efcidj_global < 0: - return - icj = map_efc2iefc_in[worldid, efcidj_global] - uj = Jaref_in[worldid, icj] * fri[j - 1] - tt += uj * uj - u[j] = uj - - if tt <= 0.0: - t = 0.0 - else: - t = wp.sqrt(tt) - t = wp.max(t, types.MJ_MINVAL) - ttt = wp.max(t * t * t, types.MJ_MINVAL) - - # Accumulate cone correction into ih - for dim1id in range(condim): - if dim1id == 0: - ic1 = ic0 - else: - efcid1_global = contact_efc_address_in[conid, dim1id] - if efcid1_global < 0: - return - ic1 = map_efc2iefc_in[worldid, efcid1_global] - - ui = u[dim1id] - - for dim2id in range(dim1id + 1): - if dim2id == 0: - ic2 = ic0 - else: - efcid2_global = contact_efc_address_in[conid, dim2id] - if efcid2_global < 0: - return - ic2 = map_efc2iefc_in[worldid, efcid2_global] - - uj = u[dim2id] - - # Cone correction matrix - if dim1id == 0 and dim2id == 0: - hcone = 1.0 - elif dim1id == 0: - hcone = -math.safe_div(mu, t) * uj - elif dim2id == 0: - hcone = -math.safe_div(mu, t) * ui - else: - hcone = mu * math.safe_div(n, ttt) * ui * uj - if dim1id == dim2id: - hcone += mu2 - mu * math.safe_div(n, t) - - # Scale by dm * friction - if dim1id == 0: - fri1 = mu - else: - fri1 = fri[dim1id - 1] - if dim2id == 0: - fri2 = mu - else: - fri2 = fri[dim2id - 1] - - hcone *= dm * fri1 * fri2 - - if hcone == 0.0: - continue - - # Accumulate J1^T * hcone * J2 into ih (lower triangle) - if is_sparse: - if dim1id == dim2id: - rownnz = iefc_J_rownnz_in[worldid, ic1] - rowadr = iefc_J_rowadr_in[worldid, ic1] - for k1 in range(rownnz): - adr1 = rowadr + k1 - J1 = iefc_J_in[worldid, 0, adr1] - i = iefc_J_colind_in[worldid, 0, adr1] - - for k2 in range(k1 + 1): - adr2 = rowadr + k2 - J2 = iefc_J_in[worldid, 0, adr2] - j = iefc_J_colind_in[worldid, 0, adr2] - - val = hcone * J1 * J2 - wp.atomic_add(ih_out[worldid, i], j, val) - if i != j: - wp.atomic_add(ih_out[worldid, j], i, val) - else: - rownnz1 = iefc_J_rownnz_in[worldid, ic1] - rowadr1 = iefc_J_rowadr_in[worldid, ic1] - rownnz2 = iefc_J_rownnz_in[worldid, ic2] - rowadr2 = iefc_J_rowadr_in[worldid, ic2] - - for k1 in range(rownnz1): - adr1 = rowadr1 + k1 - J1 = iefc_J_in[worldid, 0, adr1] - i = iefc_J_colind_in[worldid, 0, adr1] - - for k2 in range(rownnz2): - adr2 = rowadr2 + k2 - J2 = iefc_J_in[worldid, 0, adr2] - j = iefc_J_colind_in[worldid, 0, adr2] - - val = hcone * J1 * J2 - if i == j: - wp.atomic_add(ih_out[worldid, i], j, val * 2.0) - else: - wp.atomic_add(ih_out[worldid, i], j, val) - wp.atomic_add(ih_out[worldid, j], i, val) - else: - for i in range(inv): - J1i = iefc_J_in[worldid, ic1, idofadr + i] - if J1i == 0.0: - continue - for jj in range(i + 1): - J2j = iefc_J_in[worldid, ic2, idofadr + jj] - if J2j == 0.0: - continue - val = hcone * J1i * J2j - wp.atomic_add(ih_out[worldid, idofadr + i], idofadr + jj, val) - if i != jj: - wp.atomic_add(ih_out[worldid, idofadr + jj], idofadr + i, val) - - if dim1id != dim2id: - # Swap-pair contribution: hcone * J[ic2, i] * J[ic1, j]. - # Together with the loop above this gives the full - # hcone * (J[ic1, i] * J[ic2, j] + J[ic2, i] * J[ic1, j]) - # contribution to cell (i, j). - for i in range(inv): - J2i = iefc_J_in[worldid, ic2, idofadr + i] - if J2i == 0.0: - continue - for jj in range(i + 1): - J1j = iefc_J_in[worldid, ic1, idofadr + jj] - if J1j == 0.0: - continue - val = hcone * J2i * J1j - wp.atomic_add(ih_out[worldid, idofadr + i], idofadr + jj, val) - if i != jj: - wp.atomic_add(ih_out[worldid, idofadr + jj], idofadr + i, val) - - -@wp.kernel -def _cholesky_factorize_solve_island( - # Data in: - nisland_in: wp.array[int], - island_idofadr_in: wp.array2d[int], - island_nv_in: wp.array2d[int], - # In: - grad_in: wp.array2d[float], - ih_in: wp.array3d[float], - island_done_in: wp.array2d[bool], - # Out: - Mgrad_out: wp.array2d[float], -): - """Per-island Cholesky factorize and solve: Mgrad = H⁻¹ @ grad. - - One thread per (world, island). Performs dense in-place Cholesky factorization - on the island's inv x inv subblock of ih, then forward/backward substitution. + Inactive DOFs are frozen (qacc_smooth set to 0). Reads the sparse Model inertia. """ - worldid, islandid = wp.tid() - - if islandid >= nisland_in[worldid]: - return - if island_done_in[worldid, islandid]: - return - - inv = island_nv_in[worldid, islandid] - - if inv == 0: - return - - adr = island_idofadr_in[worldid, islandid] - # Cholesky factorization in-place: L such that H = L @ L^T - for i in range(inv): - for j in range(i + 1): - s = ih_in[worldid, adr + i, adr + j] - for k in range(j): - s -= ih_in[worldid, adr + i, adr + k] * ih_in[worldid, adr + j, adr + k] - if i == j: - if s <= 1e-6: - s = 1e-6 - ih_in[worldid, adr + i, adr + j] = wp.sqrt(s) - else: - div = ih_in[worldid, adr + j, adr + j] - ih_in[worldid, adr + i, adr + j] = s / wp.max(1e-6, div) - - # Forward substitution: L @ y = grad => y - for i in range(inv): - s = grad_in[worldid, adr + i] - for k in range(i): - s -= ih_in[worldid, adr + i, adr + k] * Mgrad_out[worldid, adr + k] - Mgrad_out[worldid, adr + i] = s / wp.max(1e-6, ih_in[worldid, adr + i, adr + i]) - - # Backward substitution: L^T @ x = y => x = Mgrad - for i_rev in range(inv): - i = inv - 1 - i_rev - s = Mgrad_out[worldid, adr + i] - for k in range(i + 1, inv): - s -= ih_in[worldid, adr + k, adr + i] * Mgrad_out[worldid, adr + k] - Mgrad_out[worldid, adr + i] = s / wp.max(types.MJ_MINVAL, ih_in[worldid, adr + i, adr + i]) - - -def _init_context_island(m: types.Model, d: types.Data, ctx: IslandSolverContext, nsolving: wp.array): - """Initialize island solver context.""" - # Init per-island scalars - d.solver_niter.zero_() - enable_sleep = bool(m.opt.enableflags & types.EnableBit.SLEEP) wp.launch( - _solve_init_efc_island(enable_sleep), - dim=(d.nworld, m.ntree), - inputs=[ - m.ntree, - d.nisland, - d.tree_awake, - d.tree_island, - ], - outputs=[ctx.cost, ctx.search_dot, ctx.done, ctx.solver_niter, nsolving], + _init_compact_inertia, + dim=(d.nworld, d.nvmax_pad, d.nvmax_pad), + inputs=[d.ncdof], + outputs=[d.cM], ) - - # Jaref = iefc_J @ iacc - iefc_aref wp.launch( - _solve_init_jaref_island, - dim=(d.nworld, d.njmax), - inputs=[ - m.is_sparse, - d.nefc, - d.island_idofadr, - d.island_nv, - d.njmax, - d.efc.iJ_rownnz, - d.efc.iJ_rowadr, - d.efc.iJ_colind, - d.efc.iJ, - d.iqacc, - d.efc.iaref, - d.efc_islandid, - ctx.done, - ], - outputs=[ctx.Jaref], - ) - - # iMa = M @ iacc (all islands in parallel) - support.mul_m_island( - m, - d, - ctx.Ma, - d.iqacc, - d.nidof, - d.map_idof2dof, - d.map_dof2idof, - d.dof_islandid, - ) - - # Update constraint - _update_constraint_island(m, d, ctx) - - # Update gradient - if m.opt.solver == types.SolverType.NEWTON: - _update_gradient_incremental_island(m, d, ctx) - else: - _update_gradient_island(m, d, ctx) - - -@event_scope -def _update_constraint_island(m: types.Model, d: types.Data, ctx: IslandSolverContext): - """Update constraint arrays for island solver.""" - # Save prev cost, zero cost/gauss - wp.launch( - _update_constraint_init_cost, - dim=d.nworld * m.ntree, - inputs=[ctx.cost.reshape(-1), ctx.done.reshape(-1)], - outputs=[ - ctx.gauss.reshape(-1), - ctx.cost.reshape(-1), - ctx.prev_cost.reshape(-1), - ], - ) - - # Compute force, state, cost per EFC - wp.launch( - _update_constraint_efc_island, - dim=(d.nworld, d.njmax), - inputs=[ - m.opt.impratio_invsqrt, - d.nefc, - d.contact.friction, - d.contact.dim, - d.contact.efc_address, - d.island_nefc, - d.island_ne, - d.island_nf, - d.island_efcadr, - d.map_efc2iefc, - d.njmax, - d.nacon, - d.efc.itype, - d.efc.iid, - d.efc.iD, - d.efc.ifrictionloss, - ctx.Jaref, - d.efc_islandid, - ctx.done, - ], - outputs=[d.efc.iforce, d.efc.istate, ctx.cost], - ) - - # qfrc_constraint = J^T @ force - if m.is_sparse: - d.iqfrc_constraint.zero_() - wp.launch( - _update_constraint_init_qfrc_constraint_sparse_island, - dim=(d.nworld, d.njmax), - inputs=[ - d.nefc, - d.njmax, - d.efc.iJ_rownnz, - d.efc.iJ_rowadr, - d.efc.iJ_colind, - d.efc.iJ, - d.efc.iforce, - d.efc_islandid, - ctx.done, - ], - outputs=[d.iqfrc_constraint], - ) - else: - wp.launch( - _update_constraint_init_qfrc_constraint_dense_island, - dim=(d.nworld, m.nv), - inputs=[ - d.nefc, - d.nidof, - d.island_nefc, - d.island_efcadr, - d.njmax, - d.efc.iJ, - d.efc.iforce, - d.dof_islandid, - ctx.done, - ], - outputs=[d.iqfrc_constraint], - ) - - # Gauss cost - wp.launch( - _update_constraint_gauss_cost_island, + _gather_M_sparse, dim=(d.nworld, m.nv), - inputs=[ - d.nidof, - d.iqacc, - d.iqfrc_smooth, - d.iqacc_smooth, - ctx.Ma, - d.dof_islandid, - ctx.done, - ], - outputs=[ctx.gauss, ctx.cost], + inputs=[m.M_rownnz, m.M_rowadr, m.M_colind, d.M, d.dof_cdof], + outputs=[d.cM], ) - - -@event_scope -def _update_gradient_island(m: types.Model, d: types.Data, ctx: IslandSolverContext): - """Update gradient for island solver.""" - # Zero grad_dot per island - ctx.grad_dot.zero_() - - # grad = Ma - frc_smooth - frc_constraint, accumulate grad_dot wp.launch( - _update_gradient_grad_island, - dim=(d.nworld, m.nv), - inputs=[ - d.nidof, - d.iqfrc_smooth, - d.iqfrc_constraint, - ctx.Ma, - d.dof_islandid, - ctx.done, - ], - outputs=[ctx.grad, ctx.grad_dot], + _gather_rhs_compact, + dim=(d.nworld, d.nvmax_pad), + inputs=[d.cdof_dof, d.qfrc_smooth], + outputs=[d.crhs], ) - - # CG preconditioner: Mgrad = M^{-1} @ grad (direct solve) - support.solve_m_island( - m, - d, - ctx.Mgrad, - ctx.grad, - d.nidof, - d.map_idof2dof, + wp.launch_tiled( + _update_gradient_cholesky_blocked(types.TILE_SIZE_JTDAJ_DENSE, d.nvmax_pad, False), + dim=d.nworld, + inputs=[wp.empty(0, dtype=bool), d.crhs, d.cM, d.cqLD], + outputs=[d.cx], + block_dim=m.block_dim.update_gradient_cholesky_blocked, ) + wp.launch(_scatter_solution, dim=(d.nworld, m.nv), inputs=[d.dof_cdof, d.cx], outputs=[d.qacc_smooth]) -@event_scope -def _update_gradient_incremental_island(m: types.Model, d: types.Data, ctx: IslandSolverContext): - """Full Newton gradient update for islands: build H, factorize, solve.""" - # Zero grad_dot per island - ctx.grad_dot.zero_() - - # grad = Ma - frc_smooth - frc_constraint, accumulate grad_dot - wp.launch( - _update_gradient_grad_island, - dim=(d.nworld, m.nv), - inputs=[ - d.nidof, - d.iqfrc_smooth, - d.iqfrc_constraint, - ctx.Ma, - d.dof_islandid, - ctx.done, - ], - outputs=[ctx.grad, ctx.grad_dot], - ) - - # Build H = qM + Jᵀ·D·J - ctx.h.zero_() - - # JTDAJ - wp.launch( - _update_gradient_JTDAJ_island, - dim=(d.nworld, d.njmax), - inputs=[ - m.is_sparse, - d.nefc, - d.njmax, - d.island_idofadr, - d.island_nv, - d.efc.iJ_rownnz, - d.efc.iJ_rowadr, - d.efc.iJ_colind, - d.efc.iJ, - d.efc.iD, - d.efc.istate, - d.efc_islandid, - ctx.done, - ], - outputs=[ctx.h], - ) - - # Add mass matrix - if m.is_sparse: - wp.launch( - _update_gradient_set_h_M_sparse_island, - dim=(d.nworld, m.M_fullm_i.shape[0]), - inputs=[ - m.M_fullm_i, - m.M_fullm_j, - m.M_elemid, - d.nidof, - d.M, - d.dof_island, - d.map_dof2idof, - ctx.done, - ], - outputs=[ctx.h], - ) - else: - wp.launch( - _update_gradient_set_h_M_dense_island, - dim=(d.nworld, m.nv), - inputs=[ - m.nv, - d.nidof, - d.M, - d.map_idof2dof, - d.dof_islandid, - ctx.done, - ], - outputs=[ctx.h], - ) - - # Elliptic cone correction: JTCJ - if m.opt.cone == types.ConeType.ELLIPTIC and d.naconmax > 0: - wp.launch( - _update_gradient_JTCJ_island, - dim=d.naconmax, - inputs=[ - m.opt.impratio_invsqrt, - m.is_sparse, - d.nacon, - d.contact.friction, - d.contact.dim, - d.contact.efc_address, - d.contact.worldid, - d.island_idofadr, - d.naconmax, - d.nidof, - d.map_efc2iefc, - d.island_nv, - d.efc.iJ_rownnz, - d.efc.iJ_rowadr, - d.efc.iJ_colind, - d.efc.iJ, - d.efc.iD, - d.efc.istate, - ctx.Jaref, - d.efc_islandid, - ctx.done, - ], - outputs=[ctx.h], - ) - - # Cholesky factorize and solve: Mgrad = H⁻¹ @ grad - wp.launch( - _cholesky_factorize_solve_island, - dim=(d.nworld, m.ntree), - inputs=[ - d.nisland, - d.island_idofadr, - d.island_nv, - ctx.grad, - ctx.h, - ctx.done, - ], - outputs=[ctx.Mgrad], - ) - - -@event_scope -def _linesearch_island(m: types.Model, d: types.Data, ctx: IslandSolverContext): - """Linesearch for island solver.""" - # mv = M @ search (all islands) - support.mul_m_island( - m, - d, - ctx.mv, - ctx.search, - d.nidof, - d.map_idof2dof, - d.map_dof2idof, - d.dof_islandid, - island_done=ctx.done, - ) - - # jv = J @ search - wp.launch( - _linesearch_jv_island, - dim=(d.nworld, d.njmax), - inputs=[ - m.is_sparse, - d.nefc, - d.nidof, - d.island_idofadr, - d.island_nv, - d.njmax, - d.efc.iJ_rownnz, - d.efc.iJ_rowadr, - d.efc.iJ_colind, - d.efc.iJ, - ctx.search, - d.efc_islandid, - ctx.done, - ], - outputs=[ctx.jv], - ) - - # linesearch - wp.launch( - _linesearch_kernel_island, - dim=(d.nworld, m.ntree), - inputs=[ - m.opt.tolerance, - m.opt.ls_tolerance, - m.opt.ls_iterations, - m.opt.impratio_invsqrt, - m.stat.meaninertia, - d.nefc, - d.nisland, - d.contact.friction, - d.contact.dim, - d.contact.efc_address, - d.nidof, - d.island_nv, - d.island_ne, - d.island_nf, - d.island_efcadr, - d.island_nefc, - d.map_efc2iefc, - d.njmax, - d.nacon, - d.island_idofadr, - d.efc.itype, - d.efc.iid, - d.efc.iD, - d.efc.ifrictionloss, - ctx.Jaref, - ctx.jv, - ctx.mv, - ctx.search, - d.iqfrc_smooth, - ctx.Ma, - ctx.search_dot, - ctx.gauss, - ctx.done, - ], - outputs=[ctx.alpha], - ) - - # Update iacc, iMa - wp.launch( - _linesearch_qacc_ma_island, - dim=(d.nworld, m.nv), - inputs=[ - d.nidof, - ctx.search, - ctx.mv, - ctx.alpha, - d.dof_islandid, - ctx.done, - ], - outputs=[d.iqacc, ctx.Ma], - ) - - # Update Jaref - wp.launch( - _linesearch_jaref_island, - dim=(d.nworld, d.njmax), - inputs=[ - d.nefc, - d.njmax, - ctx.jv, - ctx.alpha, - d.efc_islandid, - ctx.done, - ], - outputs=[ctx.Jaref], - ) - - -@event_scope -def _solver_iteration_island( - m: types.Model, - d: types.Data, - ctx: IslandSolverContext, - nsolving: wp.array[int], +@wp.kernel +def _gather_dof_vecs_compact( + # Data in: + qacc_warmstart_in: wp.array2d[float], + qfrc_smooth_in: wp.array2d[float], + qacc_smooth_in: wp.array2d[float], + cdof_dof_in: wp.array2d[int], + # Out: + qfrc_smooth_c_out: wp.array2d[float], + qacc_smooth_c_out: wp.array2d[float], + qacc_warmstart_c_out: wp.array2d[float], ): - """One iteration of island solver for all islands in parallel.""" - _linesearch_island(m, d, ctx) - - is_newton = m.opt.solver == types.SolverType.NEWTON - is_cg = not is_newton - - # Save prev_grad, prev_Mgrad for CG - if is_cg: - wp.launch( - _solve_prev_grad_Mgrad_island, - dim=(d.nworld, m.nv), - inputs=[d.nidof, ctx.grad, ctx.Mgrad, d.dof_islandid, ctx.done], - outputs=[ctx.prev_grad, ctx.prev_Mgrad], - ) - - # Update constraint - _update_constraint_island(m, d, ctx) - - # Update gradient - if is_newton: - _update_gradient_incremental_island(m, d, ctx) + worldid, ci = wp.tid() + dof = cdof_dof_in[worldid, ci] + if dof >= 0: + qfrc_smooth_c_out[worldid, ci] = qfrc_smooth_in[worldid, dof] + qacc_smooth_c_out[worldid, ci] = qacc_smooth_in[worldid, dof] + qacc_warmstart_c_out[worldid, ci] = qacc_warmstart_in[worldid, dof] else: - _update_gradient_island(m, d, ctx) + qfrc_smooth_c_out[worldid, ci] = 0.0 + qacc_smooth_c_out[worldid, ci] = 0.0 + qacc_warmstart_c_out[worldid, ci] = 0.0 - # Polak-Ribière beta (CG only) - if is_cg: - wp.launch( - _solve_beta_island_zero, - dim=(d.nworld, m.ntree), - inputs=[d.nisland], - outputs=[ctx.beta, ctx.beta_den], - ) - wp.launch( - _solve_beta_island_accumulate, - dim=(d.nworld, m.nv), - inputs=[ - d.nidof, - d.dof_islandid, - ctx.grad, - ctx.Mgrad, - ctx.prev_grad, - ctx.prev_Mgrad, - ctx.done, - ], - outputs=[ctx.beta, ctx.beta_den], - ) - wp.launch( - _solve_beta_island_finalize, - dim=(d.nworld, m.ntree), - inputs=[d.nisland, ctx.beta, ctx.beta_den, ctx.done], - outputs=[ctx.beta], - ) - # Zero search_dot - wp.launch( - _solve_zero_search_dot, - dim=d.nworld * m.ntree, - inputs=[ctx.done.reshape(-1)], - outputs=[ctx.search_dot.reshape(-1)], +@wp.kernel +def _scatter_dof_vecs( + # Data in: + dof_cdof_in: wp.array2d[int], + # In: + qacc_c_in: wp.array2d[float], + qfrc_constraint_c_in: wp.array2d[float], + # Data out: + qacc_out: wp.array2d[float], + qfrc_constraint_out: wp.array2d[float], +): + worldid, i = wp.tid() + ci = dof_cdof_in[worldid, i] + if ci >= 0: + qacc_out[worldid, i] = qacc_c_in[worldid, ci] + qfrc_constraint_out[worldid, i] = qfrc_constraint_c_in[worldid, ci] + else: + qacc_out[worldid, i] = 0.0 + qfrc_constraint_out[worldid, i] = 0.0 + + +@wp.kernel +def _gather_J_sparse( + # Data in: + nefc_in: wp.array[int], + dof_cdof_in: wp.array2d[int], + # In: + J_rownnz_in: wp.array2d[int], + J_rowadr_in: wp.array2d[int], + J_colind_in: wp.array3d[int], + J_in: wp.array3d[float], + # Out: + J_c_out: wp.array3d[float], +): + worldid, efcid = wp.tid() + if efcid >= nefc_in[worldid]: + return + rowadr = J_rowadr_in[worldid, efcid] + for k in range(J_rownnz_in[worldid, efcid]): + adr = rowadr + k + cj = dof_cdof_in[worldid, J_colind_in[worldid, 0, adr]] + if cj >= 0: + J_c_out[worldid, efcid, cj] = J_in[worldid, 0, adr] + + +@wp.kernel +def _gather_J_dense( + # Data in: + nefc_in: wp.array[int], + dof_cdof_in: wp.array2d[int], + # In: + J_in: wp.array3d[float], + # Out: + J_c_out: wp.array3d[float], +): + worldid, efcid = wp.tid() + if efcid >= nefc_in[worldid]: + return + nv = dof_cdof_in.shape[1] + for j in range(nv): + cj = dof_cdof_in[worldid, j] + if cj >= 0: + J_c_out[worldid, efcid, cj] = J_in[worldid, efcid, j] + + +@event_scope +def solve_compact(m: types.Model, d: types.Data): + """Run the dense Newton constraint solver in compacted DOF space. + + Gathers the active-DOF inertia, constraint Jacobian, and smooth/warmstart vectors + into nvmax_pad-sized dense workspaces, runs the stock dense Newton solver on a + shallow-replaced (m, d) at nvmax_pad, then scatters qacc/qfrc_constraint back. + Inactive DOFs are frozen to 0. Reads the sparse Model inertia and constraint J. + """ + _compact_gather(m, d) + + # shallow-replace (m, d) so the stock dense Newton solver runs at nvmax_pad. + # Keep graph-conditional early-exit on CUDA (matches baseline: stops at convergence + # instead of running all iterations); fall back to the plain loop on CPU. + nvp = d.nvmax_pad + gc = m.opt.graph_conditional and wp.get_device().is_cuda + opt2 = dataclasses.replace(m.opt, graph_conditional=gc, tolerance=d.ctol, ls_tolerance=d.cls_tol) + m2 = dataclasses.replace( + m, opt=opt2, nv=nvp, nv_pad=nvp, is_sparse=False, dof_tri_row=d.cdof_tri_row, dof_tri_col=d.cdof_tri_col + ) + efc2 = dataclasses.replace(d.efc, J=d.cJ, Ma=d.cMa) + d2 = dataclasses.replace( + d, + M=d.cM, + qfrc_smooth=d.cqfrc_smooth, + qacc_smooth=d.cqacc_smooth, + qacc_warmstart=d.cqacc_warmstart, + qacc=d.cqacc, + qfrc_constraint=d.cqfrc_constraint, + efc=efc2, ) - # Search update + sctx = _create_solver_context(m2, d2) + _solve(m2, d2, sctx, compact=True) + + _compact_scatter(m, d) + + +@event_scope +def _compact_gather(m: types.Model, d: types.Data): + nvp = d.nvmax_pad + # gather compacted dense inertia (identity-padded tail) wp.launch( - _solve_search_update_island, + _init_compact_inertia, + dim=(d.nworld, nvp, nvp), + inputs=[d.ncdof], + outputs=[d.cM], + ) + + wp.launch( + _gather_M_sparse, dim=(d.nworld, m.nv), + inputs=[m.M_rownnz, m.M_rowadr, m.M_colind, d.M, d.dof_cdof], + outputs=[d.cM], + ) + # gather compacted dense constraint Jacobian (active columns only) + d.cJ.zero_() + if m.is_sparse: + wp.launch( + _gather_J_sparse, + dim=(d.nworld, d.njmax), + inputs=[d.nefc, d.dof_cdof, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, d.efc.J], + outputs=[d.cJ], + ) + else: + wp.launch( + _gather_J_dense, + dim=(d.nworld, d.njmax), + inputs=[d.nefc, d.dof_cdof, d.efc.J], + outputs=[d.cJ], + ) + # gather compacted DOF-space vectors in a single launch + wp.launch( + _gather_dof_vecs_compact, + dim=(d.nworld, nvp), inputs=[ - m.opt.solver, - d.nidof, - ctx.Mgrad, - ctx.search, - ctx.beta, - d.dof_islandid, - ctx.done, + d.qacc_warmstart, + d.qfrc_smooth, + d.qacc_smooth, + d.cdof_dof, + ], + outputs=[ + d.cqfrc_smooth, + d.cqacc_smooth, + d.cqacc_warmstart, ], - outputs=[ctx.search, ctx.search_dot], ) - # Convergence check - d.solver_niter.zero_() + +@event_scope +def _compact_scatter(m: types.Model, d: types.Data): + # scatter results back to full DOF space (inactive frozen to 0) in one launch wp.launch( - _solve_done_island, - dim=(d.nworld, m.ntree), - inputs=[ - m.opt.tolerance, - m.opt.iterations, - m.stat.meaninertia, - d.nisland, - d.island_nv, - ctx.grad_dot, - ctx.cost, - ctx.prev_cost, - ctx.done, - ], - outputs=[d.solver_niter, ctx.done, ctx.solver_niter, nsolving], + _scatter_dof_vecs, + dim=(d.nworld, m.nv), + inputs=[d.dof_cdof, d.cqacc, d.cqfrc_constraint], + outputs=[d.qacc, d.qfrc_constraint], ) + + # Refresh full d.efc.Ma = M @ qacc. The integrators (Euler/implicit damping) use Ma as + # the RHS; the compact solve only populated the compacted Ma, so recompute it in full + # space. Inactive DOFs have qacc=0 so their Ma is 0 and they stay frozen. + support.mul_m(m, d, d.efc.Ma, d.qacc) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py index 97fb1f39..ab7d5c62 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py @@ -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, ], ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py index 941e1507..7884e999 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py @@ -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 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 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml b/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml index eac2e265..29db9ded 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name="mujoco-warp" -version = "3.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", ] diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py b/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py index 5193c9a4..388b9c48 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py @@ -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) diff --git a/mjx/mujoco/mjx/warp/__init__.py b/mjx/mujoco/mjx/warp/__init__.py index 28ebd35f..1270fe0a 100644 --- a/mjx/mujoco/mjx/warp/__init__.py +++ b/mjx/mujoco/mjx/warp/__init__.py @@ -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() diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py index bb020a16..4c061f64 100644 --- a/mjx/mujoco/mjx/warp/collision_driver.py +++ b/mjx/mujoco/mjx/warp/collision_driver.py @@ -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 diff --git a/mjx/mujoco/mjx/warp/ffi.py b/mjx/mujoco/mjx/warp/ffi.py index 9e956991..99ef9b6f 100644 --- a/mjx/mujoco/mjx/warp/ffi.py +++ b/mjx/mujoco/mjx/warp/ffi.py @@ -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' diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index c9e386ea..9440af51 100644 --- a/mjx/mujoco/mjx/warp/forward.py +++ b/mjx/mujoco/mjx/warp/forward.py @@ -46,13 +46,14 @@ _cb = mjwp_types.Callback( **{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init} ) + @ffi.format_args_for_warp def _forward_shim( # Model nworld: int, + M_colind: wp.array[int], M_elemid: wp.array2d[int], - M_fullm_i: wp.array[int], - M_fullm_j: wp.array[int], + M_hinit_i: wp.array[int], M_mulm_col: wp.array[int], M_mulm_madr: wp.array[int], M_mulm_rowadr: wp.array[int], @@ -165,15 +166,23 @@ def _forward_shim( flex_elemdataadr: wp.array[int], flex_elemedge: wp.array[int], flex_elemedgeadr: 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], @@ -189,6 +198,9 @@ def _forward_shim( flexedge_J_rownnz: wp.array[int], flexedge_invweight0: wp.array[float], flexedge_length0: wp.array[float], + 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], @@ -213,6 +225,7 @@ def _forward_shim( geom_solmix: wp.array2d[float], geom_solref: wp.array2d[wp.vec2], geom_type: wp.array[int], + has_flex_selfcollide: bool, has_fluid: bool, has_sdf_geom: bool, hfield_adr: wp.array[int], @@ -248,6 +261,7 @@ def _forward_shim( light_poscom0: wp.array2d[wp.vec3], light_targetbodyid: wp.array[int], mat_rgba: wp.array2d[wp.vec4], + max_flex_dim: int, max_ten_J_rownnz: int, mesh_face: wp.array[wp.vec3i], mesh_faceadr: wp.array[int], @@ -266,6 +280,7 @@ def _forward_shim( mesh_polyvert: wp.array[int], mesh_polyvertadr: wp.array[int], mesh_polyvertnum: wp.array[int], + mesh_pos: wp.array[wp.vec3], mesh_quat: wp.array[wp.quat], mesh_vert: wp.array[wp.vec3], mesh_vertadr: wp.array[int], @@ -280,10 +295,9 @@ def _forward_shim( nflex: int, nflexedge: int, nflexelem: int, - nflexshelldata: int, + nflexevpair: int, nflexvert: int, ngeom: int, - ngravcomp: int, nhistory: int, njnt: int, nlight: int, @@ -291,6 +305,7 @@ def _forward_shim( nmaxmeshdeg: int, nmaxpolygon: int, nmaxpyramid: int, + nmesh: int, nmeshface: int, nrangefinder: int, nsensorcollision: int, @@ -319,7 +334,15 @@ def _forward_shim( plugin: wp.array[int], plugin_attr: wp.array[mjwp_types.vec_pluginattr], qLD_all_updates: wp.array[wp.vec3i], + qLD_block_adr: wp.array[int], + qLD_block_total: int, + qLD_dof_dense: wp.array[int], + qLD_dof_simple: wp.array[int], + qLD_has_dense: bool, + qLD_has_simple: bool, + qLD_has_sparse: bool, qLD_level_offsets: wp.array[int], + qLD_simple_dofs: wp.array[int], qLD_updates: tuple[wp.array[wp.vec3i], ...], qpos0: wp.array2d[float], qpos_spring: wp.array2d[float], @@ -410,7 +433,6 @@ def _forward_shim( opt__graph_conditional: bool, opt__gravity: wp.array[wp.vec3], opt__impratio_invsqrt: wp.array[float], - opt__integrator: int, opt__iterations: int, opt__ls_iterations: int, opt__ls_tolerance: wp.array[float], @@ -422,6 +444,7 @@ def _forward_shim( opt__timestep: wp.array[float], opt__tolerance: wp.array[float], opt__viscosity: wp.array[float], + opt__warn_overflow: bool, opt__wind: wp.array[wp.vec3], stat__meaninertia: wp.array[float], # Data @@ -429,7 +452,9 @@ def _forward_shim( naconmax: int, njmax: int, njmax_nnz: int, - M: wp.array3d[float], + nvmax: int, + nvmax_pad: int, + M: wp.array2d[float], act: wp.array2d[float], act_dot: wp.array2d[float], actuator_force: wp.array2d[float], @@ -438,23 +463,42 @@ def _forward_shim( actuator_velocity: wp.array2d[float], body_awake: wp.array2d[int], body_awake_ind: wp.array2d[int], + cJ: wp.array3d[float], + cM: wp.array3d[float], + cMa: wp.array2d[float], cacc: wp.array2d[wp.spatial_vector], cam_xmat: wp.array2d[wp.mat33], cam_xpos: wp.array2d[wp.vec3], cdof: wp.array2d[wp.spatial_vector], + cdof_dof: wp.array2d[int], cdof_dot: wp.array2d[wp.spatial_vector], + cdof_tri_col: wp.array[int], + cdof_tri_row: wp.array[int], cfrc_ext: wp.array2d[wp.spatial_vector], cfrc_int: wp.array2d[wp.spatial_vector], cinert: wp.array2d[mjwp_types.vec10], + cls_tol: wp.array[float], + cqLD: wp.array3d[float], + cqacc: wp.array2d[float], + cqacc_smooth: wp.array2d[float], + cqacc_warmstart: wp.array2d[float], + cqfrc_constraint: wp.array2d[float], + cqfrc_smooth: wp.array2d[float], crb: wp.array2d[mjwp_types.vec10], + crhs: wp.array3d[float], + ctol: wp.array[float], ctrl: wp.array2d[float], cvel: wp.array2d[wp.spatial_vector], + cx: wp.array3d[float], dof_awake_ind: wp.array2d[int], + dof_cdof: wp.array2d[int], dof_island: wp.array2d[int], dof_islandid: wp.array2d[int], efc_islandid: wp.array2d[int], energy: wp.array[wp.vec2], eq_active: wp.array2d[bool], + flex_aabb_max: wp.array2d[wp.vec3], + flex_aabb_min: wp.array2d[wp.vec3], flexedge_J: wp.array2d[float], flexedge_length: wp.array2d[float], flexedge_velocity: wp.array2d[float], @@ -462,13 +506,9 @@ def _forward_shim( geom_xmat: wp.array2d[wp.mat33], geom_xpos: wp.array2d[wp.vec3], history: wp.array2d[float], - iqacc: wp.array2d[float], - iqacc_smooth: wp.array2d[float], - iqfrc_constraint: wp.array2d[float], - iqfrc_smooth: wp.array2d[float], island_dofadr: wp.array2d[int], - island_efcadr: wp.array2d[int], island_idofadr: wp.array2d[int], + island_iefcadr: wp.array2d[int], island_ne: wp.array2d[int], island_nefc: wp.array2d[int], island_nf: wp.array2d[int], @@ -486,6 +526,7 @@ def _forward_shim( moment_rownnz: wp.array2d[int], nacon: wp.array[int], nbody_awake: wp.array[int], + ncdof: wp.array[int], ncollision: wp.array[int], ne: wp.array[int], nefc: wp.array[int], @@ -495,7 +536,8 @@ def _forward_shim( nl: wp.array[int], ntree_awake: wp.array[int], nv_awake: wp.array[int], - qLD: wp.array3d[float], + overflow: wp.array[int], + qLD: wp.array2d[float], qLDiagInv: wp.array2d[float], qacc: wp.array2d[float], qacc_smooth: wp.array2d[float], @@ -541,6 +583,7 @@ def _forward_shim( 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], @@ -564,19 +607,11 @@ def _forward_shim( efc__aref: wp.array2d[float], efc__force: wp.array2d[float], efc__frictionloss: wp.array2d[float], - efc__iD: wp.array2d[float], - efc__iJ: wp.array3d[float], - efc__iJ_colind: wp.array3d[int], - efc__iJ_rowadr: wp.array2d[int], - efc__iJ_rownnz: wp.array2d[int], - efc__iaref: wp.array2d[float], efc__id: wp.array2d[int], - efc__iforce: wp.array2d[float], - efc__ifrictionloss: wp.array2d[float], - efc__iid: wp.array2d[int], efc__island: wp.array2d[int], - efc__istate: wp.array2d[int], - efc__itype: wp.array2d[int], + efc__jtdaj_adr: wp.array2d[int], + efc__jtdaj_nblock: wp.array[int], + efc__jtdaj_nrow: wp.array2d[int], efc__margin: wp.array2d[float], efc__pos: wp.array2d[float], efc__state: wp.array2d[int], @@ -588,9 +623,9 @@ def _forward_shim( _m.callback = _cb _d.efc = _e _d.contact = _c + _m.M_colind = M_colind _m.M_elemid = M_elemid - _m.M_fullm_i = M_fullm_i - _m.M_fullm_j = M_fullm_j + _m.M_hinit_i = M_hinit_i _m.M_mulm_col = M_mulm_col _m.M_mulm_madr = M_mulm_madr _m.M_mulm_rowadr = M_mulm_rowadr @@ -703,15 +738,23 @@ def _forward_shim( _m.flex_elemdataadr = flex_elemdataadr _m.flex_elemedge = flex_elemedge _m.flex_elemedgeadr = flex_elemedgeadr + _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 @@ -727,6 +770,9 @@ def _forward_shim( _m.flexedge_J_rownnz = flexedge_J_rownnz _m.flexedge_invweight0 = flexedge_invweight0 _m.flexedge_length0 = flexedge_length0 + _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 @@ -751,6 +797,7 @@ def _forward_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_fluid = has_fluid _m.has_sdf_geom = has_sdf_geom _m.hfield_adr = hfield_adr @@ -786,6 +833,7 @@ def _forward_shim( _m.light_poscom0 = light_poscom0 _m.light_targetbodyid = light_targetbodyid _m.mat_rgba = mat_rgba + _m.max_flex_dim = max_flex_dim _m.max_ten_J_rownnz = max_ten_J_rownnz _m.mesh_face = mesh_face _m.mesh_faceadr = mesh_faceadr @@ -804,6 +852,7 @@ def _forward_shim( _m.mesh_polyvert = mesh_polyvert _m.mesh_polyvertadr = mesh_polyvertadr _m.mesh_polyvertnum = mesh_polyvertnum + _m.mesh_pos = mesh_pos _m.mesh_quat = mesh_quat _m.mesh_vert = mesh_vert _m.mesh_vertadr = mesh_vertadr @@ -818,10 +867,9 @@ def _forward_shim( _m.nflex = nflex _m.nflexedge = nflexedge _m.nflexelem = nflexelem - _m.nflexshelldata = nflexshelldata + _m.nflexevpair = nflexevpair _m.nflexvert = nflexvert _m.ngeom = ngeom - _m.ngravcomp = ngravcomp _m.nhistory = nhistory _m.njnt = njnt _m.nlight = nlight @@ -829,6 +877,7 @@ def _forward_shim( _m.nmaxmeshdeg = nmaxmeshdeg _m.nmaxpolygon = nmaxpolygon _m.nmaxpyramid = nmaxpyramid + _m.nmesh = nmesh _m.nmeshface = nmeshface _m.nrangefinder = nrangefinder _m.nsensorcollision = nsensorcollision @@ -859,7 +908,6 @@ def _forward_shim( _m.opt.graph_conditional = opt__graph_conditional _m.opt.gravity = opt__gravity _m.opt.impratio_invsqrt = opt__impratio_invsqrt - _m.opt.integrator = opt__integrator _m.opt.iterations = opt__iterations _m.opt.ls_iterations = opt__ls_iterations _m.opt.ls_tolerance = opt__ls_tolerance @@ -871,6 +919,7 @@ def _forward_shim( _m.opt.timestep = opt__timestep _m.opt.tolerance = opt__tolerance _m.opt.viscosity = opt__viscosity + _m.opt.warn_overflow = opt__warn_overflow _m.opt.wind = opt__wind _m.pair_dim = pair_dim _m.pair_friction = pair_friction @@ -882,7 +931,15 @@ def _forward_shim( _m.plugin = plugin _m.plugin_attr = plugin_attr _m.qLD_all_updates = qLD_all_updates + _m.qLD_block_adr = qLD_block_adr + _m.qLD_block_total = qLD_block_total + _m.qLD_dof_dense = qLD_dof_dense + _m.qLD_dof_simple = qLD_dof_simple + _m.qLD_has_dense = qLD_has_dense + _m.qLD_has_simple = qLD_has_simple + _m.qLD_has_sparse = qLD_has_sparse _m.qLD_level_offsets = qLD_level_offsets + _m.qLD_simple_dofs = qLD_simple_dofs _m.qLD_updates = qLD_updates _m.qpos0 = qpos0 _m.qpos_spring = qpos_spring @@ -971,17 +1028,25 @@ def _forward_shim( _d.actuator_velocity = actuator_velocity _d.body_awake = body_awake _d.body_awake_ind = body_awake_ind + _d.cJ = cJ + _d.cM = cM + _d.cMa = cMa _d.cacc = cacc _d.cam_xmat = cam_xmat _d.cam_xpos = cam_xpos _d.cdof = cdof + _d.cdof_dof = cdof_dof _d.cdof_dot = cdof_dot + _d.cdof_tri_col = cdof_tri_col + _d.cdof_tri_row = cdof_tri_row _d.cfrc_ext = cfrc_ext _d.cfrc_int = cfrc_int _d.cinert = cinert + _d.cls_tol = cls_tol _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 @@ -995,10 +1060,20 @@ def _forward_shim( _d.contact.type = contact__type _d.contact.vert = contact__vert _d.contact.worldid = contact__worldid + _d.cqLD = cqLD + _d.cqacc = cqacc + _d.cqacc_smooth = cqacc_smooth + _d.cqacc_warmstart = cqacc_warmstart + _d.cqfrc_constraint = cqfrc_constraint + _d.cqfrc_smooth = cqfrc_smooth _d.crb = crb + _d.crhs = crhs + _d.ctol = ctol _d.ctrl = ctrl _d.cvel = cvel + _d.cx = cx _d.dof_awake_ind = dof_awake_ind + _d.dof_cdof = dof_cdof _d.dof_island = dof_island _d.dof_islandid = dof_islandid _d.efc.D = efc__D @@ -1011,19 +1086,11 @@ def _forward_shim( _d.efc.aref = efc__aref _d.efc.force = efc__force _d.efc.frictionloss = efc__frictionloss - _d.efc.iD = efc__iD - _d.efc.iJ = efc__iJ - _d.efc.iJ_colind = efc__iJ_colind - _d.efc.iJ_rowadr = efc__iJ_rowadr - _d.efc.iJ_rownnz = efc__iJ_rownnz - _d.efc.iaref = efc__iaref _d.efc.id = efc__id - _d.efc.iforce = efc__iforce - _d.efc.ifrictionloss = efc__ifrictionloss - _d.efc.iid = efc__iid _d.efc.island = efc__island - _d.efc.istate = efc__istate - _d.efc.itype = efc__itype + _d.efc.jtdaj_adr = efc__jtdaj_adr + _d.efc.jtdaj_nblock = efc__jtdaj_nblock + _d.efc.jtdaj_nrow = efc__jtdaj_nrow _d.efc.margin = efc__margin _d.efc.pos = efc__pos _d.efc.state = efc__state @@ -1032,6 +1099,8 @@ def _forward_shim( _d.efc_islandid = efc_islandid _d.energy = energy _d.eq_active = eq_active + _d.flex_aabb_max = flex_aabb_max + _d.flex_aabb_min = flex_aabb_min _d.flexedge_J = flexedge_J _d.flexedge_length = flexedge_length _d.flexedge_velocity = flexedge_velocity @@ -1039,13 +1108,9 @@ def _forward_shim( _d.geom_xmat = geom_xmat _d.geom_xpos = geom_xpos _d.history = history - _d.iqacc = iqacc - _d.iqacc_smooth = iqacc_smooth - _d.iqfrc_constraint = iqfrc_constraint - _d.iqfrc_smooth = iqfrc_smooth _d.island_dofadr = island_dofadr - _d.island_efcadr = island_efcadr _d.island_idofadr = island_idofadr + _d.island_iefcadr = island_iefcadr _d.island_ne = island_ne _d.island_nefc = island_nefc _d.island_nf = island_nf @@ -1065,6 +1130,7 @@ def _forward_shim( _d.nacon = nacon _d.naconmax = naconmax _d.nbody_awake = nbody_awake + _d.ncdof = ncdof _d.ncollision = ncollision _d.ne = ne _d.nefc = nefc @@ -1076,6 +1142,9 @@ def _forward_shim( _d.nl = nl _d.ntree_awake = ntree_awake _d.nv_awake = nv_awake + _d.nvmax = nvmax + _d.nvmax_pad = nvmax_pad + _d.overflow = overflow _d.qLD = qLD _d.qLDiagInv = qLDiagInv _d.qacc = qacc @@ -1133,21 +1202,32 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'actuator_velocity': d._impl.actuator_velocity.shape, 'body_awake': d._impl.body_awake.shape, 'body_awake_ind': d._impl.body_awake_ind.shape, + 'cJ': d._impl.cJ.shape, + 'cM': d._impl.cM.shape, 'cacc': d._impl.cacc.shape, 'cam_xmat': d.cam_xmat.shape, 'cam_xpos': d.cam_xpos.shape, 'cdof': d.cdof.shape, + 'cdof_dof': d._impl.cdof_dof.shape, 'cdof_dot': d.cdof_dot.shape, 'cfrc_ext': d._impl.cfrc_ext.shape, 'cfrc_int': d._impl.cfrc_int.shape, 'cinert': d._impl.cinert.shape, + 'cqacc_smooth': d._impl.cqacc_smooth.shape, + 'cqacc_warmstart': d._impl.cqacc_warmstart.shape, + 'cqfrc_smooth': d._impl.cqfrc_smooth.shape, 'crb': d._impl.crb.shape, + 'crhs': d._impl.crhs.shape, 'cvel': d.cvel.shape, + 'cx': d._impl.cx.shape, 'dof_awake_ind': d._impl.dof_awake_ind.shape, + 'dof_cdof': d._impl.dof_cdof.shape, 'dof_island': d._impl.dof_island.shape, 'dof_islandid': d._impl.dof_islandid.shape, 'efc_islandid': d._impl.efc_islandid.shape, 'energy': d._impl.energy.shape, + 'flex_aabb_max': d._impl.flex_aabb_max.shape, + 'flex_aabb_min': d._impl.flex_aabb_min.shape, 'flexedge_J': d._impl.flexedge_J.shape, 'flexedge_length': d._impl.flexedge_length.shape, 'flexedge_velocity': d._impl.flexedge_velocity.shape, @@ -1155,13 +1235,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'geom_xmat': d.geom_xmat.shape, 'geom_xpos': d.geom_xpos.shape, 'history': d.history.shape, - 'iqacc': d._impl.iqacc.shape, - 'iqacc_smooth': d._impl.iqacc_smooth.shape, - 'iqfrc_constraint': d._impl.iqfrc_constraint.shape, - 'iqfrc_smooth': d._impl.iqfrc_smooth.shape, 'island_dofadr': d._impl.island_dofadr.shape, - 'island_efcadr': d._impl.island_efcadr.shape, 'island_idofadr': d._impl.island_idofadr.shape, + 'island_iefcadr': d._impl.island_iefcadr.shape, 'island_ne': d._impl.island_ne.shape, 'island_nefc': d._impl.island_nefc.shape, 'island_nf': d._impl.island_nf.shape, @@ -1177,6 +1253,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'moment_rownnz': d._impl.moment_rownnz.shape, 'nacon': d._impl.nacon.shape, 'nbody_awake': d._impl.nbody_awake.shape, + 'ncdof': d._impl.ncdof.shape, 'ncollision': d._impl.ncollision.shape, 'ne': d._impl.ne.shape, 'nefc': d._impl.nefc.shape, @@ -1186,6 +1263,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'nl': d._impl.nl.shape, 'ntree_awake': d._impl.ntree_awake.shape, 'nv_awake': d._impl.nv_awake.shape, + 'overflow': d._impl.overflow.shape, 'qLD': d._impl.qLD.shape, 'qLDiagInv': d._impl.qLDiagInv.shape, 'qacc': d.qacc.shape, @@ -1227,6 +1305,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): '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, @@ -1250,19 +1329,11 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'efc__aref': d._impl.efc__aref.shape, 'efc__force': d._impl.efc__force.shape, 'efc__frictionloss': d._impl.efc__frictionloss.shape, - 'efc__iD': d._impl.efc__iD.shape, - 'efc__iJ': d._impl.efc__iJ.shape, - 'efc__iJ_colind': d._impl.efc__iJ_colind.shape, - 'efc__iJ_rowadr': d._impl.efc__iJ_rowadr.shape, - 'efc__iJ_rownnz': d._impl.efc__iJ_rownnz.shape, - 'efc__iaref': d._impl.efc__iaref.shape, 'efc__id': d._impl.efc__id.shape, - 'efc__iforce': d._impl.efc__iforce.shape, - 'efc__ifrictionloss': d._impl.efc__ifrictionloss.shape, - 'efc__iid': d._impl.efc__iid.shape, 'efc__island': d._impl.efc__island.shape, - 'efc__istate': d._impl.efc__istate.shape, - 'efc__itype': d._impl.efc__itype.shape, + 'efc__jtdaj_adr': d._impl.efc__jtdaj_adr.shape, + 'efc__jtdaj_nblock': d._impl.efc__jtdaj_nblock.shape, + 'efc__jtdaj_nrow': d._impl.efc__jtdaj_nrow.shape, 'efc__margin': d._impl.efc__margin.shape, 'efc__pos': d._impl.efc__pos.shape, 'efc__state': d._impl.efc__state.shape, @@ -1271,7 +1342,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): } jf = ffi.jax_callable_variadic_tuple( _forward_shim, - num_outputs=143, + num_outputs=145, output_dims=output_dims, vmap_method=None, in_out_argnames=set([ @@ -1283,21 +1354,32 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'actuator_velocity', 'body_awake', 'body_awake_ind', + 'cJ', + 'cM', 'cacc', 'cam_xmat', 'cam_xpos', 'cdof', + 'cdof_dof', 'cdof_dot', 'cfrc_ext', 'cfrc_int', 'cinert', + 'cqacc_smooth', + 'cqacc_warmstart', + 'cqfrc_smooth', 'crb', + 'crhs', 'cvel', + 'cx', 'dof_awake_ind', + 'dof_cdof', 'dof_island', 'dof_islandid', 'efc_islandid', 'energy', + 'flex_aabb_max', + 'flex_aabb_min', 'flexedge_J', 'flexedge_length', 'flexedge_velocity', @@ -1305,13 +1387,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'geom_xmat', 'geom_xpos', 'history', - 'iqacc', - 'iqacc_smooth', - 'iqfrc_constraint', - 'iqfrc_smooth', 'island_dofadr', - 'island_efcadr', 'island_idofadr', + 'island_iefcadr', 'island_ne', 'island_nefc', 'island_nf', @@ -1327,6 +1405,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'moment_rownnz', 'nacon', 'nbody_awake', + 'ncdof', 'ncollision', 'ne', 'nefc', @@ -1336,6 +1415,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'nl', 'ntree_awake', 'nv_awake', + 'overflow', 'qLD', 'qLDiagInv', 'qacc', @@ -1377,6 +1457,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'contact__dim', 'contact__dist', 'contact__efc_address', + 'contact__elem', 'contact__flex', 'contact__frame', 'contact__friction', @@ -1400,19 +1481,11 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'efc__aref', 'efc__force', 'efc__frictionloss', - 'efc__iD', - 'efc__iJ', - 'efc__iJ_colind', - 'efc__iJ_rowadr', - 'efc__iJ_rownnz', - 'efc__iaref', 'efc__id', - 'efc__iforce', - 'efc__ifrictionloss', - 'efc__iid', 'efc__island', - 'efc__istate', - 'efc__itype', + 'efc__jtdaj_adr', + 'efc__jtdaj_nblock', + 'efc__jtdaj_nrow', 'efc__margin', 'efc__pos', 'efc__state', @@ -1603,9 +1676,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): ) out = jf( d.qpos.shape[0], + m.M_colind, m._impl.M_elemid, - m._impl.M_fullm_i, - m._impl.M_fullm_j, + m._impl.M_hinit_i, m._impl.M_mulm_col, m._impl.M_mulm_madr, m._impl.M_mulm_rowadr, @@ -1718,15 +1791,23 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.flex_elemdataadr, m._impl.flex_elemedge, m._impl.flex_elemedgeadr, + 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, @@ -1742,6 +1823,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.flexedge_J_rownnz, m._impl.flexedge_invweight0, m._impl.flexedge_length0, + 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, @@ -1766,6 +1850,7 @@ def _forward_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_fluid, m._impl.has_sdf_geom, m.hfield_adr, @@ -1801,6 +1886,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.light_poscom0, m._impl.light_targetbodyid, m.mat_rgba, + m._impl.max_flex_dim, m._impl.max_ten_J_rownnz, m.mesh_face, m.mesh_faceadr, @@ -1819,6 +1905,7 @@ def _forward_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_quat, m.mesh_vert, m.mesh_vertadr, @@ -1833,10 +1920,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.nflex, m._impl.nflexedge, m._impl.nflexelem, - m._impl.nflexshelldata, + m._impl.nflexevpair, m._impl.nflexvert, m.ngeom, - m.ngravcomp, m.nhistory, m.njnt, m.nlight, @@ -1844,6 +1930,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.nmaxmeshdeg, m._impl.nmaxpolygon, m._impl.nmaxpyramid, + m.nmesh, m.nmeshface, m._impl.nrangefinder, m._impl.nsensorcollision, @@ -1872,7 +1959,15 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.plugin, m._impl.plugin_attr, m._impl.qLD_all_updates, + m._impl.qLD_block_adr, + m._impl.qLD_block_total, + m._impl.qLD_dof_dense, + m._impl.qLD_dof_simple, + m._impl.qLD_has_dense, + m._impl.qLD_has_simple, + m._impl.qLD_has_sparse, m._impl.qLD_level_offsets, + m._impl.qLD_simple_dofs, m._impl.qLD_updates, m.qpos0, m.qpos_spring, @@ -1963,7 +2058,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.opt._impl.graph_conditional, m.opt.gravity, m.opt._impl.impratio_invsqrt, - m.opt.integrator, m.opt.iterations, m.opt.ls_iterations, m.opt.ls_tolerance, @@ -1975,12 +2069,15 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.opt.timestep, m.opt.tolerance, m.opt.viscosity, + m.opt._impl.warn_overflow, m.opt.wind, m.stat.meaninertia, d._impl.naccdmax, d._impl.naconmax, d._impl.njmax, d._impl.njmax_nnz, + d._impl.nvmax, + d._impl.nvmax_pad, d._impl.M, d.act, d.act_dot, @@ -1990,23 +2087,42 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.actuator_velocity, d._impl.body_awake, d._impl.body_awake_ind, + d._impl.cJ, + d._impl.cM, + d._impl.cMa, d._impl.cacc, d.cam_xmat, d.cam_xpos, d.cdof, + d._impl.cdof_dof, d.cdof_dot, + d._impl.cdof_tri_col, + d._impl.cdof_tri_row, d._impl.cfrc_ext, d._impl.cfrc_int, d._impl.cinert, + d._impl.cls_tol, + d._impl.cqLD, + d._impl.cqacc, + d._impl.cqacc_smooth, + d._impl.cqacc_warmstart, + d._impl.cqfrc_constraint, + d._impl.cqfrc_smooth, d._impl.crb, + d._impl.crhs, + d._impl.ctol, d.ctrl, d.cvel, + d._impl.cx, d._impl.dof_awake_ind, + d._impl.dof_cdof, d._impl.dof_island, d._impl.dof_islandid, d._impl.efc_islandid, d._impl.energy, d.eq_active, + d._impl.flex_aabb_max, + d._impl.flex_aabb_min, d._impl.flexedge_J, d._impl.flexedge_length, d._impl.flexedge_velocity, @@ -2014,13 +2130,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d.geom_xmat, d.geom_xpos, d.history, - d._impl.iqacc, - d._impl.iqacc_smooth, - d._impl.iqfrc_constraint, - d._impl.iqfrc_smooth, d._impl.island_dofadr, - d._impl.island_efcadr, d._impl.island_idofadr, + d._impl.island_iefcadr, d._impl.island_ne, d._impl.island_nefc, d._impl.island_nf, @@ -2038,6 +2150,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.moment_rownnz, d._impl.nacon, d._impl.nbody_awake, + d._impl.ncdof, d._impl.ncollision, d._impl.ne, d._impl.nefc, @@ -2047,6 +2160,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.nl, d._impl.ntree_awake, d._impl.nv_awake, + d._impl.overflow, d._impl.qLD, d._impl.qLDiagInv, d.qacc, @@ -2093,6 +2207,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 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, @@ -2116,19 +2231,11 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.efc__aref, d._impl.efc__force, d._impl.efc__frictionloss, - d._impl.efc__iD, - d._impl.efc__iJ, - d._impl.efc__iJ_colind, - d._impl.efc__iJ_rowadr, - d._impl.efc__iJ_rownnz, - d._impl.efc__iaref, d._impl.efc__id, - d._impl.efc__iforce, - d._impl.efc__ifrictionloss, - d._impl.efc__iid, d._impl.efc__island, - d._impl.efc__istate, - d._impl.efc__itype, + d._impl.efc__jtdaj_adr, + d._impl.efc__jtdaj_nblock, + d._impl.efc__jtdaj_nrow, d._impl.efc__margin, d._impl.efc__pos, d._impl.efc__state, @@ -2144,141 +2251,143 @@ def _forward_jax_impl(m: types.Model, d: types.Data): '_impl.actuator_velocity': out[5], '_impl.body_awake': out[6], '_impl.body_awake_ind': out[7], - '_impl.cacc': out[8], - 'cam_xmat': out[9], - 'cam_xpos': out[10], - 'cdof': out[11], - 'cdof_dot': out[12], - '_impl.cfrc_ext': out[13], - '_impl.cfrc_int': out[14], - '_impl.cinert': out[15], - '_impl.crb': out[16], - 'cvel': out[17], - '_impl.dof_awake_ind': out[18], - '_impl.dof_island': out[19], - '_impl.dof_islandid': out[20], - '_impl.efc_islandid': out[21], - '_impl.energy': out[22], - '_impl.flexedge_J': out[23], - '_impl.flexedge_length': out[24], - '_impl.flexedge_velocity': out[25], - '_impl.flexvert_xpos': out[26], - 'geom_xmat': out[27], - 'geom_xpos': out[28], - 'history': out[29], - '_impl.iqacc': out[30], - '_impl.iqacc_smooth': out[31], - '_impl.iqfrc_constraint': out[32], - '_impl.iqfrc_smooth': out[33], - '_impl.island_dofadr': out[34], - '_impl.island_efcadr': out[35], - '_impl.island_idofadr': out[36], - '_impl.island_ne': out[37], - '_impl.island_nefc': out[38], - '_impl.island_nf': out[39], - '_impl.island_nv': out[40], - '_impl.light_xdir': out[41], - '_impl.light_xpos': out[42], - '_impl.map_dof2idof': out[43], - '_impl.map_efc2iefc': out[44], - '_impl.map_idof2dof': out[45], - '_impl.map_iefc2efc': out[46], - '_impl.moment_colind': out[47], - '_impl.moment_rowadr': out[48], - '_impl.moment_rownnz': out[49], - '_impl.nacon': out[50], - '_impl.nbody_awake': out[51], - '_impl.ncollision': out[52], - '_impl.ne': out[53], - '_impl.nefc': out[54], - '_impl.nf': out[55], - '_impl.nidof': out[56], - '_impl.nisland': out[57], - '_impl.nl': out[58], - '_impl.ntree_awake': out[59], - '_impl.nv_awake': out[60], - '_impl.qLD': out[61], - '_impl.qLDiagInv': out[62], - 'qacc': out[63], - 'qacc_smooth': out[64], - 'qfrc_actuator': out[65], - 'qfrc_bias': out[66], - 'qfrc_constraint': out[67], - '_impl.qfrc_damper': out[68], - 'qfrc_fluid': out[69], - 'qfrc_gravcomp': out[70], - 'qfrc_passive': out[71], - 'qfrc_smooth': out[72], - '_impl.qfrc_spring': out[73], - 'qvel': out[74], - 'sensordata': out[75], - 'site_xmat': out[76], - 'site_xpos': out[77], - '_impl.solver_niter': out[78], - '_impl.subtree_angmom': out[79], - 'subtree_com': out[80], - '_impl.subtree_linvel': out[81], - '_impl.ten_J': out[82], - 'ten_length': out[83], - '_impl.ten_velocity': out[84], - '_impl.ten_wrapadr': out[85], - '_impl.ten_wrapnum': out[86], - '_impl.tree_asleep': out[87], - '_impl.tree_awake': out[88], - '_impl.tree_island': out[89], - '_impl.wrap_obj': out[90], - '_impl.wrap_xpos': out[91], - 'xanchor': out[92], - 'xaxis': out[93], - 'ximat': out[94], - 'xipos': out[95], - 'xmat': out[96], - 'xpos': out[97], - 'xquat': out[98], - '_impl.contact__dim': out[99], - '_impl.contact__dist': out[100], - '_impl.contact__efc_address': out[101], - '_impl.contact__flex': out[102], - '_impl.contact__frame': out[103], - '_impl.contact__friction': out[104], - '_impl.contact__geom': out[105], - '_impl.contact__geomcollisionid': out[106], - '_impl.contact__includemargin': out[107], - '_impl.contact__pos': out[108], - '_impl.contact__solimp': out[109], - '_impl.contact__solref': out[110], - '_impl.contact__solreffriction': out[111], - '_impl.contact__type': out[112], - '_impl.contact__vert': out[113], - '_impl.contact__worldid': out[114], - '_impl.efc__D': out[115], - '_impl.efc__J': out[116], - '_impl.efc__J_colind': out[117], - '_impl.efc__J_rowadr': out[118], - '_impl.efc__J_rownnz': out[119], - '_impl.efc__Jqvel': out[120], - '_impl.efc__Ma': out[121], - '_impl.efc__aref': out[122], - '_impl.efc__force': out[123], - '_impl.efc__frictionloss': out[124], - '_impl.efc__iD': out[125], - '_impl.efc__iJ': out[126], - '_impl.efc__iJ_colind': out[127], - '_impl.efc__iJ_rowadr': out[128], - '_impl.efc__iJ_rownnz': out[129], - '_impl.efc__iaref': out[130], - '_impl.efc__id': out[131], - '_impl.efc__iforce': out[132], - '_impl.efc__ifrictionloss': out[133], - '_impl.efc__iid': out[134], - '_impl.efc__island': out[135], - '_impl.efc__istate': out[136], - '_impl.efc__itype': out[137], - '_impl.efc__margin': out[138], - '_impl.efc__pos': out[139], - '_impl.efc__state': out[140], - '_impl.efc__type': out[141], - '_impl.efc__vel': out[142], + '_impl.cJ': out[8], + '_impl.cM': out[9], + '_impl.cacc': out[10], + 'cam_xmat': out[11], + 'cam_xpos': out[12], + 'cdof': out[13], + '_impl.cdof_dof': out[14], + 'cdof_dot': out[15], + '_impl.cfrc_ext': out[16], + '_impl.cfrc_int': out[17], + '_impl.cinert': out[18], + '_impl.cqacc_smooth': out[19], + '_impl.cqacc_warmstart': out[20], + '_impl.cqfrc_smooth': out[21], + '_impl.crb': out[22], + '_impl.crhs': out[23], + 'cvel': out[24], + '_impl.cx': out[25], + '_impl.dof_awake_ind': out[26], + '_impl.dof_cdof': out[27], + '_impl.dof_island': out[28], + '_impl.dof_islandid': out[29], + '_impl.efc_islandid': out[30], + '_impl.energy': out[31], + '_impl.flex_aabb_max': out[32], + '_impl.flex_aabb_min': out[33], + '_impl.flexedge_J': out[34], + '_impl.flexedge_length': out[35], + '_impl.flexedge_velocity': out[36], + '_impl.flexvert_xpos': out[37], + 'geom_xmat': out[38], + 'geom_xpos': out[39], + 'history': out[40], + '_impl.island_dofadr': out[41], + '_impl.island_idofadr': out[42], + '_impl.island_iefcadr': out[43], + '_impl.island_ne': out[44], + '_impl.island_nefc': out[45], + '_impl.island_nf': out[46], + '_impl.island_nv': out[47], + '_impl.light_xdir': out[48], + '_impl.light_xpos': out[49], + '_impl.map_dof2idof': out[50], + '_impl.map_efc2iefc': out[51], + '_impl.map_idof2dof': out[52], + '_impl.map_iefc2efc': out[53], + '_impl.moment_colind': out[54], + '_impl.moment_rowadr': out[55], + '_impl.moment_rownnz': out[56], + '_impl.nacon': out[57], + '_impl.nbody_awake': out[58], + '_impl.ncdof': out[59], + '_impl.ncollision': out[60], + '_impl.ne': out[61], + '_impl.nefc': out[62], + '_impl.nf': out[63], + '_impl.nidof': out[64], + '_impl.nisland': out[65], + '_impl.nl': out[66], + '_impl.ntree_awake': out[67], + '_impl.nv_awake': out[68], + '_impl.overflow': out[69], + '_impl.qLD': out[70], + '_impl.qLDiagInv': out[71], + 'qacc': out[72], + 'qacc_smooth': out[73], + 'qfrc_actuator': out[74], + 'qfrc_bias': out[75], + 'qfrc_constraint': out[76], + '_impl.qfrc_damper': out[77], + 'qfrc_fluid': out[78], + 'qfrc_gravcomp': out[79], + 'qfrc_passive': out[80], + 'qfrc_smooth': out[81], + '_impl.qfrc_spring': out[82], + 'qvel': out[83], + 'sensordata': out[84], + 'site_xmat': out[85], + 'site_xpos': out[86], + '_impl.solver_niter': out[87], + '_impl.subtree_angmom': out[88], + 'subtree_com': out[89], + '_impl.subtree_linvel': out[90], + '_impl.ten_J': out[91], + 'ten_length': out[92], + '_impl.ten_velocity': out[93], + '_impl.ten_wrapadr': out[94], + '_impl.ten_wrapnum': out[95], + '_impl.tree_asleep': out[96], + '_impl.tree_awake': out[97], + '_impl.tree_island': out[98], + '_impl.wrap_obj': out[99], + '_impl.wrap_xpos': out[100], + 'xanchor': out[101], + 'xaxis': out[102], + 'ximat': out[103], + 'xipos': out[104], + 'xmat': out[105], + 'xpos': out[106], + 'xquat': out[107], + '_impl.contact__dim': out[108], + '_impl.contact__dist': out[109], + '_impl.contact__efc_address': out[110], + '_impl.contact__elem': out[111], + '_impl.contact__flex': out[112], + '_impl.contact__frame': out[113], + '_impl.contact__friction': out[114], + '_impl.contact__geom': out[115], + '_impl.contact__geomcollisionid': out[116], + '_impl.contact__includemargin': out[117], + '_impl.contact__pos': out[118], + '_impl.contact__solimp': out[119], + '_impl.contact__solref': out[120], + '_impl.contact__solreffriction': out[121], + '_impl.contact__type': out[122], + '_impl.contact__vert': out[123], + '_impl.contact__worldid': out[124], + '_impl.efc__D': out[125], + '_impl.efc__J': out[126], + '_impl.efc__J_colind': out[127], + '_impl.efc__J_rowadr': out[128], + '_impl.efc__J_rownnz': out[129], + '_impl.efc__Jqvel': out[130], + '_impl.efc__Ma': out[131], + '_impl.efc__aref': out[132], + '_impl.efc__force': out[133], + '_impl.efc__frictionloss': out[134], + '_impl.efc__id': out[135], + '_impl.efc__island': out[136], + '_impl.efc__jtdaj_adr': out[137], + '_impl.efc__jtdaj_nblock': out[138], + '_impl.efc__jtdaj_nrow': out[139], + '_impl.efc__margin': out[140], + '_impl.efc__pos': out[141], + '_impl.efc__state': out[142], + '_impl.efc__type': out[143], + '_impl.efc__vel': out[144], }) return d @@ -2304,9 +2413,11 @@ def _step_shim( D_diag: wp.array[int], D_rowadr: wp.array[int], D_rownnz: wp.array[int], + M_colind: wp.array[int], M_elemid: wp.array2d[int], M_fullm_i: wp.array[int], M_fullm_j: wp.array[int], + M_hinit_i: wp.array[int], M_mulm_col: wp.array[int], M_mulm_madr: wp.array[int], M_mulm_rowadr: wp.array[int], @@ -2343,7 +2454,9 @@ def _step_shim( body_branches: wp.array[int], body_dofadr: wp.array[int], body_dofnum: wp.array[int], + body_fluid_box_adr: wp.array[int], body_fluid_ellipsoid: wp.array[bool], + body_fluid_ellipsoid_adr: wp.array[int], body_geomadr: wp.array[int], body_geomnum: wp.array[int], body_gravcomp: wp.array2d[float], @@ -2419,15 +2532,23 @@ def _step_shim( flex_elemdataadr: wp.array[int], flex_elemedge: wp.array[int], flex_elemedgeadr: 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], @@ -2443,6 +2564,9 @@ def _step_shim( flexedge_J_rownnz: wp.array[int], flexedge_invweight0: wp.array[float], flexedge_length0: wp.array[float], + 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], @@ -2467,6 +2591,7 @@ def _step_shim( geom_solmix: wp.array2d[float], geom_solref: wp.array2d[wp.vec2], geom_type: wp.array[int], + has_flex_selfcollide: bool, has_fluid: bool, has_sdf_geom: bool, hfield_adr: wp.array[int], @@ -2503,6 +2628,7 @@ def _step_shim( light_targetbodyid: wp.array[int], mapM2D: wp.array[int], mat_rgba: wp.array2d[wp.vec4], + max_flex_dim: int, max_ten_J_rownnz: int, mesh_face: wp.array[wp.vec3i], mesh_faceadr: wp.array[int], @@ -2521,6 +2647,7 @@ def _step_shim( mesh_polyvert: wp.array[int], mesh_polyvertadr: wp.array[int], mesh_polyvertnum: wp.array[int], + mesh_pos: wp.array[wp.vec3], mesh_quat: wp.array[wp.quat], mesh_vert: wp.array[wp.vec3], mesh_vertadr: wp.array[int], @@ -2537,10 +2664,9 @@ def _step_shim( nflex: int, nflexedge: int, nflexelem: int, - nflexshelldata: int, + nflexevpair: int, nflexvert: int, ngeom: int, - ngravcomp: int, nhistory: int, njnt: int, nlight: int, @@ -2548,6 +2674,7 @@ def _step_shim( nmaxmeshdeg: int, nmaxpolygon: int, nmaxpyramid: int, + nmesh: int, nmeshface: int, nrangefinder: int, nsensorcollision: int, @@ -2578,7 +2705,15 @@ def _step_shim( qD_fullm_i: wp.array[int], qD_fullm_j: wp.array[int], qLD_all_updates: wp.array[wp.vec3i], + qLD_block_adr: wp.array[int], + qLD_block_total: int, + qLD_dof_dense: wp.array[int], + qLD_dof_simple: wp.array[int], + qLD_has_dense: bool, + qLD_has_simple: bool, + qLD_has_sparse: bool, qLD_level_offsets: wp.array[int], + qLD_simple_dofs: wp.array[int], qLD_updates: tuple[wp.array[wp.vec3i], ...], qpos0: wp.array2d[float], qpos_spring: wp.array2d[float], @@ -2682,6 +2817,7 @@ def _step_shim( opt__timestep: wp.array[float], opt__tolerance: wp.array[float], opt__viscosity: wp.array[float], + opt__warn_overflow: bool, opt__wind: wp.array[wp.vec3], stat__meaninertia: wp.array[float], # Data @@ -2689,7 +2825,9 @@ def _step_shim( naconmax: int, njmax: int, njmax_nnz: int, - M: wp.array3d[float], + nvmax: int, + nvmax_pad: int, + M: wp.array2d[float], act: wp.array2d[float], act_dot: wp.array2d[float], actuator_force: wp.array2d[float], @@ -2698,23 +2836,42 @@ def _step_shim( actuator_velocity: wp.array2d[float], body_awake: wp.array2d[int], body_awake_ind: wp.array2d[int], + cJ: wp.array3d[float], + cM: wp.array3d[float], + cMa: wp.array2d[float], cacc: wp.array2d[wp.spatial_vector], cam_xmat: wp.array2d[wp.mat33], cam_xpos: wp.array2d[wp.vec3], cdof: wp.array2d[wp.spatial_vector], + cdof_dof: wp.array2d[int], cdof_dot: wp.array2d[wp.spatial_vector], + cdof_tri_col: wp.array[int], + cdof_tri_row: wp.array[int], cfrc_ext: wp.array2d[wp.spatial_vector], cfrc_int: wp.array2d[wp.spatial_vector], cinert: wp.array2d[mjwp_types.vec10], + cls_tol: wp.array[float], + cqLD: wp.array3d[float], + cqacc: wp.array2d[float], + cqacc_smooth: wp.array2d[float], + cqacc_warmstart: wp.array2d[float], + cqfrc_constraint: wp.array2d[float], + cqfrc_smooth: wp.array2d[float], crb: wp.array2d[mjwp_types.vec10], + crhs: wp.array3d[float], + ctol: wp.array[float], ctrl: wp.array2d[float], cvel: wp.array2d[wp.spatial_vector], + cx: wp.array3d[float], dof_awake_ind: wp.array2d[int], + dof_cdof: wp.array2d[int], dof_island: wp.array2d[int], dof_islandid: wp.array2d[int], efc_islandid: wp.array2d[int], energy: wp.array[wp.vec2], eq_active: wp.array2d[bool], + flex_aabb_max: wp.array2d[wp.vec3], + flex_aabb_min: wp.array2d[wp.vec3], flexedge_J: wp.array2d[float], flexedge_length: wp.array2d[float], flexedge_velocity: wp.array2d[float], @@ -2722,13 +2879,9 @@ def _step_shim( geom_xmat: wp.array2d[wp.mat33], geom_xpos: wp.array2d[wp.vec3], history: wp.array2d[float], - iqacc: wp.array2d[float], - iqacc_smooth: wp.array2d[float], - iqfrc_constraint: wp.array2d[float], - iqfrc_smooth: wp.array2d[float], island_dofadr: wp.array2d[int], - island_efcadr: wp.array2d[int], island_idofadr: wp.array2d[int], + island_iefcadr: wp.array2d[int], island_ne: wp.array2d[int], island_nefc: wp.array2d[int], island_nf: wp.array2d[int], @@ -2746,6 +2899,7 @@ def _step_shim( moment_rownnz: wp.array2d[int], nacon: wp.array[int], nbody_awake: wp.array[int], + ncdof: wp.array[int], ncollision: wp.array[int], ne: wp.array[int], nefc: wp.array[int], @@ -2755,9 +2909,10 @@ def _step_shim( nl: wp.array[int], ntree_awake: wp.array[int], nv_awake: wp.array[int], - qLD: wp.array3d[float], + overflow: wp.array[int], + qLD: wp.array2d[float], qLDiagInv: wp.array2d[float], - qLU: wp.array3d[float], + qLU: wp.array2d[float], qacc: wp.array2d[float], qacc_smooth: wp.array2d[float], qacc_warmstart: wp.array2d[float], @@ -2802,6 +2957,7 @@ def _step_shim( 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], @@ -2825,19 +2981,11 @@ def _step_shim( efc__aref: wp.array2d[float], efc__force: wp.array2d[float], efc__frictionloss: wp.array2d[float], - efc__iD: wp.array2d[float], - efc__iJ: wp.array3d[float], - efc__iJ_colind: wp.array3d[int], - efc__iJ_rowadr: wp.array2d[int], - efc__iJ_rownnz: wp.array2d[int], - efc__iaref: wp.array2d[float], efc__id: wp.array2d[int], - efc__iforce: wp.array2d[float], - efc__ifrictionloss: wp.array2d[float], - efc__iid: wp.array2d[int], efc__island: wp.array2d[int], - efc__istate: wp.array2d[int], - efc__itype: wp.array2d[int], + efc__jtdaj_adr: wp.array2d[int], + efc__jtdaj_nblock: wp.array[int], + efc__jtdaj_nrow: wp.array2d[int], efc__margin: wp.array2d[float], efc__pos: wp.array2d[float], efc__state: wp.array2d[int], @@ -2853,9 +3001,11 @@ def _step_shim( _m.D_diag = D_diag _m.D_rowadr = D_rowadr _m.D_rownnz = D_rownnz + _m.M_colind = M_colind _m.M_elemid = M_elemid _m.M_fullm_i = M_fullm_i _m.M_fullm_j = M_fullm_j + _m.M_hinit_i = M_hinit_i _m.M_mulm_col = M_mulm_col _m.M_mulm_madr = M_mulm_madr _m.M_mulm_rowadr = M_mulm_rowadr @@ -2892,7 +3042,9 @@ def _step_shim( _m.body_branches = body_branches _m.body_dofadr = body_dofadr _m.body_dofnum = body_dofnum + _m.body_fluid_box_adr = body_fluid_box_adr _m.body_fluid_ellipsoid = body_fluid_ellipsoid + _m.body_fluid_ellipsoid_adr = body_fluid_ellipsoid_adr _m.body_geomadr = body_geomadr _m.body_geomnum = body_geomnum _m.body_gravcomp = body_gravcomp @@ -2968,15 +3120,23 @@ def _step_shim( _m.flex_elemdataadr = flex_elemdataadr _m.flex_elemedge = flex_elemedge _m.flex_elemedgeadr = flex_elemedgeadr + _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 @@ -2992,6 +3152,9 @@ def _step_shim( _m.flexedge_J_rownnz = flexedge_J_rownnz _m.flexedge_invweight0 = flexedge_invweight0 _m.flexedge_length0 = flexedge_length0 + _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 @@ -3016,6 +3179,7 @@ def _step_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_fluid = has_fluid _m.has_sdf_geom = has_sdf_geom _m.hfield_adr = hfield_adr @@ -3052,6 +3216,7 @@ def _step_shim( _m.light_targetbodyid = light_targetbodyid _m.mapM2D = mapM2D _m.mat_rgba = mat_rgba + _m.max_flex_dim = max_flex_dim _m.max_ten_J_rownnz = max_ten_J_rownnz _m.mesh_face = mesh_face _m.mesh_faceadr = mesh_faceadr @@ -3070,6 +3235,7 @@ def _step_shim( _m.mesh_polyvert = mesh_polyvert _m.mesh_polyvertadr = mesh_polyvertadr _m.mesh_polyvertnum = mesh_polyvertnum + _m.mesh_pos = mesh_pos _m.mesh_quat = mesh_quat _m.mesh_vert = mesh_vert _m.mesh_vertadr = mesh_vertadr @@ -3086,10 +3252,9 @@ def _step_shim( _m.nflex = nflex _m.nflexedge = nflexedge _m.nflexelem = nflexelem - _m.nflexshelldata = nflexshelldata + _m.nflexevpair = nflexevpair _m.nflexvert = nflexvert _m.ngeom = ngeom - _m.ngravcomp = ngravcomp _m.nhistory = nhistory _m.njnt = njnt _m.nlight = nlight @@ -3097,6 +3262,7 @@ def _step_shim( _m.nmaxmeshdeg = nmaxmeshdeg _m.nmaxpolygon = nmaxpolygon _m.nmaxpyramid = nmaxpyramid + _m.nmesh = nmesh _m.nmeshface = nmeshface _m.nrangefinder = nrangefinder _m.nsensorcollision = nsensorcollision @@ -3140,6 +3306,7 @@ def _step_shim( _m.opt.timestep = opt__timestep _m.opt.tolerance = opt__tolerance _m.opt.viscosity = opt__viscosity + _m.opt.warn_overflow = opt__warn_overflow _m.opt.wind = opt__wind _m.pair_dim = pair_dim _m.pair_friction = pair_friction @@ -3153,7 +3320,15 @@ def _step_shim( _m.qD_fullm_i = qD_fullm_i _m.qD_fullm_j = qD_fullm_j _m.qLD_all_updates = qLD_all_updates + _m.qLD_block_adr = qLD_block_adr + _m.qLD_block_total = qLD_block_total + _m.qLD_dof_dense = qLD_dof_dense + _m.qLD_dof_simple = qLD_dof_simple + _m.qLD_has_dense = qLD_has_dense + _m.qLD_has_simple = qLD_has_simple + _m.qLD_has_sparse = qLD_has_sparse _m.qLD_level_offsets = qLD_level_offsets + _m.qLD_simple_dofs = qLD_simple_dofs _m.qLD_updates = qLD_updates _m.qpos0 = qpos0 _m.qpos_spring = qpos_spring @@ -3242,17 +3417,25 @@ def _step_shim( _d.actuator_velocity = actuator_velocity _d.body_awake = body_awake _d.body_awake_ind = body_awake_ind + _d.cJ = cJ + _d.cM = cM + _d.cMa = cMa _d.cacc = cacc _d.cam_xmat = cam_xmat _d.cam_xpos = cam_xpos _d.cdof = cdof + _d.cdof_dof = cdof_dof _d.cdof_dot = cdof_dot + _d.cdof_tri_col = cdof_tri_col + _d.cdof_tri_row = cdof_tri_row _d.cfrc_ext = cfrc_ext _d.cfrc_int = cfrc_int _d.cinert = cinert + _d.cls_tol = cls_tol _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 @@ -3266,10 +3449,20 @@ def _step_shim( _d.contact.type = contact__type _d.contact.vert = contact__vert _d.contact.worldid = contact__worldid + _d.cqLD = cqLD + _d.cqacc = cqacc + _d.cqacc_smooth = cqacc_smooth + _d.cqacc_warmstart = cqacc_warmstart + _d.cqfrc_constraint = cqfrc_constraint + _d.cqfrc_smooth = cqfrc_smooth _d.crb = crb + _d.crhs = crhs + _d.ctol = ctol _d.ctrl = ctrl _d.cvel = cvel + _d.cx = cx _d.dof_awake_ind = dof_awake_ind + _d.dof_cdof = dof_cdof _d.dof_island = dof_island _d.dof_islandid = dof_islandid _d.efc.D = efc__D @@ -3282,19 +3475,11 @@ def _step_shim( _d.efc.aref = efc__aref _d.efc.force = efc__force _d.efc.frictionloss = efc__frictionloss - _d.efc.iD = efc__iD - _d.efc.iJ = efc__iJ - _d.efc.iJ_colind = efc__iJ_colind - _d.efc.iJ_rowadr = efc__iJ_rowadr - _d.efc.iJ_rownnz = efc__iJ_rownnz - _d.efc.iaref = efc__iaref _d.efc.id = efc__id - _d.efc.iforce = efc__iforce - _d.efc.ifrictionloss = efc__ifrictionloss - _d.efc.iid = efc__iid _d.efc.island = efc__island - _d.efc.istate = efc__istate - _d.efc.itype = efc__itype + _d.efc.jtdaj_adr = efc__jtdaj_adr + _d.efc.jtdaj_nblock = efc__jtdaj_nblock + _d.efc.jtdaj_nrow = efc__jtdaj_nrow _d.efc.margin = efc__margin _d.efc.pos = efc__pos _d.efc.state = efc__state @@ -3303,6 +3488,8 @@ def _step_shim( _d.efc_islandid = efc_islandid _d.energy = energy _d.eq_active = eq_active + _d.flex_aabb_max = flex_aabb_max + _d.flex_aabb_min = flex_aabb_min _d.flexedge_J = flexedge_J _d.flexedge_length = flexedge_length _d.flexedge_velocity = flexedge_velocity @@ -3310,13 +3497,9 @@ def _step_shim( _d.geom_xmat = geom_xmat _d.geom_xpos = geom_xpos _d.history = history - _d.iqacc = iqacc - _d.iqacc_smooth = iqacc_smooth - _d.iqfrc_constraint = iqfrc_constraint - _d.iqfrc_smooth = iqfrc_smooth _d.island_dofadr = island_dofadr - _d.island_efcadr = island_efcadr _d.island_idofadr = island_idofadr + _d.island_iefcadr = island_iefcadr _d.island_ne = island_ne _d.island_nefc = island_nefc _d.island_nf = island_nf @@ -3336,6 +3519,7 @@ def _step_shim( _d.nacon = nacon _d.naconmax = naconmax _d.nbody_awake = nbody_awake + _d.ncdof = ncdof _d.ncollision = ncollision _d.ne = ne _d.nefc = nefc @@ -3347,6 +3531,9 @@ def _step_shim( _d.nl = nl _d.ntree_awake = ntree_awake _d.nv_awake = nv_awake + _d.nvmax = nvmax + _d.nvmax_pad = nvmax_pad + _d.overflow = overflow _d.qLD = qLD _d.qLDiagInv = qLDiagInv _d.qLU = qLU @@ -3406,21 +3593,32 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'actuator_velocity': d._impl.actuator_velocity.shape, 'body_awake': d._impl.body_awake.shape, 'body_awake_ind': d._impl.body_awake_ind.shape, + 'cJ': d._impl.cJ.shape, + 'cM': d._impl.cM.shape, 'cacc': d._impl.cacc.shape, 'cam_xmat': d.cam_xmat.shape, 'cam_xpos': d.cam_xpos.shape, 'cdof': d.cdof.shape, + 'cdof_dof': d._impl.cdof_dof.shape, 'cdof_dot': d.cdof_dot.shape, 'cfrc_ext': d._impl.cfrc_ext.shape, 'cfrc_int': d._impl.cfrc_int.shape, 'cinert': d._impl.cinert.shape, + 'cqacc_smooth': d._impl.cqacc_smooth.shape, + 'cqacc_warmstart': d._impl.cqacc_warmstart.shape, + 'cqfrc_smooth': d._impl.cqfrc_smooth.shape, 'crb': d._impl.crb.shape, + 'crhs': d._impl.crhs.shape, 'cvel': d.cvel.shape, + 'cx': d._impl.cx.shape, 'dof_awake_ind': d._impl.dof_awake_ind.shape, + 'dof_cdof': d._impl.dof_cdof.shape, 'dof_island': d._impl.dof_island.shape, 'dof_islandid': d._impl.dof_islandid.shape, 'efc_islandid': d._impl.efc_islandid.shape, 'energy': d._impl.energy.shape, + 'flex_aabb_max': d._impl.flex_aabb_max.shape, + 'flex_aabb_min': d._impl.flex_aabb_min.shape, 'flexedge_J': d._impl.flexedge_J.shape, 'flexedge_length': d._impl.flexedge_length.shape, 'flexedge_velocity': d._impl.flexedge_velocity.shape, @@ -3428,13 +3626,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'geom_xmat': d.geom_xmat.shape, 'geom_xpos': d.geom_xpos.shape, 'history': d.history.shape, - 'iqacc': d._impl.iqacc.shape, - 'iqacc_smooth': d._impl.iqacc_smooth.shape, - 'iqfrc_constraint': d._impl.iqfrc_constraint.shape, - 'iqfrc_smooth': d._impl.iqfrc_smooth.shape, 'island_dofadr': d._impl.island_dofadr.shape, - 'island_efcadr': d._impl.island_efcadr.shape, 'island_idofadr': d._impl.island_idofadr.shape, + 'island_iefcadr': d._impl.island_iefcadr.shape, 'island_ne': d._impl.island_ne.shape, 'island_nefc': d._impl.island_nefc.shape, 'island_nf': d._impl.island_nf.shape, @@ -3450,6 +3644,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'moment_rownnz': d._impl.moment_rownnz.shape, 'nacon': d._impl.nacon.shape, 'nbody_awake': d._impl.nbody_awake.shape, + 'ncdof': d._impl.ncdof.shape, 'ncollision': d._impl.ncollision.shape, 'ne': d._impl.ne.shape, 'nefc': d._impl.nefc.shape, @@ -3459,6 +3654,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'nl': d._impl.nl.shape, 'ntree_awake': d._impl.ntree_awake.shape, 'nv_awake': d._impl.nv_awake.shape, + 'overflow': d._impl.overflow.shape, 'qLD': d._impl.qLD.shape, 'qLDiagInv': d._impl.qLDiagInv.shape, 'qLU': d._impl.qLU.shape, @@ -3504,6 +3700,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): '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, @@ -3527,19 +3724,11 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'efc__aref': d._impl.efc__aref.shape, 'efc__force': d._impl.efc__force.shape, 'efc__frictionloss': d._impl.efc__frictionloss.shape, - 'efc__iD': d._impl.efc__iD.shape, - 'efc__iJ': d._impl.efc__iJ.shape, - 'efc__iJ_colind': d._impl.efc__iJ_colind.shape, - 'efc__iJ_rowadr': d._impl.efc__iJ_rowadr.shape, - 'efc__iJ_rownnz': d._impl.efc__iJ_rownnz.shape, - 'efc__iaref': d._impl.efc__iaref.shape, 'efc__id': d._impl.efc__id.shape, - 'efc__iforce': d._impl.efc__iforce.shape, - 'efc__ifrictionloss': d._impl.efc__ifrictionloss.shape, - 'efc__iid': d._impl.efc__iid.shape, 'efc__island': d._impl.efc__island.shape, - 'efc__istate': d._impl.efc__istate.shape, - 'efc__itype': d._impl.efc__itype.shape, + 'efc__jtdaj_adr': d._impl.efc__jtdaj_adr.shape, + 'efc__jtdaj_nblock': d._impl.efc__jtdaj_nblock.shape, + 'efc__jtdaj_nrow': d._impl.efc__jtdaj_nrow.shape, 'efc__margin': d._impl.efc__margin.shape, 'efc__pos': d._impl.efc__pos.shape, 'efc__state': d._impl.efc__state.shape, @@ -3548,7 +3737,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): } jf = ffi.jax_callable_variadic_tuple( _step_shim, - num_outputs=148, + num_outputs=150, output_dims=output_dims, vmap_method=None, in_out_argnames=set([ @@ -3561,21 +3750,32 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'actuator_velocity', 'body_awake', 'body_awake_ind', + 'cJ', + 'cM', 'cacc', 'cam_xmat', 'cam_xpos', 'cdof', + 'cdof_dof', 'cdof_dot', 'cfrc_ext', 'cfrc_int', 'cinert', + 'cqacc_smooth', + 'cqacc_warmstart', + 'cqfrc_smooth', 'crb', + 'crhs', 'cvel', + 'cx', 'dof_awake_ind', + 'dof_cdof', 'dof_island', 'dof_islandid', 'efc_islandid', 'energy', + 'flex_aabb_max', + 'flex_aabb_min', 'flexedge_J', 'flexedge_length', 'flexedge_velocity', @@ -3583,13 +3783,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'geom_xmat', 'geom_xpos', 'history', - 'iqacc', - 'iqacc_smooth', - 'iqfrc_constraint', - 'iqfrc_smooth', 'island_dofadr', - 'island_efcadr', 'island_idofadr', + 'island_iefcadr', 'island_ne', 'island_nefc', 'island_nf', @@ -3605,6 +3801,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'moment_rownnz', 'nacon', 'nbody_awake', + 'ncdof', 'ncollision', 'ne', 'nefc', @@ -3614,6 +3811,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'nl', 'ntree_awake', 'nv_awake', + 'overflow', 'qLD', 'qLDiagInv', 'qLU', @@ -3659,6 +3857,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'contact__dim', 'contact__dist', 'contact__efc_address', + 'contact__elem', 'contact__flex', 'contact__frame', 'contact__friction', @@ -3682,19 +3881,11 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'efc__aref', 'efc__force', 'efc__frictionloss', - 'efc__iD', - 'efc__iJ', - 'efc__iJ_colind', - 'efc__iJ_rowadr', - 'efc__iJ_rownnz', - 'efc__iaref', 'efc__id', - 'efc__iforce', - 'efc__ifrictionloss', - 'efc__iid', 'efc__island', - 'efc__istate', - 'efc__itype', + 'efc__jtdaj_adr', + 'efc__jtdaj_nblock', + 'efc__jtdaj_nrow', 'efc__margin', 'efc__pos', 'efc__state', @@ -3893,9 +4084,11 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.D_diag, m._impl.D_rowadr, m._impl.D_rownnz, + m.M_colind, m._impl.M_elemid, m._impl.M_fullm_i, m._impl.M_fullm_j, + m._impl.M_hinit_i, m._impl.M_mulm_col, m._impl.M_mulm_madr, m._impl.M_mulm_rowadr, @@ -3932,7 +4125,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.body_branches, m.body_dofadr, m.body_dofnum, + m._impl.body_fluid_box_adr, m._impl.body_fluid_ellipsoid, + m._impl.body_fluid_ellipsoid_adr, m.body_geomadr, m.body_geomnum, m.body_gravcomp, @@ -4008,15 +4203,23 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.flex_elemdataadr, m._impl.flex_elemedge, m._impl.flex_elemedgeadr, + 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, @@ -4032,6 +4235,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.flexedge_J_rownnz, m._impl.flexedge_invweight0, m._impl.flexedge_length0, + 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, @@ -4056,6 +4262,7 @@ def _step_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_fluid, m._impl.has_sdf_geom, m.hfield_adr, @@ -4092,6 +4299,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.light_targetbodyid, m._impl.mapM2D, m.mat_rgba, + m._impl.max_flex_dim, m._impl.max_ten_J_rownnz, m.mesh_face, m.mesh_faceadr, @@ -4110,6 +4318,7 @@ def _step_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_quat, m.mesh_vert, m.mesh_vertadr, @@ -4126,10 +4335,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.nflex, m._impl.nflexedge, m._impl.nflexelem, - m._impl.nflexshelldata, + m._impl.nflexevpair, m._impl.nflexvert, m.ngeom, - m.ngravcomp, m.nhistory, m.njnt, m.nlight, @@ -4137,6 +4345,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.nmaxmeshdeg, m._impl.nmaxpolygon, m._impl.nmaxpyramid, + m.nmesh, m.nmeshface, m._impl.nrangefinder, m._impl.nsensorcollision, @@ -4167,7 +4376,15 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.qD_fullm_i, m._impl.qD_fullm_j, m._impl.qLD_all_updates, + m._impl.qLD_block_adr, + m._impl.qLD_block_total, + m._impl.qLD_dof_dense, + m._impl.qLD_dof_simple, + m._impl.qLD_has_dense, + m._impl.qLD_has_simple, + m._impl.qLD_has_sparse, m._impl.qLD_level_offsets, + m._impl.qLD_simple_dofs, m._impl.qLD_updates, m.qpos0, m.qpos_spring, @@ -4271,12 +4488,15 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.opt.timestep, m.opt.tolerance, m.opt.viscosity, + m.opt._impl.warn_overflow, m.opt.wind, m.stat.meaninertia, d._impl.naccdmax, d._impl.naconmax, d._impl.njmax, d._impl.njmax_nnz, + d._impl.nvmax, + d._impl.nvmax_pad, d._impl.M, d.act, d.act_dot, @@ -4286,23 +4506,42 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.actuator_velocity, d._impl.body_awake, d._impl.body_awake_ind, + d._impl.cJ, + d._impl.cM, + d._impl.cMa, d._impl.cacc, d.cam_xmat, d.cam_xpos, d.cdof, + d._impl.cdof_dof, d.cdof_dot, + d._impl.cdof_tri_col, + d._impl.cdof_tri_row, d._impl.cfrc_ext, d._impl.cfrc_int, d._impl.cinert, + d._impl.cls_tol, + d._impl.cqLD, + d._impl.cqacc, + d._impl.cqacc_smooth, + d._impl.cqacc_warmstart, + d._impl.cqfrc_constraint, + d._impl.cqfrc_smooth, d._impl.crb, + d._impl.crhs, + d._impl.ctol, d.ctrl, d.cvel, + d._impl.cx, d._impl.dof_awake_ind, + d._impl.dof_cdof, d._impl.dof_island, d._impl.dof_islandid, d._impl.efc_islandid, d._impl.energy, d.eq_active, + d._impl.flex_aabb_max, + d._impl.flex_aabb_min, d._impl.flexedge_J, d._impl.flexedge_length, d._impl.flexedge_velocity, @@ -4310,13 +4549,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): d.geom_xmat, d.geom_xpos, d.history, - d._impl.iqacc, - d._impl.iqacc_smooth, - d._impl.iqfrc_constraint, - d._impl.iqfrc_smooth, d._impl.island_dofadr, - d._impl.island_efcadr, d._impl.island_idofadr, + d._impl.island_iefcadr, d._impl.island_ne, d._impl.island_nefc, d._impl.island_nf, @@ -4334,6 +4569,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.moment_rownnz, d._impl.nacon, d._impl.nbody_awake, + d._impl.ncdof, d._impl.ncollision, d._impl.ne, d._impl.nefc, @@ -4343,6 +4579,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.nl, d._impl.ntree_awake, d._impl.nv_awake, + d._impl.overflow, d._impl.qLD, d._impl.qLDiagInv, d._impl.qLU, @@ -4390,6 +4627,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 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, @@ -4413,19 +4651,11 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.efc__aref, d._impl.efc__force, d._impl.efc__frictionloss, - d._impl.efc__iD, - d._impl.efc__iJ, - d._impl.efc__iJ_colind, - d._impl.efc__iJ_rowadr, - d._impl.efc__iJ_rownnz, - d._impl.efc__iaref, d._impl.efc__id, - d._impl.efc__iforce, - d._impl.efc__ifrictionloss, - d._impl.efc__iid, d._impl.efc__island, - d._impl.efc__istate, - d._impl.efc__itype, + d._impl.efc__jtdaj_adr, + d._impl.efc__jtdaj_nblock, + d._impl.efc__jtdaj_nrow, d._impl.efc__margin, d._impl.efc__pos, d._impl.efc__state, @@ -4442,145 +4672,147 @@ def _step_jax_impl(m: types.Model, d: types.Data): '_impl.actuator_velocity': out[6], '_impl.body_awake': out[7], '_impl.body_awake_ind': out[8], - '_impl.cacc': out[9], - 'cam_xmat': out[10], - 'cam_xpos': out[11], - 'cdof': out[12], - 'cdof_dot': out[13], - '_impl.cfrc_ext': out[14], - '_impl.cfrc_int': out[15], - '_impl.cinert': out[16], - '_impl.crb': out[17], - 'cvel': out[18], - '_impl.dof_awake_ind': out[19], - '_impl.dof_island': out[20], - '_impl.dof_islandid': out[21], - '_impl.efc_islandid': out[22], - '_impl.energy': out[23], - '_impl.flexedge_J': out[24], - '_impl.flexedge_length': out[25], - '_impl.flexedge_velocity': out[26], - '_impl.flexvert_xpos': out[27], - 'geom_xmat': out[28], - 'geom_xpos': out[29], - 'history': out[30], - '_impl.iqacc': out[31], - '_impl.iqacc_smooth': out[32], - '_impl.iqfrc_constraint': out[33], - '_impl.iqfrc_smooth': out[34], - '_impl.island_dofadr': out[35], - '_impl.island_efcadr': out[36], - '_impl.island_idofadr': out[37], - '_impl.island_ne': out[38], - '_impl.island_nefc': out[39], - '_impl.island_nf': out[40], - '_impl.island_nv': out[41], - '_impl.light_xdir': out[42], - '_impl.light_xpos': out[43], - '_impl.map_dof2idof': out[44], - '_impl.map_efc2iefc': out[45], - '_impl.map_idof2dof': out[46], - '_impl.map_iefc2efc': out[47], - '_impl.moment_colind': out[48], - '_impl.moment_rowadr': out[49], - '_impl.moment_rownnz': out[50], - '_impl.nacon': out[51], - '_impl.nbody_awake': out[52], - '_impl.ncollision': out[53], - '_impl.ne': out[54], - '_impl.nefc': out[55], - '_impl.nf': out[56], - '_impl.nidof': out[57], - '_impl.nisland': out[58], - '_impl.nl': out[59], - '_impl.ntree_awake': out[60], - '_impl.nv_awake': out[61], - '_impl.qLD': out[62], - '_impl.qLDiagInv': out[63], - '_impl.qLU': out[64], - 'qacc': out[65], - 'qacc_smooth': out[66], - 'qacc_warmstart': out[67], - 'qfrc_actuator': out[68], - 'qfrc_bias': out[69], - 'qfrc_constraint': out[70], - '_impl.qfrc_damper': out[71], - 'qfrc_fluid': out[72], - 'qfrc_gravcomp': out[73], - 'qfrc_passive': out[74], - 'qfrc_smooth': out[75], - '_impl.qfrc_spring': out[76], - 'qpos': out[77], - 'qvel': out[78], - 'sensordata': out[79], - 'site_xmat': out[80], - 'site_xpos': out[81], - '_impl.solver_niter': out[82], - '_impl.subtree_angmom': out[83], - 'subtree_com': out[84], - '_impl.subtree_linvel': out[85], - '_impl.ten_J': out[86], - 'ten_length': out[87], - '_impl.ten_velocity': out[88], - '_impl.ten_wrapadr': out[89], - '_impl.ten_wrapnum': out[90], - 'time': out[91], - '_impl.tree_asleep': out[92], - '_impl.tree_awake': out[93], - '_impl.tree_island': out[94], - '_impl.wrap_obj': out[95], - '_impl.wrap_xpos': out[96], - 'xanchor': out[97], - 'xaxis': out[98], - 'ximat': out[99], - 'xipos': out[100], - 'xmat': out[101], - 'xpos': out[102], - 'xquat': out[103], - '_impl.contact__dim': out[104], - '_impl.contact__dist': out[105], - '_impl.contact__efc_address': out[106], - '_impl.contact__flex': out[107], - '_impl.contact__frame': out[108], - '_impl.contact__friction': out[109], - '_impl.contact__geom': out[110], - '_impl.contact__geomcollisionid': out[111], - '_impl.contact__includemargin': out[112], - '_impl.contact__pos': out[113], - '_impl.contact__solimp': out[114], - '_impl.contact__solref': out[115], - '_impl.contact__solreffriction': out[116], - '_impl.contact__type': out[117], - '_impl.contact__vert': out[118], - '_impl.contact__worldid': out[119], - '_impl.efc__D': out[120], - '_impl.efc__J': out[121], - '_impl.efc__J_colind': out[122], - '_impl.efc__J_rowadr': out[123], - '_impl.efc__J_rownnz': out[124], - '_impl.efc__Jqvel': out[125], - '_impl.efc__Ma': out[126], - '_impl.efc__aref': out[127], - '_impl.efc__force': out[128], - '_impl.efc__frictionloss': out[129], - '_impl.efc__iD': out[130], - '_impl.efc__iJ': out[131], - '_impl.efc__iJ_colind': out[132], - '_impl.efc__iJ_rowadr': out[133], - '_impl.efc__iJ_rownnz': out[134], - '_impl.efc__iaref': out[135], - '_impl.efc__id': out[136], - '_impl.efc__iforce': out[137], - '_impl.efc__ifrictionloss': out[138], - '_impl.efc__iid': out[139], - '_impl.efc__island': out[140], - '_impl.efc__istate': out[141], - '_impl.efc__itype': out[142], - '_impl.efc__margin': out[143], - '_impl.efc__pos': out[144], - '_impl.efc__state': out[145], - '_impl.efc__type': out[146], - '_impl.efc__vel': out[147], + '_impl.cJ': out[9], + '_impl.cM': out[10], + '_impl.cacc': out[11], + 'cam_xmat': out[12], + 'cam_xpos': out[13], + 'cdof': out[14], + '_impl.cdof_dof': out[15], + 'cdof_dot': out[16], + '_impl.cfrc_ext': out[17], + '_impl.cfrc_int': out[18], + '_impl.cinert': out[19], + '_impl.cqacc_smooth': out[20], + '_impl.cqacc_warmstart': out[21], + '_impl.cqfrc_smooth': out[22], + '_impl.crb': out[23], + '_impl.crhs': out[24], + 'cvel': out[25], + '_impl.cx': out[26], + '_impl.dof_awake_ind': out[27], + '_impl.dof_cdof': out[28], + '_impl.dof_island': out[29], + '_impl.dof_islandid': out[30], + '_impl.efc_islandid': out[31], + '_impl.energy': out[32], + '_impl.flex_aabb_max': out[33], + '_impl.flex_aabb_min': out[34], + '_impl.flexedge_J': out[35], + '_impl.flexedge_length': out[36], + '_impl.flexedge_velocity': out[37], + '_impl.flexvert_xpos': out[38], + 'geom_xmat': out[39], + 'geom_xpos': out[40], + 'history': out[41], + '_impl.island_dofadr': out[42], + '_impl.island_idofadr': out[43], + '_impl.island_iefcadr': out[44], + '_impl.island_ne': out[45], + '_impl.island_nefc': out[46], + '_impl.island_nf': out[47], + '_impl.island_nv': out[48], + '_impl.light_xdir': out[49], + '_impl.light_xpos': out[50], + '_impl.map_dof2idof': out[51], + '_impl.map_efc2iefc': out[52], + '_impl.map_idof2dof': out[53], + '_impl.map_iefc2efc': out[54], + '_impl.moment_colind': out[55], + '_impl.moment_rowadr': out[56], + '_impl.moment_rownnz': out[57], + '_impl.nacon': out[58], + '_impl.nbody_awake': out[59], + '_impl.ncdof': out[60], + '_impl.ncollision': out[61], + '_impl.ne': out[62], + '_impl.nefc': out[63], + '_impl.nf': out[64], + '_impl.nidof': out[65], + '_impl.nisland': out[66], + '_impl.nl': out[67], + '_impl.ntree_awake': out[68], + '_impl.nv_awake': out[69], + '_impl.overflow': out[70], + '_impl.qLD': out[71], + '_impl.qLDiagInv': out[72], + '_impl.qLU': out[73], + 'qacc': out[74], + 'qacc_smooth': out[75], + 'qacc_warmstart': out[76], + 'qfrc_actuator': out[77], + 'qfrc_bias': out[78], + 'qfrc_constraint': out[79], + '_impl.qfrc_damper': out[80], + 'qfrc_fluid': out[81], + 'qfrc_gravcomp': out[82], + 'qfrc_passive': out[83], + 'qfrc_smooth': out[84], + '_impl.qfrc_spring': out[85], + 'qpos': out[86], + 'qvel': out[87], + 'sensordata': out[88], + 'site_xmat': out[89], + 'site_xpos': out[90], + '_impl.solver_niter': out[91], + '_impl.subtree_angmom': out[92], + 'subtree_com': out[93], + '_impl.subtree_linvel': out[94], + '_impl.ten_J': out[95], + 'ten_length': out[96], + '_impl.ten_velocity': out[97], + '_impl.ten_wrapadr': out[98], + '_impl.ten_wrapnum': out[99], + 'time': out[100], + '_impl.tree_asleep': out[101], + '_impl.tree_awake': out[102], + '_impl.tree_island': out[103], + '_impl.wrap_obj': out[104], + '_impl.wrap_xpos': out[105], + 'xanchor': out[106], + 'xaxis': out[107], + 'ximat': out[108], + 'xipos': out[109], + 'xmat': out[110], + 'xpos': out[111], + 'xquat': out[112], + '_impl.contact__dim': out[113], + '_impl.contact__dist': out[114], + '_impl.contact__efc_address': out[115], + '_impl.contact__elem': out[116], + '_impl.contact__flex': out[117], + '_impl.contact__frame': out[118], + '_impl.contact__friction': out[119], + '_impl.contact__geom': out[120], + '_impl.contact__geomcollisionid': out[121], + '_impl.contact__includemargin': out[122], + '_impl.contact__pos': out[123], + '_impl.contact__solimp': out[124], + '_impl.contact__solref': out[125], + '_impl.contact__solreffriction': out[126], + '_impl.contact__type': out[127], + '_impl.contact__vert': out[128], + '_impl.contact__worldid': out[129], + '_impl.efc__D': out[130], + '_impl.efc__J': out[131], + '_impl.efc__J_colind': out[132], + '_impl.efc__J_rowadr': out[133], + '_impl.efc__J_rownnz': out[134], + '_impl.efc__Jqvel': out[135], + '_impl.efc__Ma': out[136], + '_impl.efc__aref': out[137], + '_impl.efc__force': out[138], + '_impl.efc__frictionloss': out[139], + '_impl.efc__id': out[140], + '_impl.efc__island': out[141], + '_impl.efc__jtdaj_adr': out[142], + '_impl.efc__jtdaj_nblock': out[143], + '_impl.efc__jtdaj_nrow': out[144], + '_impl.efc__margin': out[145], + '_impl.efc__pos': out[146], + '_impl.efc__state': out[147], + '_impl.efc__type': out[148], + '_impl.efc__vel': out[149], }) return d diff --git a/mjx/mujoco/mjx/warp/forward_test.py b/mjx/mujoco/mjx/warp/forward_test.py index 95d49e0f..720a3444 100644 --- a/mjx/mujoco/mjx/warp/forward_test.py +++ b/mjx/mujoco/mjx/warp/forward_test.py @@ -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) diff --git a/mjx/mujoco/mjx/warp/render.py b/mjx/mujoco/mjx/warp/render.py index 5ab3ae69..cbfcec54 100644 --- a/mjx/mujoco/mjx/warp/render.py +++ b/mjx/mujoco/mjx/warp/render.py @@ -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} ) diff --git a/mjx/mujoco/mjx/warp/smooth.py b/mjx/mujoco/mjx/warp/smooth.py index bf26a123..ce765d8b 100644 --- a/mjx/mujoco/mjx/warp/smooth.py +++ b/mjx/mujoco/mjx/warp/smooth.py @@ -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} ) diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index 3474b32c..18d3449f 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -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},