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
128 changes: 3 additions & 125 deletions scripts/benchmark/rl/benchmark_offpolicy_collector_active.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,19 +101,6 @@
"dr_reset_obs_get_body_pose_ms",
"dr_reset_observation_compute_obs_ms",
"dr_reset_observation_internal_gap_ms",
# Manager-Based (MBA) reset sub-steps (see RESET_DONE_DETAIL_TIMING_KEYS in
# src/unilab/base/np_env.py). Task-dependent per-term event timings with the
# mba_reset_event_term_<name>_ms prefix are discovered dynamically at sample
# time (JSON/console only; not part of the fixed CSV header).
"mba_reset_total_ms",
"mba_reset_curriculum_ms",
"mba_reset_event_apply_ms",
"mba_reset_command_reset_ms",
"mba_reset_set_state_ms",
"mba_reset_manager_reset_ms",
"mba_reset_command_compute_ms",
"mba_reset_obs_build_ms",
"mba_reset_internal_gap_ms",
# Backend set_state internals (see BACKEND_SET_STATE_DETAIL_TIMING_KEYS in
# src/unilab/base/np_env.py). All backends emit the same key set for column
# stability; sub-keys that don't apply report 0.0.
Expand All @@ -133,42 +120,9 @@
"set_state_pool_reset_ms",
"set_state_state_scatter_ms",
"set_state_internal_gap_ms",
# Manager-Based (MBA) update_state blocks, reported every step by
# ManagerBasedRlEnv.update_state() (see UPDATE_STATE_DETAIL_TIMING_KEYS in
# src/unilab/base/np_env.py). Task-dependent per-term and per-method keys
# (mba_obs_term_*, mba_reward_term_*, mba_getter_<method>_ms) are discovered
# dynamically at sample time (JSON/console only; not part of the fixed CSV
# header).
"mba_update_total_ms",
"mba_termination_ms",
"mba_reward_ms",
"mba_metrics_ms",
"mba_events_ms",
"mba_command_ms",
"mba_obs_compute_ms",
"mba_obs_map_ms",
"mba_termination_getter_ms",
"mba_reward_getter_ms",
"mba_metrics_getter_ms",
"mba_events_getter_ms",
"mba_command_getter_ms",
"mba_obs_compute_getter_ms",
"mba_obs_map_getter_ms",
"mba_update_internal_gap_ms",
"mba_getters_total_ms",
"mba_obs_noise_ms",
"mba_obs_clip_scale_ms",
"mba_obs_nan_check_ms",
"mba_obs_delay_ms",
"mba_obs_history_ms",
"mba_obs_concat_ms",
)
NP_ENV_STEP_COUNT_KEYS = ("reset_done_count",)
NP_ENV_STEP_SAMPLE_KEYS = (*NP_ENV_STEP_TIMING_KEYS, *NP_ENV_STEP_COUNT_KEYS)
# Task-dependent MBA timings emitted by ManagerBasedRlEnv (per-term reset event
# timings, per-term obs/reward timings, per-method getter timings); discovered
# dynamically from info["timing"] since term/method names vary per task.
MBA_DYNAMIC_KEY_PREFIX = "mba_"
NP_RANDOM_PROFILE_FUNCTIONS = (
"uniform",
"randint",
Expand Down Expand Up @@ -206,15 +160,6 @@
("dr_reset_obs_get_body_pose_ms", "dr_reset_obs_get_body_pose_ms"),
("dr_reset_observation_compute_obs_ms", "dr_reset_observation_compute_obs_ms"),
("dr_reset_observation_internal_gap_ms", "dr_reset_observation_internal_gap_ms"),
("mba_reset_total_ms", "mba_reset_total_ms"),
("mba_reset_curriculum_ms", "mba_reset_curriculum_ms"),
("mba_reset_event_apply_ms", "mba_reset_event_apply_ms"),
("mba_reset_command_reset_ms", "mba_reset_command_reset_ms"),
("mba_reset_set_state_ms", "mba_reset_set_state_ms"),
("mba_reset_manager_reset_ms", "mba_reset_manager_reset_ms"),
("mba_reset_command_compute_ms", "mba_reset_command_compute_ms"),
("mba_reset_obs_build_ms", "mba_reset_obs_build_ms"),
("mba_reset_internal_gap_ms", "mba_reset_internal_gap_ms"),
("set_state_mask_ms", "set_state_mask_ms"),
("set_state_data_slice_ms", "set_state_data_slice_ms"),
("set_state_data_reset_ms", "set_state_data_reset_ms"),
Expand All @@ -231,29 +176,6 @@
("set_state_pool_reset_ms", "set_state_pool_reset_ms"),
("set_state_state_scatter_ms", "set_state_state_scatter_ms"),
("set_state_internal_gap_ms", "set_state_internal_gap_ms"),
("mba_update_total_ms", "mba_update_total_ms"),
("mba_termination_ms", "mba_termination_ms"),
("mba_reward_ms", "mba_reward_ms"),
("mba_metrics_ms", "mba_metrics_ms"),
("mba_events_ms", "mba_events_ms"),
("mba_command_ms", "mba_command_ms"),
("mba_obs_compute_ms", "mba_obs_compute_ms"),
("mba_obs_map_ms", "mba_obs_map_ms"),
("mba_termination_getter_ms", "mba_termination_getter_ms"),
("mba_reward_getter_ms", "mba_reward_getter_ms"),
("mba_metrics_getter_ms", "mba_metrics_getter_ms"),
("mba_events_getter_ms", "mba_events_getter_ms"),
("mba_command_getter_ms", "mba_command_getter_ms"),
("mba_obs_compute_getter_ms", "mba_obs_compute_getter_ms"),
("mba_obs_map_getter_ms", "mba_obs_map_getter_ms"),
("mba_update_internal_gap_ms", "mba_update_internal_gap_ms"),
("mba_getters_total_ms", "mba_getters_total_ms"),
("mba_obs_noise_ms", "mba_obs_noise_ms"),
("mba_obs_clip_scale_ms", "mba_obs_clip_scale_ms"),
("mba_obs_nan_check_ms", "mba_obs_nan_check_ms"),
("mba_obs_delay_ms", "mba_obs_delay_ms"),
("mba_obs_history_ms", "mba_obs_history_ms"),
("mba_obs_concat_ms", "mba_obs_concat_ms"),
)


Expand Down Expand Up @@ -685,13 +607,6 @@ def _run_active_window_case(
)
else:
env_step_timing_values["env_step_internal_gap_ms"] = None
# Task-dependent MBA dynamic keys (reset event terms, obs/reward
# terms, per-method getter timings); zero-filled by NpEnv on steps
# where they don't apply, once the key set has been discovered.
for key, value in _timing.items():
if key.startswith(MBA_DYNAMIC_KEY_PREFIX) and key not in env_step_timing_values:
env_step_timing_values[key] = value

phase_start_ns = time.perf_counter_ns()
next_obs_np, next_critic_np = split_obs_dict(state.obs)
next_obs_np = np.asarray(next_obs_np, dtype=np.float32)
Expand Down Expand Up @@ -1223,17 +1138,14 @@ def _format_np_env_value(result: CollectorResult, key: str, *, digits: int = 1)
def _format_set_state_sub_ms(result: CollectorResult, key: str) -> str:
"""Format a backend set_state sub-timing as ``ms (%of set_state)``.

The percentage is relative to the outer set_state wall-clock measurement —
``dr_reset_set_state_ms`` on the direct path, ``mba_reset_set_state_ms`` on
the MBA path — not to env_step_ms, so a reader can see which sub-step
dominates set_state.
The percentage is relative to ``dr_reset_set_state_ms`` (the outer
wall-clock measurement in DomainRandomizationManager), not to env_step_ms,
so a reader can see which sub-step dominates set_state.
"""
stat = result.env_step_timing_ms_per_vector_step.get(key)
if stat is None:
return "n/a"
outer = result.env_step_timing_ms_per_vector_step.get("dr_reset_set_state_ms")
if outer is None or outer.mean_ms <= 0.0:
outer = result.env_step_timing_ms_per_vector_step.get("mba_reset_set_state_ms")
if outer is None or outer.mean_ms <= 0.0:
return f"{stat.mean_ms:.3f} (n/a)"
pct = 100.0 * stat.mean_ms / outer.mean_ms
Expand Down Expand Up @@ -1621,29 +1533,6 @@ def _print_result(result: CollectorResult) -> None:
"dr_reset_observation_getters_ms",
"dr_reset_obs_get_body_pose_ms",
"dr_reset_observation_compute_obs_ms",
"mba_reset_total_ms",
"mba_reset_curriculum_ms",
"mba_reset_event_apply_ms",
"mba_reset_command_reset_ms",
"mba_reset_set_state_ms",
"mba_reset_manager_reset_ms",
"mba_reset_command_compute_ms",
"mba_reset_obs_build_ms",
"mba_reset_internal_gap_ms",
"mba_update_total_ms",
"mba_termination_ms",
"mba_reward_ms",
"mba_metrics_ms",
"mba_events_ms",
"mba_command_ms",
"mba_obs_compute_ms",
"mba_obs_map_ms",
"mba_update_internal_gap_ms",
"mba_getters_total_ms",
"mba_obs_noise_ms",
"mba_obs_clip_scale_ms",
"mba_obs_nan_check_ms",
"mba_obs_concat_ms",
):
stat = result.env_step_timing_ms_per_vector_step.get(key)
if stat is None:
Expand All @@ -1653,17 +1542,6 @@ def _print_result(result: CollectorResult) -> None:
f"pct_env={_env_step_child_env_pct(result, stat.mean_ms):5.1f}% "
f"pct_active={_env_step_child_pct(result, stat.mean_ms):5.1f}%"
)
# Task-dependent MBA dynamic keys (reset event terms, obs/reward terms,
# per-method getter timings) not already printed above.
for key in sorted(result.env_step_timing_ms_per_vector_step):
if not key.startswith(MBA_DYNAMIC_KEY_PREFIX) or key in NP_ENV_STEP_SAMPLE_KEYS:
continue
stat = result.env_step_timing_ms_per_vector_step[key]
print(
f" {('np_env_' + key):<18} mean={stat.mean_ms:8.3f} ms "
f"pct_env={_env_step_child_env_pct(result, stat.mean_ms):5.1f}% "
f"pct_active={_env_step_child_pct(result, stat.mean_ms):5.1f}%"
)


def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace:
Expand Down
58 changes: 14 additions & 44 deletions src/unilab/base/entity.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,6 @@
from __future__ import annotations

import re
import time
from collections.abc import Iterator, Mapping, Sequence
from contextlib import contextmanager
from dataclasses import dataclass
Expand Down Expand Up @@ -134,28 +133,6 @@ def _resolve_matching_names(
return [item[1] for item in matches], [item[2] for item in matches]


class GetterTimingRecorder:
"""Per-step wall-time accumulator for leaf backend state getters.

One recorder per :class:`EntityData` instance; the owning env aggregates and
resets the recorders around its ``update_state`` pass (see
``ManagerBasedRlEnv``). Timing is always on for actual backend reads: cache
hits add no synthetic getter time.
"""

def __init__(self) -> None:
self.method_ms: dict[str, float] = {}
self.total_ms: float = 0.0

def record(self, method: str, elapsed_ms: float) -> None:
self.method_ms[method] = self.method_ms.get(method, 0.0) + elapsed_ms
self.total_ms += elapsed_ms

def reset(self) -> None:
self.method_ms.clear()
self.total_ms = 0.0


_StateReadKey = tuple[str, tuple[int, ...] | None]


Expand Down Expand Up @@ -246,11 +223,8 @@ def __init__(
self._actuator_ctrl_range = actuator_ctrl_range
self._control_buffer = control_buffer
self._state_read_cache = state_read_cache
# Always-on leaf getter timing; aggregated/reset per update_state by the
# owning env (MBA update_state instrumentation, issue #1256).
self.getter_timing = GetterTimingRecorder()

def _timed_getter(
def _cached_getter(
self,
method: str,
fn: Any,
Expand All @@ -261,11 +235,7 @@ def _timed_getter(
cached = self._state_read_cache.get(key)
if cached is not None:
return cached
t0 = time.perf_counter_ns()
try:
value = fn(*args)
finally:
self.getter_timing.record(method, (time.perf_counter_ns() - t0) / 1e6)
value = fn(*args)
self._state_read_cache.put(key, value)
return value

Expand All @@ -280,7 +250,7 @@ def _require(self, value, capability: str):
@property
def root_link_pos_w(self) -> np.ndarray:
ids = self._require(self._root_body_ids, "root body state")
return self._timed_getter(
return self._cached_getter(
"body_pos_w",
self._backend.get_body_pos_w,
ids,
Expand All @@ -290,7 +260,7 @@ def root_link_pos_w(self) -> np.ndarray:
@property
def root_link_quat_w(self) -> np.ndarray:
ids = self._require(self._root_body_ids, "root body state")
return self._timed_getter(
return self._cached_getter(
"body_quat_w",
self._backend.get_body_quat_w,
ids,
Expand All @@ -300,7 +270,7 @@ def root_link_quat_w(self) -> np.ndarray:
@property
def root_link_lin_vel_w(self) -> np.ndarray:
ids = self._require(self._root_body_ids, "root body state")
return self._timed_getter(
return self._cached_getter(
"body_lin_vel_w",
self._backend.get_body_lin_vel_w,
ids,
Expand All @@ -310,7 +280,7 @@ def root_link_lin_vel_w(self) -> np.ndarray:
@property
def root_link_ang_vel_w(self) -> np.ndarray:
ids = self._require(self._root_body_ids, "root body state")
return self._timed_getter(
return self._cached_getter(
"body_ang_vel_w",
self._backend.get_body_ang_vel_w,
ids,
Expand All @@ -320,7 +290,7 @@ def root_link_ang_vel_w(self) -> np.ndarray:
@property
def root_link_lin_vel_b(self) -> np.ndarray:
ids = self._require(self._root_body_ids, "root body state")
return self._timed_getter(
return self._cached_getter(
"body_lin_vel_b",
self._backend.get_body_lin_vel_b,
ids,
Expand All @@ -330,7 +300,7 @@ def root_link_lin_vel_b(self) -> np.ndarray:
@property
def root_link_ang_vel_b(self) -> np.ndarray:
ids = self._require(self._root_body_ids, "root body state")
return self._timed_getter(
return self._cached_getter(
"body_ang_vel_b",
self._backend.get_body_ang_vel_b,
ids,
Expand Down Expand Up @@ -375,12 +345,12 @@ def default_root_state(self) -> np.ndarray:
@property
def joint_pos(self) -> np.ndarray:
index = self._require(self._joint_pos_index, "joint position")
return self._timed_getter("dof_pos", self._backend.get_dof_pos)[:, index]
return self._cached_getter("dof_pos", self._backend.get_dof_pos)[:, index]

@property
def joint_vel(self) -> np.ndarray:
index = self._require(self._joint_vel_index, "joint velocity")
return self._timed_getter("dof_vel", self._backend.get_dof_vel)[:, index]
return self._cached_getter("dof_vel", self._backend.get_dof_vel)[:, index]

@property
def joint_pos_biased(self) -> np.ndarray:
Expand Down Expand Up @@ -410,7 +380,7 @@ def encoder_bias(self) -> np.ndarray:
@property
def body_link_pos_w(self) -> np.ndarray:
ids = self._require(self._body_ids, "body state")
return self._timed_getter(
return self._cached_getter(
"body_pos_w",
self._backend.get_body_pos_w,
ids,
Expand All @@ -420,7 +390,7 @@ def body_link_pos_w(self) -> np.ndarray:
@property
def body_link_quat_w(self) -> np.ndarray:
ids = self._require(self._body_ids, "body state")
return self._timed_getter(
return self._cached_getter(
"body_quat_w",
self._backend.get_body_quat_w,
ids,
Expand All @@ -430,7 +400,7 @@ def body_link_quat_w(self) -> np.ndarray:
@property
def body_link_lin_vel_w(self) -> np.ndarray:
ids = self._require(self._body_ids, "body state")
return self._timed_getter(
return self._cached_getter(
"body_lin_vel_w",
self._backend.get_body_lin_vel_w,
ids,
Expand All @@ -440,7 +410,7 @@ def body_link_lin_vel_w(self) -> np.ndarray:
@property
def body_link_ang_vel_w(self) -> np.ndarray:
ids = self._require(self._body_ids, "body state")
return self._timed_getter(
return self._cached_getter(
"body_ang_vel_w",
self._backend.get_body_ang_vel_w,
ids,
Expand Down
Loading
Loading