Source code for unilab.base.np_env

from __future__ import annotations

import abc
import dataclasses
import time
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from os import PathLike
from typing import TYPE_CHECKING, Any, Optional, Tuple, cast

import gymnasium as gym
import numpy as np
from unisim.backend.base import (
    BackendPlayRenderPlan,
    CameraCfg,
    DebugOverlayGetter,
    SimBackend,
)

from unilab.base.base import ABEnv, EnvCfg, EnvPlayCapabilities
from unilab.base.cpu_runtime import apply_env_cpu_runtime
from unilab.base.scene import SceneCfg
from unilab.dr import DomainRandomizationManager, DomainRandomizationProvider
from unilab.dtype_config import get_global_dtype

if TYPE_CHECKING:
    from unilab.utils.nan_guard import NanGuard


RESET_DONE_DETAIL_TIMING_KEYS = (
    "reset_done_count",
    "reset_done_terminal_obs_ms",
    "reset_done_reset_call_ms",
    "reset_done_obs_scatter_ms",
    "reset_done_info_scatter_ms",
    "reset_done_internal_gap_ms",
    "dr_reset_total_ms",
    "dr_reset_plan_ms",
    "dr_reset_payload_filter_ms",
    "dr_reset_set_state_ms",
    "dr_reset_build_observation_ms",
    "dr_reset_internal_gap_ms",
    "dr_reset_observation_getters_ms",
    "dr_reset_obs_get_motion_ms",
    "dr_reset_obs_get_local_linvel_ms",
    "dr_reset_obs_get_gyro_ms",
    "dr_reset_obs_get_gravity_ms",
    "dr_reset_obs_get_dof_pos_ms",
    "dr_reset_obs_get_dof_vel_ms",
    "dr_reset_obs_get_body_pose_ms",
    "dr_reset_observation_compute_obs_ms",
    "dr_reset_observation_internal_gap_ms",
    # Backend-internal set_state sub-timings. All backends report the same key
    # set for column stability; sub-keys that don't apply report 0.0.
    "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",
    "set_state_reset_upload_ms",
    "set_state_reset_forward_ms",
    "set_state_host_cache_refresh_ms",
    "set_state_internal_gap_ms",
)

# Subset of RESET_DONE_DETAIL_TIMING_KEYS that comes from the backend's
# set_state() timing dict. DR manager merges these keys 1-to-1 from the
# backend return value.
BACKEND_SET_STATE_DETAIL_TIMING_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",
    "set_state_reset_upload_ms",
    "set_state_reset_forward_ms",
    "set_state_host_cache_refresh_ms",
    "set_state_internal_gap_ms",
)


[docs] @dataclass class NpEnvState: obs: dict[str, np.ndarray] reward: np.ndarray terminated: np.ndarray truncated: np.ndarray info: dict[str, Any] final_observation: dict[str, np.ndarray] | None = None
[docs] def replace(self, **updates: Any) -> "NpEnvState": return dataclasses.replace(self, **updates)
[docs] class NpEnv(ABEnv): """Backend-agnostic numpy environment base class."""
[docs] def __init__(self, cfg: EnvCfg, backend: SimBackend, num_envs: int): # Cold-path process confinement for envs that own an explicit CPU # block (multi-rank DP collectors): keeps host-side NumPy/Numba compute # inside the same CPUs the backend pool workers are pinned to. Runs # before managers/materialization so Numba's lazily-launched pool # inherits the confined mask. apply_env_cpu_runtime(cfg.cpu_ids) self._cfg = cfg self._backend: SimBackend = backend self._num_envs = num_envs self._state: Optional[NpEnvState] = None self._truncated_scratch: np.ndarray = np.zeros((self._num_envs,), dtype=bool) self._final_observation_scratch: dict[str, np.ndarray] | None = None self.step_counter = 0 self._dr_manager: DomainRandomizationManager | None = None self._init_randomization_applied = False self._nan_guard: NanGuard | None = None self._autoreset = True self._autoreset_reset_active = False self._nan_guard_model_file = self._resolve_nan_guard_model_file() self._rgb_array_renderer_ready = False
@property def cfg(self) -> EnvCfg: return self._cfg @property def num_envs(self) -> int: return self._num_envs @property def state(self) -> Optional[NpEnvState]: return self._state @property def obs_groups_spec(self) -> dict[str, int]: """Return observation group dimensions, e.g. {"obs": 98, "critic": 101}. Subclasses MUST override this property. """ raise NotImplementedError @property def observation_space(self) -> gym.Space: total = sum(self.obs_groups_spec.values()) return gym.spaces.Box(-np.inf, np.inf, shape=(total,), dtype=np.float64)
[docs] def init_state(self) -> NpEnvState: dtype = get_global_dtype() obs = { k: np.zeros((self._num_envs, d), dtype=dtype) for k, d in self.obs_groups_spec.items() } reward = np.zeros((self._num_envs,), dtype=dtype) terminated = np.ones((self._num_envs,), dtype=bool) truncated = np.zeros((self._num_envs,), dtype=bool) steps = self._initial_episode_steps() info: dict = {"steps": steps} self._state = NpEnvState(obs, reward, terminated, truncated, info) self._reset_done_envs() self._clear_step_final_observation() return self._state
def _initial_episode_steps(self) -> np.ndarray: """Return initial per-env episode counters. Existing monolithic tasks keep their randomized initialization. A lifecycle with different public semantics can override this cold-path hook without duplicating :meth:`init_state` or the autoreset machinery. """ if self._cfg.max_episode_steps: return np.random.randint( 0, self._cfg.max_episode_steps, size=(self._num_envs,), dtype=np.uint32, ) return np.zeros((self._num_envs,), dtype=np.uint32)
[docs] def step(self, actions: np.ndarray) -> NpEnvState: step_t0 = time.perf_counter() if self._state is None: self.init_state() assert self._state is not None t0 = time.perf_counter() ctrl = self.apply_action(actions, self._state) apply_action_time = time.perf_counter() - t0 if self._dr_manager is not None: self._dr_manager.apply_interval_randomization_if_due(self.step_counter) self._state.truncated.fill(False) self._clear_step_final_observation() if self._nan_guard is not None: bad_ctrl_ids = self._nan_guard.check_ctrl(ctrl, step=self.step_counter) if bad_ctrl_ids is not None: self._nan_guard.dump( bad_ctrl_ids, self._nan_guard_model_file, self.step_counter, ) t0 = time.perf_counter() backend_result = self._backend.step(ctrl, self._cfg.sim_substeps) step_core_time = time.perf_counter() - t0 t0 = time.perf_counter() self._state = self.update_state(self._state) update_state_time = time.perf_counter() - t0 self._state.info["steps"] += 1 self.step_counter += 1 truncated = self._compute_truncated(self._state) np.logical_or(self._state.truncated, truncated, out=self._state.truncated) done = self._state.terminated | self._state.truncated t0 = time.perf_counter() did_reset = self._autoreset and np.any(done) if did_reset: self._reset_done_envs() reset_done_time = time.perf_counter() - t0 timing = self._state.info.setdefault("timing", {}) if not did_reset: self._clear_reset_done_detail_timing(timing) timing["env_step_total_ms"] = (time.perf_counter() - step_t0) * 1000.0 timing["apply_action_ms"] = apply_action_time * 1000.0 timing["step_core_ms"] = step_core_time * 1000.0 timing["update_state_ms"] = update_state_time * 1000.0 timing["reset_done_ms"] = reset_done_time * 1000.0 if backend_result is not None: backend_timing = backend_result.get("timing") if backend_timing: for k, v in backend_timing.items(): timing[f"backend_{k}"] = v if self._nan_guard is not None: self._nan_guard.capture( self.get_physics_state_snapshot() if self.play_capabilities.supports_physics_state_playback else None ) nan_ids = self._nan_guard.check( self._state.obs, self._state.reward, step=self.step_counter ) if nan_ids is not None: self._nan_guard.dump(nan_ids, self._nan_guard_model_file, self.step_counter) np.nan_to_num(self._state.reward, copy=False, nan=0.0, posinf=0.0, neginf=0.0) return self._state
def _reset_done_envs(self) -> None: assert self._state is not None reset_t0 = time.perf_counter() detail_timing = {key: 0.0 for key in RESET_DONE_DETAIL_TIMING_KEYS} done = self._state.terminated | self._state.truncated if not np.any(done): timing = self._state.info.setdefault("timing", {}) self._clear_reset_done_detail_timing(timing) return t0 = time.perf_counter() env_indices = np.flatnonzero(done).astype(np.int32) detail_timing["reset_done_count"] = float(len(env_indices)) self._state.info["steps"][env_indices] = 0 final_observation = self._ensure_final_observation_scratch() compat_final_observation, compat_terminal_mask = ( self._ensure_final_observation_compat_buffers() ) compat_terminal_mask[:] = False compat_terminal_mask[env_indices] = True for key in self._state.obs: final_observation[key][env_indices] = self._state.obs[key][env_indices] compat_final_observation[key][env_indices] = final_observation[key][env_indices] self._state.final_observation = final_observation detail_timing["reset_done_terminal_obs_ms"] = (time.perf_counter() - t0) * 1000.0 t0 = time.perf_counter() self._autoreset_reset_active = True try: new_obs, info1 = self.reset(env_indices) finally: self._autoreset_reset_active = False detail_timing["reset_done_reset_call_ms"] = (time.perf_counter() - t0) * 1000.0 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] detail_timing["reset_done_obs_scatter_ms"] = (time.perf_counter() - t0) * 1000.0 t0 = time.perf_counter() if info1: for key, value in info1.items(): if key not in self._state.info: if isinstance(value, np.ndarray): full_shape = (self._num_envs,) + value.shape[1:] self._state.info[key] = np.zeros(full_shape, dtype=value.dtype) self._state.info[key][env_indices] = value else: self._state.info[key] = value elif isinstance(value, np.ndarray): self._state.info[key][env_indices] = value detail_timing["reset_done_info_scatter_ms"] = (time.perf_counter() - t0) * 1000.0 reset_done_total_ms = (time.perf_counter() - reset_t0) * 1000.0 measured_reset_done_ms = ( detail_timing["reset_done_terminal_obs_ms"] + detail_timing["reset_done_reset_call_ms"] + detail_timing["reset_done_obs_scatter_ms"] + detail_timing["reset_done_info_scatter_ms"] ) detail_timing["reset_done_internal_gap_ms"] = reset_done_total_ms - measured_reset_done_ms timing = self._state.info.setdefault("timing", {}) self._clear_reset_done_detail_timing(timing) timing.update(detail_timing) 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: return str(scene.model_file) model_file = self._backend.get_scene_model_file() if model_file: return str(model_file) return "" def _ensure_final_observation_scratch(self) -> dict[str, np.ndarray]: assert self._state is not None obs = self._state.obs scratch = self._final_observation_scratch if scratch is None or set(scratch) != set(obs): scratch = {key: np.zeros_like(value) for key, value in obs.items()} self._final_observation_scratch = scratch else: for key, value in obs.items(): if scratch[key].shape != value.shape or scratch[key].dtype != value.dtype: scratch = { obs_key: np.zeros_like(obs_value) for obs_key, obs_value in obs.items() } self._final_observation_scratch = scratch break assert scratch is not None return scratch def _ensure_final_observation_compat_buffers( self, ) -> tuple[dict[str, np.ndarray], np.ndarray]: assert self._state is not None obs = self._state.obs compat_final_observation = self._state.info.get("final_observation") if not isinstance(compat_final_observation, dict) or set(compat_final_observation) != set( obs ): compat_final_observation = {key: np.zeros_like(value) for key, value in obs.items()} self._state.info["final_observation"] = compat_final_observation else: for key, value in obs.items(): if ( compat_final_observation[key].shape != value.shape or compat_final_observation[key].dtype != value.dtype ): compat_final_observation = { obs_key: np.zeros_like(obs_value) for obs_key, obs_value in obs.items() } self._state.info["final_observation"] = compat_final_observation break compat_terminal_mask = self._state.info.get("_final_observation") if not isinstance(compat_terminal_mask, np.ndarray) or compat_terminal_mask.shape != ( self._num_envs, ): compat_terminal_mask = np.zeros((self._num_envs,), dtype=bool) self._state.info["_final_observation"] = compat_terminal_mask return compat_final_observation, compat_terminal_mask def _clear_step_final_observation(self) -> None: assert self._state is not None self._state.final_observation = None compat_terminal_mask = self._state.info.get("_final_observation") if isinstance(compat_terminal_mask, np.ndarray): compat_terminal_mask.fill(False) def _init_domain_randomization(self, provider: "DomainRandomizationProvider") -> None: from unilab.dr import DomainRandomizationManager self._dr_manager = DomainRandomizationManager(self, provider) if not self._init_randomization_applied: self._init_randomization_applied = self._dr_manager.apply_init_randomization() self._backend.materialize()
[docs] def reset(self, env_indices: np.ndarray) -> Tuple[dict[str, np.ndarray], dict]: if self._dr_manager is None: # pragma: no cover - constructor integration error raise RuntimeError("Domain-randomization manager has not been initialized") return self._dr_manager.reset(env_indices)
def _compute_truncated(self, state: NpEnvState) -> np.ndarray: """Compute truncation conditions. By default, episodes are truncated only when the configured maximum episode length is reached. Subclasses may override this to add task-specific truncation conditions while remaining compatible with the existing done/reset contract. """ truncated = cast(np.ndarray | None, getattr(self, "_truncated_scratch", None)) if truncated is None or truncated.shape != (self._num_envs,): truncated = np.zeros((self._num_envs,), dtype=bool) self._truncated_scratch = truncated truncated.fill(False) if self._cfg.max_episode_steps: np.greater_equal(state.info["steps"], self._cfg.max_episode_steps, out=truncated) return truncated
[docs] def init_play_renderer( self, render_spacing: float | None = None, render_offset_mode: str | None = None, *, headless: bool = False, capture: bool = False, width: int = 1280, height: int = 720, camera_kwargs: CameraCfg | Mapping[str, Any] | None = None, ) -> None: """Initialize backend-native playback rendering when available.""" if capture: if not self.play_capabilities.supports_native_video_capture: raise NotImplementedError( f"{self._backend.__class__.__name__} does not support native video capture" ) elif not self.play_capabilities.supports_native_interactive_renderer: raise NotImplementedError( f"{self._backend.__class__.__name__} does not support native interactive playback" ) spacing = ( float(render_spacing) if render_spacing is not None else float(self._cfg.render_spacing) ) offset_mode = ( str(render_offset_mode) if render_offset_mode is not None else str(getattr(self._cfg, "render_offset_mode", "grid")) ) self._backend.init_renderer( spacing=spacing, offset_mode=offset_mode, headless=bool(headless), capture=bool(capture), width=int(width), height=int(height), camera_kwargs=camera_kwargs, )
[docs] def resolve_play_render_plan( self, *, play_render_mode: str | None, play_steps: int | None, output_video: str | PathLike[str] | None, ) -> BackendPlayRenderPlan: """Resolve high-level playback mode through the concrete backend.""" # Drake has no renderer of its own. Its playback contract captures # Drake state and feeds it to the shared MuJoCo renderer, so ordinary # ``auto`` playback should produce a video without requiring a # backend-specific override in every task owner YAML. Explicit # ``none`` remains available for headless runs. if self._backend.backend_type == "drake" and ( play_render_mode is None or str(play_render_mode).strip().lower() == "auto" ): play_render_mode = "record" return self._backend.resolve_play_render_plan( play_render_mode=play_render_mode, play_steps=play_steps, output_video=output_video, )
[docs] def run_playback( self, *, initialize: Callable[[], Any], step: Callable[[Any], Any], num_steps: int | None, output_video: str | PathLike[str] | None = None, render_spacing: float | None = None, render_offset_mode: str | None = None, headless: bool | None = None, record_video: bool | None = None, frame_state_getter: Callable[[], np.ndarray] | None = None, camera_kwargs: CameraCfg | Mapping[str, Any] | None = None, debug_overlay_getter: DebugOverlayGetter | None = None, on_frame: Callable[[int, np.ndarray], np.ndarray | None] | None = None, ) -> str | None: """Execute playback through the concrete backend. ``on_frame`` is declared on the env contract but the unisim ``SimBackend.run_playback`` boundary does not accept it yet; passing a callback fails closed until the upstream contract lands. Use :class:`unilab.visualization.playback_session.SnapshotPlaybackSession` for deferred rendering with per-frame callbacks today. """ if on_frame is not None: raise NotImplementedError( f"{self.__class__.__name__} cannot forward on_frame to " f"{self._backend.__class__.__name__}.run_playback yet; the unisim playback " "contract does not declare the parameter (tracked by unilabsim/wuji_unilab#21)" ) return cast( str | None, self._backend.run_playback( env=self, initialize=initialize, step=step, num_steps=num_steps, output_video=output_video, render_spacing=render_spacing, render_offset_mode=render_offset_mode, headless=headless, record_video=record_video, frame_state_getter=frame_state_getter, camera_kwargs=camera_kwargs, debug_overlay_getter=debug_overlay_getter, ), )
[docs] def render_play_frame(self) -> None: """Render one interactive playback frame through the env contract.""" if not self.play_capabilities.supports_native_interactive_renderer: raise NotImplementedError( f"{self._backend.__class__.__name__} does not support native interactive playback" ) self._backend.render()
[docs] def render(self, mode: str = "rgb_array") -> np.ndarray: """Render the current state to an RGB array through the play renderer. Lazily initializes a headless capture renderer on first use and returns one detached ``(H, W, 3)`` uint8 frame per call. Backends without native video capture fail closed with a class-named error. """ if mode != "rgb_array": raise NotImplementedError( f"{self.__class__.__name__} does not support render mode {mode!r}" ) if not self.play_capabilities.supports_native_video_capture: raise NotImplementedError( f"{self._backend.__class__.__name__} does not support native video capture, " f"so {self.__class__.__name__}.render(mode='rgb_array') is unavailable" ) if not self._rgb_array_renderer_ready: self.init_play_renderer(headless=True, capture=True) self._rgb_array_renderer_ready = True return self.capture_play_video_frame()
[docs] def capture_play_video_frame(self) -> np.ndarray: """Capture one detached RGB video frame through the env contract.""" if not self.play_capabilities.supports_native_video_capture: raise NotImplementedError( f"{self._backend.__class__.__name__} does not support native video capture" ) return cast( np.ndarray, np.asarray(self._backend.capture_video_frame(), dtype=np.uint8).copy() )
[docs] def get_physics_state_snapshot(self) -> np.ndarray: """Return a detached physics snapshot for offline playback/video export.""" if not self.play_capabilities.supports_physics_state_playback: raise NotImplementedError( f"{self._backend.__class__.__name__} does not support physics-state playback" ) physics_state = cast( np.ndarray, np.asarray(self._backend.get_physics_state(), dtype=np.float32) ) snapshot = cast(np.ndarray, physics_state.copy()) return snapshot
[docs] @abc.abstractmethod def apply_action(self, actions: np.ndarray, state: NpEnvState) -> np.ndarray: """Subclasses implement the action-to-control conversion."""
[docs] @abc.abstractmethod def update_state(self, state: NpEnvState) -> NpEnvState: """Subclasses compute observation, reward, and termination state."""
@property def play_capabilities(self) -> EnvPlayCapabilities: capabilities = self._backend.get_play_capabilities() return EnvPlayCapabilities( supports_native_interactive_renderer=capabilities.supports_native_interactive_renderer, supports_physics_state_playback=capabilities.supports_physics_state_playback, supports_native_video_capture=capabilities.supports_native_video_capture, supports_debug_overlay=capabilities.supports_debug_overlay, supports_interactive_debug_overlay=capabilities.supports_interactive_debug_overlay, )
[docs] def get_playback_model(self, env_index: int | None = None) -> Any: """Return the backend playback model for one env in a vectorized batch. Args: env_index: Optional vectorized environment index. Returns: The backend-specific playback model. """ return self._backend.get_playback_model(env_index)
[docs] def get_scene_visual_model_file(self) -> str | None: """Return the backend scene visual model file on the cold path, when available.""" return cast(str | None, self._backend.get_scene_visual_model_file())
[docs] def set_nan_guard(self, guard: "NanGuard") -> None: self._nan_guard = guard
[docs] def set_autoreset(self, enabled: bool) -> None: """Toggle automatic reset of done envs at the end of ``step``. Defaults to ``True`` (standard RL autoreset). Interactive playback can disable it so a terminated robot stays put until a manual reset. """ self._autoreset = bool(enabled)
[docs] def export_training_state(self) -> dict[str, Any]: """Export cumulative training progress, independently of episode/physics state. Task-specific curriculum state belongs to the task's explicit provider; this payload deliberately does not inspect manager or environment internals. """ return {"version": 1, "step_counter": self.step_counter}
[docs] def import_training_state(self, state: Mapping[str, Any]) -> None: """Validate and restore cumulative progress before the next control step.""" if not isinstance(state, Mapping) or set(state) != {"version", "step_counter"}: raise ValueError("NpEnv training state requires version and step_counter only") if type(state["version"]) is not int or state["version"] != 1: raise ValueError("Unsupported NpEnv training state version") counter = state["step_counter"] if type(counter) is not int or counter < 0: raise ValueError("NpEnv training step_counter must be a non-negative integer") self.step_counter = counter
[docs] def close(self) -> None: """Close the environment and release backend-owned scene assets.""" self._backend.cleanup_scene_assets()