Skip to content

Commit 5b4cb35

Browse files
committed
Fix Implicit MPM reset masks
Adapt coupled and standalone Implicit MPM reset masks to the solver's world-mask contract, and skip unsupported selective shared-world resets. Add regression tests and changelog fragments.
1 parent ee735ae commit 5b4cb35

7 files changed

Lines changed: 307 additions & 121 deletions

File tree

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,11 @@
1+
Fixed
2+
^^^^^
3+
4+
* Fixed coupled Implicit MPM environment resets raising
5+
``ValueError: world_mask has shape ...`` or
6+
``RuntimeError: Masked reset cannot selectively clear grid-backed warm
7+
starts`` when :meth:`~isaaclab_newton.physics.NewtonManager.reset_solver_state`
8+
forwarded Isaac Lab's ``(world_count,)`` mask. MPM entry ``reset`` now receives
9+
the ``(world_count + 1,)`` mask required by
10+
:meth:`newton.solvers.SolverImplicitMPM.reset` (or skips selective shared-world
11+
resets), while MJWarp entries keep the original parent mask.

source/isaaclab_contrib/isaaclab_contrib/coupling/coupler.py

Lines changed: 54 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,9 +18,14 @@
1818
NewtonCollisionPipelineCfg,
1919
NewtonSolverCfg,
2020
)
21-
from isaaclab_newton.physics.mpm_manager import NewtonMPMManager
21+
from isaaclab_newton.physics.mpm_manager import (
22+
NewtonMPMManager,
23+
adapt_world_mask_for_implicit_mpm,
24+
should_skip_implicit_mpm_masked_reset,
25+
)
2226
from isaaclab_newton.physics.newton_manager import NewtonManager
2327
from newton import CollisionPipeline, Model, ModelBuilder, ShapeFlags
28+
from newton.solvers import SolverImplicitMPM
2429
from newton.solvers.experimental.coupled import SolverCoupled, SolverCoupledADMM, SolverCoupledProxy
2530

2631
from isaaclab.physics import PhysicsManager
@@ -114,6 +119,54 @@ def _build_solver(cls, model: Model, solver_cfg: CouplerCfg) -> None:
114119
NewtonManager._supports_contact_sensors = False
115120
NewtonManager._needs_collision_pipeline = needs_collision_pipeline
116121
NewtonManager._supports_rigid_body_force_input = True
122+
cls._install_implicit_mpm_reset_mask_adapter(
123+
NewtonManager._solver,
124+
[entry.config.name for entry in resolved_entries],
125+
)
126+
127+
@classmethod
128+
def _install_implicit_mpm_reset_mask_adapter(cls, solver: object, entry_names: list[str]) -> None:
129+
"""Adapt coupled-reset masks for Implicit MPM's world-mask contract.
130+
131+
:class:`~newton.solvers.experimental.coupled.SolverCoupled` forwards the
132+
parent ``(world_count,)`` mask to every entry. MJWarp expects that shape,
133+
while :class:`~newton.solvers.SolverImplicitMPM` expects one extra bit for
134+
global world ``-1``. Shared multi-world Implicit MPM
135+
(``separate_worlds=False``) also rejects selective masks when clearing
136+
grid-backed warm starts; those resets are skipped. Wrapping only MPM
137+
entry ``reset`` methods keeps the coupled parent API unchanged.
138+
"""
139+
# SolverCoupled exposes sub-solvers via ``solver(name)`` (not ``entry_solver``).
140+
get_entry_solver = getattr(solver, "solver", None)
141+
if not callable(get_entry_solver):
142+
return
143+
144+
for entry_name in entry_names:
145+
try:
146+
entry_solver = get_entry_solver(entry_name)
147+
except (KeyError, TypeError, ValueError):
148+
continue
149+
if not isinstance(entry_solver, SolverImplicitMPM):
150+
continue
151+
original_reset = entry_solver.reset
152+
153+
def _reset(
154+
state,
155+
world_mask=None,
156+
flags=None,
157+
*,
158+
_original_reset=original_reset,
159+
_entry_solver=entry_solver,
160+
):
161+
if should_skip_implicit_mpm_masked_reset(_entry_solver, world_mask):
162+
return None
163+
return _original_reset(
164+
state,
165+
world_mask=adapt_world_mask_for_implicit_mpm(_entry_solver, world_mask),
166+
flags=flags,
167+
)
168+
169+
entry_solver.reset = _reset # type: ignore[method-assign]
117170

118171
@classmethod
119172
def _validate_config(cls, solver_cfg: CouplerCfg) -> None:

source/isaaclab_contrib/test/coupling/test_coupler.py

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -864,3 +864,41 @@ def test_admm_build_auto_detects_symmetric_contact_pairs_by_default(monkeypatch)
864864
cfg.contact_pairs = []
865865
solver = NewtonCouplerManager._build_admm_coupled_solver(model, entries, cfg)
866866
assert list(solver.coupling.contact_pairs) == []
867+
868+
869+
def test_install_implicit_mpm_reset_mask_adapter_pads_entry_masks(monkeypatch):
870+
"""Coupled MPM entry resets pad full masks and skip selective shared-world ones."""
871+
import warp as wp
872+
873+
recorded: list[list[bool] | None] = []
874+
875+
class _FakeImplicitMPM:
876+
def __init__(self):
877+
self.model = SimpleNamespace(world_count=4)
878+
self._separate_worlds = False
879+
880+
def reset(self, state, world_mask=None, flags=None):
881+
del state, flags
882+
recorded.append(None if world_mask is None else world_mask.numpy().tolist())
883+
884+
monkeypatch.setattr(coupler, "SolverImplicitMPM", _FakeImplicitMPM)
885+
mpm_solver = _FakeImplicitMPM()
886+
887+
class _FakeCoupled:
888+
def solver(self, name: str):
889+
assert name == "mpm"
890+
return mpm_solver
891+
892+
NewtonCouplerManager._install_implicit_mpm_reset_mask_adapter(_FakeCoupled(), ["mpm"])
893+
894+
# Shared multi-world selective masks are skipped (no call into original reset).
895+
mpm_solver.reset(object(), world_mask=wp.array([True, False, True, False], dtype=wp.bool, device="cpu"))
896+
assert recorded == []
897+
898+
mpm_solver.reset(object(), world_mask=wp.array([True, True, True, True], dtype=wp.bool, device="cpu"))
899+
assert recorded == [None]
900+
901+
# With separate worlds, selective masks are forwarded as (N+1,).
902+
mpm_solver._separate_worlds = True
903+
mpm_solver.reset(object(), world_mask=wp.array([True, False, True, False], dtype=wp.bool, device="cpu"))
904+
assert recorded == [None, [True, False, True, False, False]]
Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,8 @@
1+
Fixed
2+
^^^^^
3+
4+
* Fixed :meth:`~isaaclab_newton.physics.NewtonManager.reset_solver_state` under
5+
Implicit MPM rejecting Isaac Lab ``(world_count,)`` masks. The manager now
6+
expands them to the ``(world_count + 1,)`` shape required by
7+
:meth:`newton.solvers.SolverImplicitMPM.reset`, and skips selective masks when
8+
shared multi-world Implicit MPM cannot clear grid-backed warm starts per world.

source/isaaclab_newton/isaaclab_newton/physics/mpm_manager.py

Lines changed: 108 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,15 +7,94 @@
77

88
from __future__ import annotations
99

10+
import numpy as np
1011
import warp as wp
11-
from newton import BodyFlags, Contacts, Control, GeoType, Model, ModelBuilder, State
12+
from newton import BodyFlags, Contacts, Control, GeoType, Model, ModelBuilder, State, StateFlags
1213
from newton.solvers import SolverImplicitMPM
1314
from warp.fem import TemporaryStore
1415

1516
from .mpm_manager_cfg import MPMSolverCfg
1617
from .newton_manager import NewtonManager
1718

1819

20+
def adapt_world_mask_for_implicit_mpm(
21+
solver: SolverImplicitMPM,
22+
world_mask: wp.array | None,
23+
) -> wp.array | None:
24+
"""Adapt an Isaac Lab world-reset mask to Implicit MPM's mask contract.
25+
26+
Isaac Lab and most Newton solvers use a per-world mask of shape
27+
``(world_count,)``. :meth:`SolverImplicitMPM.reset` instead expects
28+
``(world_count + 1,)``, where the trailing entry selects global objects
29+
whose world index is ``-1``.
30+
31+
Args:
32+
solver: Implicit MPM solver whose model defines ``world_count``.
33+
world_mask: Optional Isaac Lab / Newton mask of shape ``(world_count,)``
34+
or an already-adapted Implicit MPM mask of shape
35+
``(world_count + 1,)``.
36+
37+
Returns:
38+
``None`` when no mask is provided or every world is selected (a full
39+
reset, including global world ``-1``). Otherwise a boolean Warp array
40+
of shape ``(world_count + 1,)`` with the global bit left ``False``.
41+
42+
Raises:
43+
ValueError: If ``world_mask`` has neither ``(world_count,)`` nor
44+
``(world_count + 1,)`` shape.
45+
"""
46+
if world_mask is None:
47+
return None
48+
49+
local_selected = _local_world_selection(solver, world_mask)
50+
if bool(np.all(local_selected)):
51+
# Full world selection is equivalent to an unmasked reset and also
52+
# covers global (world index -1) particle/collider history.
53+
return None
54+
55+
world_count = int(solver.model.world_count)
56+
if tuple(world_mask.shape) == (world_count + 1,):
57+
return world_mask
58+
59+
padded = np.zeros(world_count + 1, dtype=bool)
60+
padded[:-1] = local_selected
61+
return wp.array(padded, dtype=wp.bool, device=world_mask.device)
62+
63+
64+
def should_skip_implicit_mpm_masked_reset(
65+
solver: SolverImplicitMPM,
66+
world_mask: wp.array | None,
67+
) -> bool:
68+
"""Whether a masked Implicit MPM reset must be skipped.
69+
70+
Shared multi-world Implicit MPM (``separate_worlds=False``) rejects selective
71+
masks when clearing grid-backed warm starts. Matching
72+
:meth:`NewtonMPMManager._reset_solver_internals`, Isaac Lab skips those
73+
resets instead of raising. Full-world selection still proceeds as an
74+
unmasked reset.
75+
"""
76+
if world_mask is None:
77+
return False
78+
if bool(getattr(solver, "_separate_worlds", False)) or int(solver.model.world_count) <= 1:
79+
return False
80+
return not bool(np.all(_local_world_selection(solver, world_mask)))
81+
82+
83+
def _local_world_selection(solver: SolverImplicitMPM, world_mask: wp.array) -> np.ndarray:
84+
"""Return the per-world selection bits from an Isaac Lab or Implicit MPM mask."""
85+
world_count = int(solver.model.world_count)
86+
shape = tuple(world_mask.shape)
87+
selected = world_mask.numpy()
88+
if shape == (world_count + 1,):
89+
return selected[:-1]
90+
if shape == (world_count,):
91+
return selected
92+
raise ValueError(
93+
f"world_mask has shape {shape}, expected ({world_count},) or ({world_count + 1},) "
94+
"for SolverImplicitMPM.reset."
95+
)
96+
97+
1998
def _make_solver_config(solver_cfg: MPMSolverCfg) -> SolverImplicitMPM.Config:
2099
"""Build Newton's implicit MPM solver config from Isaac Lab's cfg."""
21100
return SolverImplicitMPM.Config(
@@ -161,6 +240,34 @@ def _reset_solver_internals(cls, world_mask: wp.array | None) -> None:
161240
world_mask: Per-world reset mask, ignored.
162241
"""
163242

243+
@classmethod
244+
def reset_solver_state(
245+
cls,
246+
state: State | None = None,
247+
world_mask: wp.array(dtype=wp.bool) | None = None,
248+
flags: StateFlags | int | None = None,
249+
) -> None:
250+
"""Reset Implicit MPM history after simulation state is rewritten.
251+
252+
Expands Isaac Lab's ``(world_count,)`` mask to the
253+
``(world_count + 1,)`` shape required by :meth:`SolverImplicitMPM.reset`
254+
before delegating to :meth:`NewtonManager.reset_solver_state`. Shared
255+
multi-world selective masks are skipped; see
256+
:func:`should_skip_implicit_mpm_masked_reset`.
257+
"""
258+
if not isinstance(cls._solver, SolverImplicitMPM):
259+
raise RuntimeError(
260+
f"{cls.__name__}.reset_solver_state requires an active SolverImplicitMPM; "
261+
f"got {type(cls._solver).__name__}."
262+
)
263+
if should_skip_implicit_mpm_masked_reset(cls._solver, world_mask):
264+
return
265+
super().reset_solver_state(
266+
state=state,
267+
world_mask=adapt_world_mask_for_implicit_mpm(cls._solver, world_mask),
268+
flags=flags,
269+
)
270+
164271
@classmethod
165272
def _solver_specific_clear(cls) -> None:
166273
"""Reset MPM-specific class state on teardown.

source/isaaclab_newton/test/physics/test_newton_manager_abstraction.py

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -632,6 +632,45 @@ def test_mpm_unsupported_cuda_graph_capture_uses_eager_execution(monkeypatch):
632632
assert NewtonManager._graph_capture_pending is False
633633

634634

635+
def test_adapt_world_mask_for_implicit_mpm_pads_selective_masks():
636+
"""Isaac Lab (world_count,) masks expand to Implicit MPM's (world_count + 1,) contract."""
637+
from isaaclab_newton.physics.mpm_manager import (
638+
adapt_world_mask_for_implicit_mpm,
639+
should_skip_implicit_mpm_masked_reset,
640+
)
641+
642+
solver = SimpleNamespace(model=SimpleNamespace(world_count=4), _separate_worlds=False)
643+
world_mask = wp.array([True, False, True, False], dtype=wp.bool, device="cpu")
644+
645+
assert should_skip_implicit_mpm_masked_reset(solver, world_mask) is True
646+
adapted = adapt_world_mask_for_implicit_mpm(solver, world_mask)
647+
648+
assert adapted is not None
649+
assert adapted.numpy().tolist() == [True, False, True, False, False]
650+
651+
solver._separate_worlds = True
652+
assert should_skip_implicit_mpm_masked_reset(solver, world_mask) is False
653+
654+
655+
def test_adapt_world_mask_for_implicit_mpm_full_selection_becomes_unmasked():
656+
"""Selecting every world is equivalent to an unmasked Implicit MPM reset."""
657+
from isaaclab_newton.physics.mpm_manager import (
658+
adapt_world_mask_for_implicit_mpm,
659+
should_skip_implicit_mpm_masked_reset,
660+
)
661+
662+
solver = SimpleNamespace(model=SimpleNamespace(world_count=3), _separate_worlds=False)
663+
world_mask = wp.array([True, True, True], dtype=wp.bool, device="cpu")
664+
665+
assert should_skip_implicit_mpm_masked_reset(solver, world_mask) is False
666+
assert adapt_world_mask_for_implicit_mpm(solver, world_mask) is None
667+
assert adapt_world_mask_for_implicit_mpm(solver, None) is None
668+
assert should_skip_implicit_mpm_masked_reset(solver, None) is False
669+
670+
empty = wp.array([False, False, False], dtype=wp.bool, device="cpu")
671+
assert should_skip_implicit_mpm_masked_reset(solver, empty) is True
672+
673+
635674
def test_cuda_graph_capture_uses_simulation_device(monkeypatch):
636675
"""CUDA graph capture should use the simulation device instead of Warp's default device."""
637676
from isaaclab.physics import PhysicsManager

0 commit comments

Comments
 (0)