Skip to content

Commit ffa8a73

Browse files
committed
Skip empty eager reset pipelines
Keep reset masks canonical while using one explicit host predicate to avoid dispatching every reset manager for an empty selection. This recovers the eager hot path until recorded reset launches can remove the boundary.
1 parent f8180c9 commit ffa8a73

2 files changed

Lines changed: 22 additions & 0 deletions

File tree

source/isaaclab_experimental/isaaclab_experimental/envs/manager_based_rl_env_warp.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -580,6 +580,12 @@ def _reset_idx(
580580
def _reset_terminated_envs(self) -> None:
581581
"""Reset terminated environments, compacting IDs only for host consumers."""
582582
reset_mask = self.termination_manager.dones_wp
583+
# The eager reset pipeline contains many small launches. Keep the mask as
584+
# the canonical selection, but use one explicit host predicate to avoid
585+
# dispatching the entire pipeline when it is empty. Record/replay can
586+
# remove this boundary once reset stages have replay-safe launch caches.
587+
if not self.reset_buf.any().item():
588+
return
583589
reset_env_ids = None
584590
if self._reset_requires_host_selection():
585591
with Timer(

source/isaaclab_experimental/test/envs/test_manager_based_rl_env_warp.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -190,3 +190,19 @@ def test_empty_reset_skips_legacy_host_boundary_and_preserves_logs() -> None:
190190
env._reset_host_pre.assert_not_called()
191191
env._reset_host_post.assert_not_called()
192192
assert env.extras["log"] == {"previous": 1.0}
193+
194+
195+
def test_empty_reset_skips_warp_pipeline_without_compacting_ids() -> None:
196+
"""An empty mask should skip eager reset stages without materializing IDs."""
197+
env = ManagerBasedRLEnvWarp.__new__(ManagerBasedRLEnvWarp)
198+
reset_mask = wp.zeros(3, dtype=wp.bool, device="cpu")
199+
env.termination_manager = SimpleNamespace(dones_wp=reset_mask)
200+
env.reset_buf = torch.zeros(3, dtype=torch.bool)
201+
env._reset_requires_host_selection = Mock(return_value=False)
202+
env._reset_idx = Mock()
203+
env.extras = {"log": {"previous": 1.0}}
204+
205+
env._reset_terminated_envs()
206+
207+
env._reset_idx.assert_not_called()
208+
assert env.extras["log"] == {"previous": 1.0}

0 commit comments

Comments
 (0)