Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion scripts/benchmark/rl/benchmark_offpolicy_collector_active.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
5 changes: 5 additions & 0 deletions src/unilab/base/backend_timing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
19 changes: 17 additions & 2 deletions src/unilab/envs/manager_based_rl_env.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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
Expand All @@ -1256,16 +1259,28 @@ def reset(
) * 1000.0

manager_state_started = time.perf_counter()
for manager in (
reset_managers = (
self.observation_manager,
self.action_manager,
self.reward_manager,
self.metrics_manager,
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
Expand Down
12 changes: 12 additions & 0 deletions src/unilab/managers/command_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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:
Expand Down
16 changes: 16 additions & 0 deletions src/unilab/tasks/motion_tracking/common/manager_terms.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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()
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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)
Expand Down
42 changes: 42 additions & 0 deletions tests/envs/test_env_configs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()