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
2 changes: 1 addition & 1 deletion docs/sphinx/source/zh_CN/5-reference/5-support_matrix.md
Original file line number Diff line number Diff line change
Expand Up @@ -100,7 +100,7 @@ uv run scripts/generate_support_matrix.py --write
| APPO (torch) | `g1_climb_tracking` (g1 climb tracking) | Tested | - | Tested |
| SAC (torch) | `g1_walk_flat` (G1 walk flat) | Tested | Tested | Tested |
| SAC (torch) | `g1_walk_rough` (G1 walk rough) | Tested | - | Tested |
| SAC (torch) | `g1_motion_tracking` (G1 motion tracking) | Tested | - | Tested |
| SAC (torch) | `g1_motion_tracking` (G1 motion tracking) | Tested | Configured | Tested |
| SAC (torch) | `g1_flip_tracking` (G1 flip tracking) | Tested | - | Registered |
| SAC (torch) | `g1_wall_flip_tracking` (G1 wall flip tracking) | Tested | - | Registered |
| SAC (torch) | `g1_23dof_flip_tracking` (g1 23dof flip tracking) | Tested | - | Registered |
Expand Down
41 changes: 41 additions & 0 deletions scripts/benchmark/rl/benchmark_offpolicy_collector_active.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,6 +127,9 @@
"set_state_qpos_convert_ms",
"set_state_pool_reset_ms",
"set_state_state_scatter_ms",
"set_state_reset_upload_ms",
"set_state_reset_forward_ms",
"set_state_host_cache_refresh_ms",
"set_state_internal_gap_ms",
)
NP_ENV_STEP_COUNT_KEYS = ("reset_done_count",)
Expand Down Expand Up @@ -183,6 +186,9 @@
("set_state_qpos_convert_ms", "set_state_qpos_convert_ms"),
("set_state_pool_reset_ms", "set_state_pool_reset_ms"),
("set_state_state_scatter_ms", "set_state_state_scatter_ms"),
("set_state_reset_upload_ms", "set_state_reset_upload_ms"),
("set_state_reset_forward_ms", "set_state_reset_forward_ms"),
("set_state_host_cache_refresh_ms", "set_state_host_cache_refresh_ms"),
("set_state_internal_gap_ms", "set_state_internal_gap_ms"),
)

Expand Down Expand Up @@ -1315,6 +1321,13 @@ def _format_dr_reset_timing_table(results: list[CollectorResult]) -> str:
("set_state_internal_gap_ms", "Gap"),
)

_SET_STATE_MJWARP_KEYS = (
("set_state_reset_upload_ms", "Reset upload"),
("set_state_reset_forward_ms", "Reset forward"),
("set_state_host_cache_refresh_ms", "Host cache refresh"),
("set_state_internal_gap_ms", "Gap"),
)


def _format_set_state_detail_table(results: list[CollectorResult]) -> str:
"""Backend set_state sub-timing table (motrix keyset).
Expand Down Expand Up @@ -1375,6 +1388,32 @@ def _format_set_state_mujoco_table(results: list[CollectorResult]) -> str:
return _format_table(headers, rows)


def _format_set_state_mjwarp_table(results: list[CollectorResult]) -> str:
"""Backend set_state sub-timing table (mjwarp keyset)."""
headers = (
"Algo",
"Task",
"Backend",
"Set state ms (%env, %active)",
*(label for _, label in _SET_STATE_MJWARP_KEYS),
)
rows = []
for result in results:
env_step = result.phase_ms_per_vector_step.get("env_step_ms")
if env_step is None:
continue
rows.append(
(
result.case.algo,
result.case.task,
result.case.runtime_sim_backend,
_format_np_env_timing(result, "dr_reset_set_state_ms"),
*(_format_set_state_sub_ms(result, key) for key, _ in _SET_STATE_MJWARP_KEYS),
)
)
return _format_table(headers, rows)


def _format_np_env_step_timing_table(results: list[CollectorResult]) -> str:
headers = (
"Algo",
Expand Down Expand Up @@ -1750,6 +1789,8 @@ def main() -> int:
print(_format_set_state_detail_table(results))
print("\nBackend set_state detail — mujoco keyset:")
print(_format_set_state_mujoco_table(results))
print("\nBackend set_state detail — mjwarp keyset:")
print(_format_set_state_mjwarp_table(results))
else:
print("No successful benchmark cases.")
return 0 if not errors else 1
Expand Down
10 changes: 10 additions & 0 deletions src/unilab/base/backend/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -823,6 +823,16 @@ def get_body_pose_w_rows(
rows = np.asarray(env_ids, dtype=np.intp)
return self.get_body_pos_w(body_ids)[rows], self.get_body_quat_w(body_ids)[rows]

def get_body_lin_vel_w_rows(self, env_ids: np.ndarray, body_ids: np.ndarray) -> np.ndarray:
"""Get selected env rows of world-frame body linear velocity."""
rows = np.asarray(env_ids, dtype=np.intp)
return self.get_body_lin_vel_w(body_ids)[rows]

def get_body_ang_vel_w_rows(self, env_ids: np.ndarray, body_ids: np.ndarray) -> np.ndarray:
"""Get selected env rows of world-frame body angular velocity."""
rows = np.asarray(env_ids, dtype=np.intp)
return self.get_body_ang_vel_w(body_ids)[rows]

# ------------------------------------------------------------------ #
# Body kinematics — baselink frame #
# ------------------------------------------------------------------ #
Expand Down
49 changes: 42 additions & 7 deletions src/unilab/base/backend/mjwarp/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -952,6 +952,26 @@ def step(self, ctrl: np.ndarray, nsteps: int = 1) -> dict[str, dict[str, float]]
timings = self._execute_host_step(ctrl_array, int(nsteps))
return {"timing": timings}

# All backends report the same set_state key set for column stability;
# sub-keys that don't apply to the mjwarp host profile report 0.0.
_SET_STATE_TIMING_ZERO_KEYS = (
"set_state_mask_ms",
"set_state_data_slice_ms",
"set_state_data_reset_ms",
"set_state_clear_forces_ms",
"set_state_geom_overrides_ms",
"set_state_reset_rand_ms",
"set_state_set_dof_vel_ms",
"set_state_set_dof_pos_ms",
"set_state_actuator_ctrl_ms",
"set_state_forward_kinematic_ms",
"set_state_refresh_pose_cache_ms",
"set_state_invalidate_velocity_ms",
"set_state_qpos_convert_ms",
"set_state_pool_reset_ms",
"set_state_state_scatter_ms",
)

def set_state(
self,
env_indices: np.ndarray,
Expand All @@ -974,9 +994,19 @@ def set_state(
raise ValueError(f"qpos must have shape {expected_qpos}, got {qpos_array.shape}")
if qvel_array.shape != expected_qvel:
raise ValueError(f"qvel must have shape {expected_qvel}, got {qvel_array.shape}")
timing: dict[str, float] = {key: 0.0 for key in self._SET_STATE_TIMING_ZERO_KEYS}
timing.update(
{
"set_state_reset_upload_ms": 0.0,
"set_state_reset_forward_ms": 0.0,
"set_state_host_cache_refresh_ms": 0.0,
"set_state_internal_gap_ms": 0.0,
}
)
if rows.size == 0:
return {"timing": {"set_state_reset_ms": 0.0, "set_state_cache_refresh_ms": 0.0}}
return {"timing": timing}

outer_t0 = time.perf_counter()
self._qpos_cache[rows] = qpos_array
self._qvel_cache[rows] = qvel_array
timings = self._execute_host_reset(
Expand All @@ -986,12 +1016,17 @@ def set_state(
qpos_array,
qvel_array,
)
return {
"timing": {
"set_state_reset_ms": timings["reset_upload_ms"] + timings["reset_forward_ms"],
"set_state_cache_refresh_ms": timings["host_cache_refresh_ms"],
}
}
timing["set_state_reset_upload_ms"] = timings["reset_upload_ms"]
timing["set_state_reset_forward_ms"] = timings["reset_forward_ms"]
timing["set_state_host_cache_refresh_ms"] = timings["host_cache_refresh_ms"]
outer_total_ms = (time.perf_counter() - outer_t0) * 1000.0
measured_ms = (
timing["set_state_reset_upload_ms"]
+ timing["set_state_reset_forward_ms"]
+ timing["set_state_host_cache_refresh_ms"]
)
timing["set_state_internal_gap_ms"] = outer_total_ms - measured_ms
return {"timing": timing}

def get_dr_capabilities(self) -> DomainRandomizationCapabilities:
"""Advertise no legacy DR until per-world model mutation is effect-tested."""
Expand Down
11 changes: 11 additions & 0 deletions src/unilab/base/backend/motrix/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -739,6 +739,9 @@ def set_state(
"set_state_qpos_convert_ms": 0.0,
"set_state_pool_reset_ms": 0.0,
"set_state_state_scatter_ms": 0.0,
"set_state_reset_upload_ms": 0.0,
"set_state_reset_forward_ms": 0.0,
"set_state_host_cache_refresh_ms": 0.0,
"set_state_internal_gap_ms": 0.0,
}
outer_t0 = time.perf_counter()
Expand Down Expand Up @@ -1164,6 +1167,14 @@ def get_body_vel_w(self, body_ids: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
velocities = np.ascontiguousarray(self._ensure_link_velocity_cache()[:, ids, :])
return velocities[:, :, :3], velocities[:, :, 3:]

def get_body_lin_vel_w_rows(self, env_ids: np.ndarray, body_ids: np.ndarray) -> np.ndarray:
rows = np.asarray(env_ids, dtype=np.intp)
return self._ensure_link_velocity_cache()[rows[:, None], self._as_body_ids(body_ids), :3] # type: ignore[no-any-return]

def get_body_ang_vel_w_rows(self, env_ids: np.ndarray, body_ids: np.ndarray) -> np.ndarray:
rows = np.asarray(env_ids, dtype=np.intp)
return self._ensure_link_velocity_cache()[rows[:, None], self._as_body_ids(body_ids), 3:] # type: ignore[no-any-return]

# ------------------------------------------------------------------ #
# Body kinematics — baselink frame #
# ------------------------------------------------------------------ #
Expand Down
11 changes: 11 additions & 0 deletions src/unilab/base/backend/mujoco/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -1081,6 +1081,9 @@ def set_state(
"set_state_qpos_convert_ms": 0.0,
"set_state_pool_reset_ms": 0.0,
"set_state_state_scatter_ms": 0.0,
"set_state_reset_upload_ms": 0.0,
"set_state_reset_forward_ms": 0.0,
"set_state_host_cache_refresh_ms": 0.0,
"set_state_internal_gap_ms": 0.0,
}
if len(env_indices) == 0:
Expand Down Expand Up @@ -1417,6 +1420,14 @@ def get_body_lin_vel_w(self, body_ids: np.ndarray) -> np.ndarray:
def get_body_ang_vel_w(self, body_ids: np.ndarray) -> np.ndarray:
return self._tracked_angvel_w_all[:, self._get_mapped_indices(body_ids), :] # type: ignore[no-any-return]

def get_body_lin_vel_w_rows(self, env_ids: np.ndarray, body_ids: np.ndarray) -> np.ndarray:
rows = np.asarray(env_ids, dtype=np.intp)
return self._tracked_linvel_w_all[rows[:, None], self._get_mapped_indices(body_ids)] # type: ignore[no-any-return]

def get_body_ang_vel_w_rows(self, env_ids: np.ndarray, body_ids: np.ndarray) -> np.ndarray:
rows = np.asarray(env_ids, dtype=np.intp)
return self._tracked_angvel_w_all[rows[:, None], self._get_mapped_indices(body_ids)] # type: ignore[no-any-return]

def get_body_state_w(
self, body_ids: np.ndarray
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
Expand Down
20 changes: 20 additions & 0 deletions src/unilab/base/entity.py
Original file line number Diff line number Diff line change
Expand Up @@ -421,6 +421,26 @@ def body_link_ang_vel_w(self) -> np.ndarray:
def body_link_pose_w(self) -> np.ndarray:
return np.concatenate((self.body_link_pos_w, self.body_link_quat_w), axis=-1)

def body_link_pos_w_rows(self, env_ids: np.ndarray) -> np.ndarray:
"""Row-scoped variant of body_link_pos_w for partial-reset rebuilds."""
ids = self._require(self._body_ids, "body state")
return self._backend.get_body_pose_w_rows(env_ids, ids)[0]

def body_link_quat_w_rows(self, env_ids: np.ndarray) -> np.ndarray:
"""Row-scoped variant of body_link_quat_w for partial-reset rebuilds."""
ids = self._require(self._body_ids, "body state")
return self._backend.get_body_pose_w_rows(env_ids, ids)[1]

def body_link_lin_vel_w_rows(self, env_ids: np.ndarray) -> np.ndarray:
"""Row-scoped variant of body_link_lin_vel_w for partial-reset rebuilds."""
ids = self._require(self._body_ids, "body state")
return self._backend.get_body_lin_vel_w_rows(env_ids, ids)

def body_link_ang_vel_w_rows(self, env_ids: np.ndarray) -> np.ndarray:
"""Row-scoped variant of body_link_ang_vel_w for partial-reset rebuilds."""
ids = self._require(self._body_ids, "body state")
return self._backend.get_body_ang_vel_w_rows(env_ids, ids)

@property
def body_link_vel_w(self) -> np.ndarray:
return np.concatenate((self.body_link_lin_vel_w, self.body_link_ang_vel_w), axis=-1)
Expand Down
24 changes: 22 additions & 2 deletions src/unilab/base/np_env.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,9 @@
"set_state_qpos_convert_ms",
"set_state_pool_reset_ms",
"set_state_state_scatter_ms",
"set_state_reset_upload_ms",
"set_state_reset_forward_ms",
"set_state_host_cache_refresh_ms",
"set_state_internal_gap_ms",
)

Expand All @@ -84,6 +87,9 @@
"set_state_qpos_convert_ms",
"set_state_pool_reset_ms",
"set_state_state_scatter_ms",
"set_state_reset_upload_ms",
"set_state_reset_forward_ms",
"set_state_host_cache_refresh_ms",
"set_state_internal_gap_ms",
)

Expand Down Expand Up @@ -287,8 +293,10 @@ def _reset_done_envs(self) -> None:
finally:
self._autoreset_reset_active = False
detail_timing["reset_done_reset_call_ms"] = (time.perf_counter() - t0) * 1000.0
if self._dr_manager is not None:
detail_timing.update(self._dr_manager.last_reset_timing_ms)
collected = self._collect_reset_backend_timing_ms()
detail_timing.update(
{key: value for key, value in collected.items() if key in detail_timing}
)
t0 = time.perf_counter()
for key in self._state.obs:
self._state.obs[key][env_indices] = new_obs[key]
Expand Down Expand Up @@ -324,6 +332,18 @@ def _clear_reset_done_detail_timing(self, timing: dict[str, Any]) -> None:
for key in RESET_DONE_DETAIL_TIMING_KEYS:
timing[key] = 0.0

def _collect_reset_backend_timing_ms(self) -> dict[str, float]:
"""Backend-sourced reset sub-timings for the last reset call.

The monolithic DR path reports through the DR manager; manager-based
envs override this to surface the reset-state transaction's set_state
timings. Keys outside RESET_DONE_DETAIL_TIMING_KEYS are dropped by the
caller so stale keys never leak into ``info["timing"]``.
"""
if self._dr_manager is not None:
return self._dr_manager.last_reset_timing_ms
return {}

def _resolve_nan_guard_model_file(self) -> str:
scene = getattr(self._cfg, "scene", None)
if isinstance(scene, SceneCfg) and scene.model_file:
Expand Down
26 changes: 25 additions & 1 deletion src/unilab/base/reset_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@

from __future__ import annotations

import time
from collections.abc import Iterator
from contextlib import contextmanager

Expand Down Expand Up @@ -55,6 +56,7 @@ def __init__(
self._randomization_dirty_masks: dict[str, np.ndarray] = {}
self._requesting_terms: set[str] = set()
self._last_commit_had_writes = False
self._last_set_state_timing_ms: dict[str, float] = {}

@property
def active(self) -> bool:
Expand All @@ -66,6 +68,17 @@ def last_commit_had_writes(self) -> bool:
"""Whether the most recent scoped commit submitted dirty rows to set_state."""
return self._last_commit_had_writes

@property
def last_set_state_timing_ms(self) -> dict[str, float]:
"""Sub-timings from the most recent commit's set_state call.

Always includes ``dr_reset_set_state_ms`` (outer wall-clock around the
backend call); backend-reported ``set_state_*_ms`` sub-keys are merged
in when the backend returns them. Empty when the last commit had no
dirty rows.
"""
return self._last_set_state_timing_ms

@contextmanager
def scoped(self, env_ids: np.ndarray) -> Iterator[ResetStateTransaction]:
"""Begin a reset transaction and commit it only after all terms succeed."""
Expand All @@ -91,6 +104,7 @@ def begin(self, env_ids: np.ndarray) -> None:
mask.fill(False)
self._requesting_terms.clear()
self._last_commit_had_writes = False
self._last_set_state_timing_ms = {}
self._active = True

def bind_body_mass_write(
Expand Down Expand Up @@ -541,12 +555,22 @@ def commit(self) -> dict | None:
assert self._qvel is not None
randomization = self._build_randomization_payload(dirty_ids)
try:
return self._backend.set_state(
set_state_t0 = time.perf_counter()
result = self._backend.set_state(
dirty_ids,
self._qpos[dirty_ids],
self._qvel[dirty_ids],
randomization=randomization,
)
timing: dict[str, float] = {
"dr_reset_set_state_ms": (time.perf_counter() - set_state_t0) * 1000.0
}
if isinstance(result, dict):
backend_timing = result.get("timing")
if isinstance(backend_timing, dict):
timing.update(backend_timing)
self._last_set_state_timing_ms = timing
return result
except (AttributeError, NotImplementedError) as exc:
terms = ", ".join(sorted(self._requesting_terms))
raise NotImplementedError(
Expand Down
Loading