From 89bcb1af4166de220c0ab7b1808032b3b7c87150 Mon Sep 17 00:00:00 2001 From: TATP-233 Date: Fri, 2 Oct 2026 22:42:17 +0800 Subject: [PATCH] feat(reset): publish call-graph diagnostics --- .../benchmark_offpolicy_collector_active.py | 9 +++- src/unilab/base/backend_timing.py | 5 +++ src/unilab/envs/manager_based_rl_env.py | 19 ++++++++- src/unilab/managers/command_manager.py | 12 ++++++ .../motion_tracking/common/manager_terms.py | 16 +++++++ tests/envs/test_env_configs.py | 42 +++++++++++++++++++ 6 files changed, 100 insertions(+), 3 deletions(-) diff --git a/scripts/benchmark/rl/benchmark_offpolicy_collector_active.py b/scripts/benchmark/rl/benchmark_offpolicy_collector_active.py index f8451edad..2de22c92a 100644 --- a/scripts/benchmark/rl/benchmark_offpolicy_collector_active.py +++ b/scripts/benchmark/rl/benchmark_offpolicy_collector_active.py @@ -158,7 +158,14 @@ "set_state_host_cache_refresh_ms", "set_state_internal_gap_ms", ) -ENV_STEP_COUNT_KEYS = ("reset_done_count",) +ENV_STEP_COUNT_KEYS = ( + "reset_done_count", + "reset_done_event_term_count", + "reset_done_command_term_count", + "reset_done_manager_reset_count", + "reset_done_observation_term_count", + "reset_done_sampler_host_transfer_count", +) ENV_STEP_SAMPLE_KEYS = (*ENV_STEP_TIMING_KEYS, *ENV_STEP_COUNT_KEYS) NP_RANDOM_PROFILE_FUNCTIONS = ( "uniform", diff --git a/src/unilab/base/backend_timing.py b/src/unilab/base/backend_timing.py index 386a40e1b..26a76b232 100644 --- a/src/unilab/base/backend_timing.py +++ b/src/unilab/base/backend_timing.py @@ -9,6 +9,11 @@ RESET_DONE_DETAIL_TIMING_KEYS = ( "reset_done_count", + "reset_done_event_term_count", + "reset_done_command_term_count", + "reset_done_manager_reset_count", + "reset_done_observation_term_count", + "reset_done_sampler_host_transfer_count", "reset_done_terminal_obs_ms", "reset_done_reset_call_ms", "reset_done_command_event_ms", diff --git a/src/unilab/envs/manager_based_rl_env.py b/src/unilab/envs/manager_based_rl_env.py index bcefad3af..53f18facf 100644 --- a/src/unilab/envs/manager_based_rl_env.py +++ b/src/unilab/envs/manager_based_rl_env.py @@ -1238,6 +1238,8 @@ def reset( else: reset_context = self._reset_state.scoped(rows) command_event_started = time.perf_counter() + event_term_count = len(self.event_manager.active_terms.get("reset", ())) + command_term_count = len(self.command_manager.active_terms) with reset_context: if "reset" in self.event_manager.available_modes: self.event_manager.apply( @@ -1247,6 +1249,7 @@ def reset( ) log.update(self.command_manager.reset(rows)) reset_timing.update(getattr(self.command_manager, "last_reset_timing_ms", {})) + reset_timing.update(self.command_manager.reset_diagnostics()) reset_commit_started = time.perf_counter() reset_timing["reset_done_reset_commit_ms"] = ( time.perf_counter() - reset_commit_started @@ -1256,7 +1259,7 @@ def reset( ) * 1000.0 manager_state_started = time.perf_counter() - for manager in ( + reset_managers = ( self.observation_manager, self.action_manager, self.reward_manager, @@ -1264,8 +1267,20 @@ def reset( self.curriculum_manager, self.event_manager, self.termination_manager, - ): + ) + for manager in reset_managers: log.update(manager.reset(rows)) + observation_term_count = sum( + len(terms) for terms in self.observation_manager.active_terms.values() + ) + reset_timing.update( + { + "reset_done_event_term_count": float(event_term_count), + "reset_done_command_term_count": float(command_term_count), + "reset_done_manager_reset_count": float(len(reset_managers)), + "reset_done_observation_term_count": float(observation_term_count), + } + ) reset_timing["reset_done_manager_state_ms"] = ( time.perf_counter() - manager_state_started ) * 1000.0 diff --git a/src/unilab/managers/command_manager.py b/src/unilab/managers/command_manager.py index b1cce4c30..b272cbc3e 100644 --- a/src/unilab/managers/command_manager.py +++ b/src/unilab/managers/command_manager.py @@ -405,6 +405,15 @@ def uses_tensor_reset_rows(self) -> bool: bool(getattr(term, "uses_tensor_reset_rows", False)) for term in self._terms.values() ) + def reset_diagnostics(self) -> dict[str, float]: + """Collect optional term-owned reset call-graph diagnostics.""" + diagnostics: dict[str, float] = {} + for term in self._terms.values(): + collector = getattr(term, "sampler_reset_diagnostics", None) + if callable(collector): + diagnostics.update(cast("dict[str, float]", collector())) + return diagnostics + def get_command(self, name: str) -> torch.Tensor: return self._validate_command(name, self._terms[name].command) @@ -477,6 +486,9 @@ def get_active_iterable_terms(self, env_idx: int) -> Sequence[tuple[str, Sequenc def reset(self, env_ids: torch.Tensor | None = None) -> dict[str, np.ndarray]: return {} + def reset_diagnostics(self) -> dict[str, float]: + return {} + def compute( self, dt: float | np.ndarray | torch.Tensor, env_ids: torch.Tensor | None = None ) -> None: diff --git a/src/unilab/tasks/motion_tracking/common/manager_terms.py b/src/unilab/tasks/motion_tracking/common/manager_terms.py index 624dd7784..d956e6f4c 100644 --- a/src/unilab/tasks/motion_tracking/common/manager_terms.py +++ b/src/unilab/tasks/motion_tracking/common/manager_terms.py @@ -497,6 +497,7 @@ def __init__(self, cfg: MotionCommandCfg, env: ManagerBasedRlEnv): self._robot_body_ang_vel_w = np.empty_like(self._body_pos_w) self._bind_read_phase = False self._tensor_all_rows = torch.arange(self.num_envs, dtype=torch.int64, device=self._device) + self._sampler_host_transfers = 0 for name in ( "error_anchor_pos", @@ -1040,8 +1041,12 @@ def reset(self, env_ids: torch.Tensor | slice | None) -> dict[str, float]: (rows.numel(), self.motion.num_joints), dtype=torch.float32, ) + self._sampler_host_transfers = 0 return CommandTerm.reset(self, rows) + def sampler_reset_diagnostics(self) -> dict[str, float]: + return {"reset_done_sampler_host_transfer_count": float(self._sampler_host_transfers)} + def _refresh_motion(self, env_ids: np.ndarray | None = None) -> None: del env_ids self._refresh_motion_torch() @@ -1067,6 +1072,7 @@ def uses_tensor_reset_rows(self) -> bool: def _resample_command(self, env_ids: torch.Tensor) -> None: host_rows = env_ids.detach().cpu().numpy() + self._sampler_host_transfers = getattr(self, "_sampler_host_transfers", 0) + 1 sampler_started = time.perf_counter() frames = self.sampler.sample_frames(host_rows) sampler_ms = (time.perf_counter() - sampler_started) * 1000.0 @@ -1194,6 +1200,7 @@ def _step_tensor_sampler(self) -> torch.Tensor: host_frames = frames.detach().cpu().numpy() host_clip_indices = clip_indices.detach().cpu().numpy() host_clip_ends = clip_ends.detach().cpu().numpy() + self._sampler_host_transfers = getattr(self, "_sampler_host_transfers", 0) + 3 self.sampler.current_frames[...] = host_frames self.sampler.current_clip_indices[...] = host_clip_indices self.sampler.current_clip_end_frames[...] = host_clip_ends @@ -1220,6 +1227,15 @@ def _sync_tensor_sampler_state(self) -> None: torch.as_tensor(self.sampler.current_clip_end_frames, device=self._device) ) + @property + def sampler_host_transfers(self) -> int: + """Count explicit sampler device-to-host transfers since the last read.""" + return self._sampler_host_transfers + + @sampler_host_transfers.setter + def sampler_host_transfers(self, value: int) -> None: + self._sampler_host_transfers = int(value) + @staticmethod def _validate_cfg(cfg: MotionCommandCfg) -> None: _validate_motion_command_cfg(cfg) diff --git a/tests/envs/test_env_configs.py b/tests/envs/test_env_configs.py index 58a7eda6a..dcf26ab2f 100644 --- a/tests/envs/test_env_configs.py +++ b/tests/envs/test_env_configs.py @@ -1211,3 +1211,45 @@ def test_g1_motion_core_manager_reset_and_step( assert all(torch.isfinite(values).all() for values in state.obs.values()) finally: env.close() + + +def test_flashsac_motion_reset_publishes_call_graph_counts() -> None: + """Selected-reset diagnostics count Manager dispatch, not only wall time.""" + ensure_registries() + _require_mjwarp_runtime() + from unilab.base import registry + from unilab.envs import ManagerBasedRlEnv + + _, override = _motion_manager_override( + "g1_motion_tracking", + "mjwarp", + config_root="flashsac", + ) + env = registry.make( + "G1MotionTrackingSAC", + num_envs=4, + sim_backend="mjwarp", + env_cfg_override=override, + ) + assert isinstance(env, ManagerBasedRlEnv) + try: + env.init_state() + # Force every row done so the selected-reset path executes deterministically. + original_compute = env._compute_truncated + + def terminate_all(state): + del state + return torch.ones((env.num_envs,), dtype=torch.bool, device=env.device) + + env._compute_truncated = terminate_all + state = env.step(torch.zeros((4, 29), dtype=torch.float32, device=env.device)) + env._compute_truncated = original_compute + timing = state.info["timing"] + assert timing["reset_done_event_term_count"] == 0.0 + assert timing["reset_done_command_term_count"] == 1.0 + assert timing["reset_done_manager_reset_count"] == 7.0 + assert timing["reset_done_observation_term_count"] == 17.0 + assert timing["reset_done_sampler_host_transfer_count"] >= 1.0 + finally: + env.close() + env._backend.close()