diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 6c488f799..b3a1dd893 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -10,6 +10,7 @@ repos: rev: v1.1.408 hooks: - id: pyright + pass_filenames: false - repo: https://github.com/pre-commit/mirrors-mypy rev: v1.15.0 diff --git a/AGENTS.md b/AGENTS.md index 6a4265670..35c76fa4e 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -32,7 +32,11 @@ UniLab 是一个 **高性能、模块化、contract 驱动** 的 RL infrastructu - SAC / TD3: `scripts/train_offpolicy.py` - env contract: `src/unilab/base/np_env.py` - backend contract: `src/unilab/base/backend/base.py` -- config schema: `src/unilab/config/structured_configs.py` +- training run helpers: `src/unilab/training/run.py` +- visualization helpers: `src/unilab/visualization/` +- env shared numeric helpers: `src/unilab/envs/common/rotation.py`, `src/unilab/envs/common/math.py` +- MLX rotation helpers: `src/unilab/algos/mlx/common/rotation.py` +- config schema: `src/unilab/structured_configs.py` - async runner: `src/unilab/ipc/async_runner.py` ## GitHub CLI (gh) 速查 diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 3b622ed93..694c5b2e7 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -17,6 +17,8 @@ Languages: English | [简体中文](docs/developers/zh_CN/CONTRIBUTING.md) - Run `make check` before code-related commits - Keep backup files, temporary exports, and legacy compatibility copies out of the source tree; do not commit artifacts such as `*.bak`, `*.tmp`, `*.old`, `*.orig`, or editor backup files ending in `~` - For user-facing workflow changes, keep `README.md`, `CONTRIBUTING.md`, and the matching localized docs under `docs/` in sync +- Do not add new owner logic under `src/unilab/utils/`; the current `src/unilab/utils/*.py` files are transition shims only and are scheduled for removal in `0.2.0` +- Name new owner modules and packages after their responsibility: prefer singular nouns, use plural only for collection-valued contracts, and reserve suffixes such as `_factory` for factory modules ## Read Before You Start diff --git a/benchmark/benchmark_fast_sac_backends.py b/benchmark/benchmark_fast_sac_backends.py index 8efef40ef..077ad01e8 100755 --- a/benchmark/benchmark_fast_sac_backends.py +++ b/benchmark/benchmark_fast_sac_backends.py @@ -15,7 +15,10 @@ try: from benchmark.core.task_names import locomotion_env_name, normalize_locomotion_task_id except ModuleNotFoundError: - from core.task_names import locomotion_env_name, normalize_locomotion_task_id + from core.task_names import ( # type: ignore[no-redef] + locomotion_env_name, + normalize_locomotion_task_id, + ) @dataclass @@ -32,7 +35,7 @@ def run_backend(task, max_iterations, backend, num_envs): import datetime from unilab.algos.torch.fast_sac.runner import FastSACRunner - from unilab.config.structured_configs import SACConfig + from unilab.structured_configs import SACConfig cfg = SACConfig() diff --git a/benchmark/benchmark_mujoco_backend_step_detail.py b/benchmark/benchmark_mujoco_backend_step_detail.py index 676274487..4771ee6ee 100644 --- a/benchmark/benchmark_mujoco_backend_step_detail.py +++ b/benchmark/benchmark_mujoco_backend_step_detail.py @@ -45,8 +45,8 @@ from matplotlib.patches import Rectangle from mujoco.batch_env import BatchEnvPool -from unilab.base.dtype_config import get_global_dtype -from unilab.utils.xml_utils import create_discardvisual_xml +from unilab.base.backend.xml import create_discardvisual_xml +from unilab.dtype_config import get_global_dtype matplotlib.use("Agg") import matplotlib.pyplot as plt diff --git a/docs/developers/adr/ADR-0004-registry-bootstrap-contract.md b/docs/developers/adr/ADR-0004-registry-bootstrap-contract.md index 7a9032cf8..cff9e8b87 100644 --- a/docs/developers/adr/ADR-0004-registry-bootstrap-contract.md +++ b/docs/developers/adr/ADR-0004-registry-bootstrap-contract.md @@ -36,6 +36,6 @@ UniLab 的 env 注册依赖 `@registry.envcfg(...)` 与 `@registry.env(...)` dec ## Evidence In Repo - Registry 入口: `src/unilab/base/registry.py` -- Bootstrap helper: `src/unilab/utils/algo_utils.py` +- Bootstrap helper: `src/unilab/base/registry.py` - Env package 入口: `src/unilab/envs/locomotion/__init__.py`, `src/unilab/envs/motion_tracking/__init__.py`, `src/unilab/envs/manipulation/__init__.py` - Bootstrap tests: `tests/utils/test_algo_utils.py`, `tests/base/test_registry.py` diff --git a/docs/developers/zh_CN/CONTRIBUTING.md b/docs/developers/zh_CN/CONTRIBUTING.md index cccef23e4..291cf2364 100644 --- a/docs/developers/zh_CN/CONTRIBUTING.md +++ b/docs/developers/zh_CN/CONTRIBUTING.md @@ -17,6 +17,8 @@ - 代码相关提交前必须运行 `make check` - 备份文件、临时导出物和历史兼容副本不要进入源码树;不要提交 `*.bak`、`*.tmp`、`*.old`、`*.orig` 或以 `~` 结尾的编辑器备份文件 - 只要改动用户可见工作流,就要同步维护顶层 `README.md`、`CONTRIBUTING.md`,以及 `docs/users/zh_CN/` 和 `docs/developers/zh_CN/` 下对应语言文档 +- 不要再往 `src/unilab/utils/` 塞新的 owner 逻辑;当前 `src/unilab/utils/*.py` 仅是过渡期 shim,计划在 `0.2.0` 删除 +- 新模块/包名应直接表达 owner 职责:默认使用单数名词;只有在语义本身就是集合契约时才使用复数;工厂模块使用 `_factory` 后缀 ## Read Before You Start diff --git a/docs/developers/zh_CN/development-standard.md b/docs/developers/zh_CN/development-standard.md index a9eef6d84..35147b23b 100644 --- a/docs/developers/zh_CN/development-standard.md +++ b/docs/developers/zh_CN/development-standard.md @@ -34,7 +34,7 @@ CPU Physics Sim ──shm──► Collector / IPC ──shm──► GPU Learne |-------|------|------|----------| | L0 Backend | `base/backend/` | `SimBackend` 物理后端抽象 | 训练逻辑、reward | | L1 Env | `envs/`, `base/np_env.py` | MDP 语义、observation、reward、reset | 调度、日志策略 | -| L2 Config & Registry | `config/`, `base/registry.py`, `conf/` | schema、task / reward 组合、注册 | 零散业务默认值 | +| L2 Config & Registry | `structured_configs.py`, `training/reward.py`, `base/registry.py`, `conf/` | schema、task / reward 组合、注册 | 零散业务默认值 | | L3 Algo & IPC | `algos/`, `ipc/` | learner、runner、collector、shared-memory 通路 | env / backend 细节 | | L4 Scripts | `scripts/` | 只做装配 | 核心业务规则 | @@ -68,7 +68,7 @@ CPU Physics Sim ──shm──► Collector / IPC ──shm──► GPU Learne ## 5. Configuration -UniLab 使用 dataclass + Hydra。schema 位于 `src/unilab/config/structured_configs.py`,运行时配置位于 `conf/{ppo,appo,offpolicy}/`。 +UniLab 使用 dataclass + Hydra。schema 位于 `src/unilab/structured_configs.py`,运行时配置位于 `conf/{ppo,appo,offpolicy}/`。 合成顺序: `{algo}/config*.yaml` -> `task=...` -> CLI override。 @@ -141,7 +141,7 @@ Env **负责** MDP 语义、observation 结构、reward、reset,以及 backend - `scripts/train_{rsl_rl,mlx_ppo,appo,offpolicy}.py` - `src/unilab/base/{registry,np_env}.py` - `src/unilab/base/backend/base.py` -- `src/unilab/config/structured_configs.py` +- `src/unilab/structured_configs.py` - `src/unilab/utils/{reward_utils,obs_utils}.py` - `src/unilab/ipc/async_runner.py` diff --git a/docs/users/zh_CN/02-simulation-backends.md b/docs/users/zh_CN/02-simulation-backends.md index 01d83e2ae..1b05c0cc0 100644 --- a/docs/users/zh_CN/02-simulation-backends.md +++ b/docs/users/zh_CN/02-simulation-backends.md @@ -76,7 +76,7 @@ uv run scripts/generate_support_matrix.py --write ### Source Index -- Registry bootstrap: `src/unilab/envs/**` decorators via `unilab.utils.algo_utils.ensure_registries()`. +- Registry bootstrap: `src/unilab/envs/**` decorators via `unilab.base.registry.ensure_registries()`. - Owner YAML scan: `conf/ppo/task/**`, `conf/appo/task/**`, `conf/offpolicy/task/**`. - Generic compose coverage: `tests/config/test_config_system.py::test_supported_task_composes`. - MLX-specific compose coverage only upgrades task owners listed in `tests/config/test_config_system.py::_PPO_MLX_TASKS`: `go1_joystick_flat`, `go2_joystick_flat`, `g1_walk_flat`. diff --git a/pyproject.toml b/pyproject.toml index a9c29dc70..ad1ffd246 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -96,6 +96,7 @@ python_version = "3.10" warn_return_any = true warn_unused_configs = true ignore_missing_imports = true +no_site_packages = true exclude = [ "src/unilab/algos/torch/rsl_rl/", # vendored third-party, tracked in .gitignore ] @@ -125,9 +126,8 @@ exclude = [ "src/unilab/algos/torch/common/ane_*", # Apple Neural Engine, macOS-only "src/unilab/base/backend/", # mujoco-uni stubs mismatch; optional backends "src/unilab/envs/", # mujoco-uni internal API, stubs mismatch - "src/unilab/utils/render_many.py", # direct mujoco C bindings - "src/unilab/utils/viser_scene.py", # optional viser dependency - "src/unilab/utils/hardware_monitor.py", # optional pynvml/psutil deps + "src/unilab/training/monitoring.py", # optional pynvml/psutil deps + "src/unilab/visualization/", # direct mujoco C bindings + optional viser deps ] reportMissingImports = "warning" reportMissingModuleSource = "none" diff --git a/scripts/motion/bones_seed_csv_to_npz.py b/scripts/motion/bones_seed_csv_to_npz.py index 8e3aea094..c3856f477 100644 --- a/scripts/motion/bones_seed_csv_to_npz.py +++ b/scripts/motion/bones_seed_csv_to_npz.py @@ -41,8 +41,8 @@ from tqdm import tqdm from unilab.assets import ASSETS_ROOT_PATH -from unilab.utils.math_utils import np_quat_angular_velocity, np_quat_ensure_continuity -from unilab.utils.xml_utils import inject_mujoco_tracking_sensors +from unilab.base.backend.xml import inject_mujoco_tracking_sensors +from unilab.envs.common.rotation import np_quat_angular_velocity, np_quat_ensure_continuity ROOT_COLUMNS = [ "Frame", diff --git a/scripts/motion/csv_to_npz.py b/scripts/motion/csv_to_npz.py index 12aa85773..0899b3531 100644 --- a/scripts/motion/csv_to_npz.py +++ b/scripts/motion/csv_to_npz.py @@ -29,8 +29,8 @@ from tqdm import tqdm from unilab.assets import ASSETS_ROOT_PATH -from unilab.utils.math_utils import np_quat_angular_velocity, np_quat_ensure_continuity -from unilab.utils.xml_utils import inject_mujoco_tracking_sensors +from unilab.base.backend.xml import inject_mujoco_tracking_sensors +from unilab.envs.common.rotation import np_quat_angular_velocity, np_quat_ensure_continuity def quat_slerp(q1: np.ndarray, q2: np.ndarray, t: float) -> np.ndarray: diff --git a/scripts/play_interactive.py b/scripts/play_interactive.py index 238033208..5e5f1b72b 100644 --- a/scripts/play_interactive.py +++ b/scripts/play_interactive.py @@ -50,18 +50,16 @@ get_entrypoint_log_root, resolve_task_checkpoint_path, ) +from unilab.training.rsl_rl import ( + RslRlVecEnvWrapper, + get_policy_obs_dims, + normalize_ppo_train_cfg, +) ensure_registries() from unilab.base import registry -from unilab.config.structured_configs import PPOConfig as _StructuredPPOConfig -from unilab.utils.rsl_rl_compat import ( - convert_config_v3_to_v4, - convert_config_v5, - is_rsl_rl_v4, - is_rsl_rl_v5, -) -from unilab.utils.rsl_rl_vec_env_wrapper import RslRlVecEnvWrapper +from unilab.structured_configs import PPOConfig as _StructuredPPOConfig PPOConfig = _StructuredPPOConfig @@ -122,8 +120,8 @@ def _infer_checkpoint_actor_input_dim(ckpt_path: str) -> int | None: def _backend_adapter(cfg: DictConfig): + from unilab.base.backend.xml import materialize_scene_visual_override from unilab.training import BackendAdapter - from unilab.utils.xml_utils import materialize_scene_visual_override return BackendAdapter( cfg, @@ -578,8 +576,7 @@ def play_interactive(args, cfg: DictConfig | None = None): "Set DISPLAY correctly, or run this command in a desktop session." ) return - actor_obs_dim = int(env.obs_groups_spec.get("obs", sum(env.obs_groups_spec.values()))) - flat_obs_dim = int(sum(env.obs_groups_spec.values())) + actor_obs_dim, flat_obs_dim = get_policy_obs_dims(env.obs_groups_spec) policy_obs_mode = args.policy_obs_mode algo_log_name = getattr(args, "algo_log_name", "rsl_rl_ppo") @@ -614,14 +611,10 @@ def play_interactive(args, cfg: DictConfig | None = None): f"{policy_obs_mode} (actor_obs={actor_obs_dim}, flat_obs={flat_obs_dim})" ) - train_cfg = _algo_config_dict(cfg) + train_cfg = normalize_ppo_train_cfg(_algo_config_dict(cfg)) if "runner" not in train_cfg: train_cfg["runner"] = {} train_cfg["runner"]["logger"] = "none" - if is_rsl_rl_v5(): - train_cfg = cast(dict[str, Any], convert_config_v5(train_cfg)) - elif is_rsl_rl_v4(): - train_cfg = convert_config_v3_to_v4(train_cfg) policy = None if args.action_mode == "policy": diff --git a/scripts/play_viser.py b/scripts/play_viser.py index 081899a0d..ccd44a82c 100644 --- a/scripts/play_viser.py +++ b/scripts/play_viser.py @@ -51,15 +51,17 @@ ensure_registries, get_entrypoint_log_root, ) -from unilab.utils.render_many import get_grid_offsets -from unilab.utils.rsl_rl_compat import ( - convert_config_v3_to_v4, - convert_config_v5, - is_rsl_rl_v4, - is_rsl_rl_v5, +from unilab.training.rsl_rl import ( + RslRlVecEnvWrapper, + get_policy_obs_dims, + normalize_ppo_train_cfg, +) +from unilab.visualization.render_many import get_grid_offsets +from unilab.visualization.viser_scene import ( + VISER_AVAILABLE, + MujocoViserScene, + build_visible_env_indices, ) -from unilab.utils.rsl_rl_vec_env_wrapper import RslRlVecEnvWrapper -from unilab.utils.viser_scene import VISER_AVAILABLE, MujocoViserScene, build_visible_env_indices ensure_registries() @@ -240,8 +242,7 @@ def play_viser(args: PlayInteractiveArgs, cfg: DictConfig) -> None: ) # --- Load policy --------------------------------------------------------- - actor_obs_dim = int(env.obs_groups_spec.get("obs", sum(env.obs_groups_spec.values()))) - flat_obs_dim = int(sum(env.obs_groups_spec.values())) + actor_obs_dim, flat_obs_dim = get_policy_obs_dims(env.obs_groups_spec) policy_obs_mode = args.policy_obs_mode algo_log_name = getattr(args, "algo_log_name", "rsl_rl_ppo") @@ -270,14 +271,10 @@ def play_viser(args: PlayInteractiveArgs, cfg: DictConfig) -> None: wrapped_env = RslRlVecEnvWrapper(env, device=device, policy_obs_mode=policy_obs_mode) - train_cfg = _algo_config_dict(cfg) + train_cfg = normalize_ppo_train_cfg(_algo_config_dict(cfg)) if "runner" not in train_cfg: train_cfg["runner"] = {} train_cfg["runner"]["logger"] = "none" - if is_rsl_rl_v5(): - train_cfg = cast(dict[str, Any], convert_config_v5(train_cfg)) - elif is_rsl_rl_v4(): - train_cfg = convert_config_v3_to_v4(train_cfg) policy = None if args.action_mode == "policy": diff --git a/scripts/train_appo.py b/scripts/train_appo.py index 3aad4db56..5664f3fb6 100644 --- a/scripts/train_appo.py +++ b/scripts/train_appo.py @@ -20,9 +20,9 @@ create_env, ensure_registries, get_log_root, - render_play_mode, ) -from unilab.utils.experiment_tracking import ExperimentTracker +from unilab.training.experiment import ExperimentTracker +from unilab.visualization import render_play_mode def build_appo_runner_kwargs( @@ -113,8 +113,6 @@ def play_appo(cfg: DictConfig, rl_cfg: dict[str, Any]) -> str | None: from rsl_rl.utils import resolve_callable from tensordict import TensorDict - from unilab.utils.rsl_rl_compat import convert_config_v3_to_v4, is_rsl_rl_v4, is_rsl_rl_v5 - env_cfg_override = BackendAdapter( cfg, root_dir=ROOT_DIR, algo_name="appo" ).build_task_env_cfg_override() @@ -136,7 +134,7 @@ def play_appo(cfg: DictConfig, rl_cfg: dict[str, Any]) -> str | None: env_cfg_override=env_cfg_override, ), ) - from unilab.utils.obs_utils import get_obs_dims + from unilab.base.observations import get_obs_dims obs_dim, critic_dim = get_obs_dims(env.obs_groups_spec) action_shape = env.action_space.shape @@ -164,11 +162,6 @@ def play_appo(cfg: DictConfig, rl_cfg: dict[str, Any]) -> str | None: elif isinstance(critic_group, dict) and "policy" in critic_group: critic_group["policy"] = critic_dim if critic_dim > 0 else obs_dim - if is_rsl_rl_v5(): - pass - elif is_rsl_rl_v4(): - rl_cfg_dict = convert_config_v3_to_v4(rl_cfg_dict) - from copy import deepcopy obs_example = torch.zeros((cfg.training.play_env_num, obs_dim), device=device) diff --git a/scripts/train_mlx_ppo.py b/scripts/train_mlx_ppo.py index 568048026..29000dd66 100644 --- a/scripts/train_mlx_ppo.py +++ b/scripts/train_mlx_ppo.py @@ -5,6 +5,7 @@ from __future__ import annotations import datetime +import importlib import math import os import pickle @@ -13,26 +14,27 @@ import time from collections import deque from pathlib import Path -from typing import Any, cast +from types import SimpleNamespace +from typing import TYPE_CHECKING, Any, cast import hydra -import mlx.core as mx import numpy as np -from mlx.utils import tree_map from omegaconf import DictConfig, OmegaConf +if TYPE_CHECKING: + from unilab.algos.mlx.ppo import MLPActorCritic, PPOTrainer + ROOT_DIR = Path(__file__).parent.parent sys.path.append(str(ROOT_DIR)) -from unilab.algos.mlx.common import EmpiricalDiscountedVariationNormalization, RolloutBuffer -from unilab.algos.mlx.ppo import MLPActorCritic, PPOConfig, PPOTrainer +from unilab.base.observations import flatten_obs_dict +from unilab.logging import OnPolicyLogger from unilab.training import ( BackendAdapter, create_env, ensure_registries, get_log_root, parse_checkpoint_path, - render_play_mode, setup_logger, ) from unilab.training import ( @@ -41,12 +43,42 @@ from unilab.training import ( get_latest_run as get_latest_run_common, ) -from unilab.utils.experiment_tracking import ExperimentTracker -from unilab.utils.obs_utils import flatten_obs_dict -from unilab.utils.onpolicy_logger import OnPolicyLogger +from unilab.training.experiment import ExperimentTracker +from unilab.visualization import render_play_mode ensure_registries() +_MLX_RUNTIME: SimpleNamespace | None = None + + +def _require_mlx_runtime() -> SimpleNamespace: + global _MLX_RUNTIME + if _MLX_RUNTIME is None: + mx = importlib.import_module("mlx.core") + tree_map = importlib.import_module("mlx.utils").tree_map + mlx_common = importlib.import_module("unilab.algos.mlx.common") + mlx_ppo = importlib.import_module("unilab.algos.mlx.ppo") + _MLX_RUNTIME = SimpleNamespace( + mx=mx, + tree_map=tree_map, + EmpiricalDiscountedVariationNormalization=( + mlx_common.EmpiricalDiscountedVariationNormalization + ), + RolloutBuffer=mlx_common.RolloutBuffer, + MLPActorCritic=mlx_ppo.MLPActorCritic, + PPOConfig=mlx_ppo.PPOConfig, + PPOTrainer=mlx_ppo.PPOTrainer, + ) + return _MLX_RUNTIME + + +mx = SimpleNamespace( + array=np.array, + float16=np.float16, + float32=np.float32, + sum=np.sum, +) + class TensorboardScalarWriter: """Minimal scalar writer based on tensorboard event files.""" @@ -86,8 +118,9 @@ def get_latest_checkpoint(run_dir: Path) -> Path | None: return cast(Path | None, get_latest_checkpoint_common(run_dir, suffix=".safetensors")) -def save_trainer_state(path: Path, trainer: PPOTrainer, iteration: int) -> None: +def save_trainer_state(path: Path, trainer: Any, iteration: int) -> None: """Save optimizer state and trainer metadata for resume.""" + tree_map = _require_mlx_runtime().tree_map payload = { "iteration": int(iteration), "learning_rate": float(trainer.learning_rate), @@ -97,8 +130,11 @@ def save_trainer_state(path: Path, trainer: PPOTrainer, iteration: int) -> None: pickle.dump(payload, f) -def load_trainer_state(path: Path, trainer: PPOTrainer, dtype=mx.float32) -> int: +def load_trainer_state(path: Path, trainer: Any, dtype: Any = None) -> int: """Load optimizer state and trainer metadata.""" + mx = _require_mlx_runtime().mx + if dtype is None: + dtype = mx.float32 with path.open("rb") as f: payload = pickle.load(f) trainer.learning_rate = float(payload.get("learning_rate", trainer.learning_rate)) @@ -106,15 +142,18 @@ def load_trainer_state(path: Path, trainer: PPOTrainer, dtype=mx.float32) -> int return int(payload.get("iteration", -1)) -def build_model(cfg, obs_dim: int, action_dim: int, dtype=mx.float32) -> MLPActorCritic: +def build_model(cfg, obs_dim: int, action_dim: int, dtype: Any = None) -> Any: """Build actor-critic model from config (expects cfg with .policy and .empirical_normalization).""" + runtime = _require_mlx_runtime() + if dtype is None: + dtype = runtime.mx.float32 policy_cfg = cfg.policy init_noise_std = float(getattr(policy_cfg, "init_noise_std", 1.0)) init_log_std = float(math.log(max(init_noise_std, 1e-6))) obs_norm = bool(getattr(cfg, "empirical_normalization", False)) noise_std_type = str(getattr(policy_cfg, "noise_std_type", "scalar")) state_dependent_std = bool(getattr(policy_cfg, "state_dependent_std", False)) - return MLPActorCritic( + return runtime.MLPActorCritic( obs_dim=obs_dim, action_dim=action_dim, actor_hidden_dims=policy_cfg.actor_hidden_dims, @@ -128,23 +167,25 @@ def build_model(cfg, obs_dim: int, action_dim: int, dtype=mx.float32) -> MLPActo ) -def get_time_limit_bootstrap_values( - state: Any, model: MLPActorCritic, model_dtype=mx.float32 -) -> mx.array | None: +def get_time_limit_bootstrap_values(state: Any, model: Any, model_dtype: Any = None) -> Any | None: """Return V(final_observation) for current timeout envs when available.""" + use_mlx_runtime = type(model).__module__.startswith("unilab.algos.mlx") + mx_mod = _require_mlx_runtime().mx if use_mlx_runtime else mx + if model_dtype is None: + model_dtype = mx_mod.float32 if not hasattr(state, "truncated"): return None timeout_mask = np.asarray(state.truncated, dtype=bool) if not np.any(timeout_mask): return None final_observation = getattr(state, "final_observation", None) + info = getattr(state, "info", None) if final_observation is None: - info = getattr(state, "info", None) if isinstance(info, dict): final_observation = info.get("final_observation") if not isinstance(final_observation, dict): return None - final_obs = mx.array(flatten_obs_dict(final_observation)) + final_obs = mx_mod.array(flatten_obs_dict(final_observation)) if getattr(final_obs, "dtype", None) != model_dtype: final_obs = final_obs.astype(model_dtype) return model.value(final_obs) @@ -156,6 +197,7 @@ def _get_log_root(cfg: DictConfig) -> Path: def play_mlx_ppo(cfg: DictConfig, dtype, use_fp16: bool, resolved_sim_backend: str) -> str | None: """Play mode for MLX PPO.""" + mx = _require_mlx_runtime().mx env_cfg_override = BackendAdapter( cfg, root_dir=ROOT_DIR, algo_name="ppo" ).build_task_env_cfg_override() @@ -316,6 +358,12 @@ def _play_step(current_obs): @hydra.main(version_base="1.3", config_path="../conf/ppo", config_name="config_mlx") def main(cfg: DictConfig) -> None: + runtime = _require_mlx_runtime() + mx = runtime.mx + PPOConfig = runtime.PPOConfig + PPOTrainer = runtime.PPOTrainer + RolloutBuffer = runtime.RolloutBuffer + EmpiricalDiscountedVariationNormalization = runtime.EmpiricalDiscountedVariationNormalization task_name = cfg.training.task_name resolved_sim_backend = cfg.training.sim_backend diff --git a/scripts/train_offpolicy.py b/scripts/train_offpolicy.py index bb2e92f13..c546806c5 100644 --- a/scripts/train_offpolicy.py +++ b/scripts/train_offpolicy.py @@ -20,12 +20,12 @@ create_env, ensure_registries, get_log_root, - render_play_mode, ) from unilab.training import ( resolve_checkpoint_path as resolve_checkpoint_path_common, ) -from unilab.utils.experiment_tracking import ExperimentTracker +from unilab.training.experiment import ExperimentTracker +from unilab.visualization import render_play_mode def default_device(torch_module, preferred: str | None = None) -> str: @@ -63,14 +63,14 @@ def extract_reset_obs(reset_result): def resolve_play_obs_dim(obs_groups_spec: dict[str, int]) -> int: - from unilab.utils.obs_utils import get_obs_dims + from unilab.base.observations import get_obs_dims obs_dim, _ = get_obs_dims(obs_groups_spec) return int(obs_dim) def extract_play_obs(obs_dict): - from unilab.utils.obs_utils import split_obs_dict + from unilab.base.observations import split_obs_dict obs_out, _ = split_obs_dict(obs_dict) return obs_out @@ -92,9 +92,10 @@ def build_runner(algo_name: str, cfg: DictConfig): raise ValueError("FlashSAC does not support training.num_gpus > 1") if algo_name == "sac": + from unilab.algos.torch.common.device import get_env_dims from unilab.algos.torch.fast_sac.learner import FastSACLearner from unilab.algos.torch.fast_sac.runner import FastSACRunner - from unilab.utils.device_utils import get_default_device, get_env_dims + from unilab.utils.device import get_default_device # Multi-GPU path if cfg.training.num_gpus > 1: @@ -108,7 +109,7 @@ def build_runner(algo_name: str, cfg: DictConfig): env_cfg_override=env_cfg_override, ) assert env.action_space.shape - from unilab.utils.obs_utils import get_obs_dims + from unilab.base.observations import get_obs_dims obs_dim, critic_dim = get_obs_dims(env.obs_groups_spec) action_dim = env.action_space.shape[0] @@ -281,7 +282,7 @@ def play_offpolicy(algo_name: str, cfg: DictConfig) -> str | None: import numpy as np import torch - from unilab.utils.algo_utils import build_actor + from unilab.algos.torch.common.actor_factory import build_actor env_cfg_override = build_offpolicy_env_cfg_override(algo_name, cfg) diff --git a/scripts/train_rsl_rl.py b/scripts/train_rsl_rl.py index 7dc456402..13d913eb1 100644 --- a/scripts/train_rsl_rl.py +++ b/scripts/train_rsl_rl.py @@ -16,6 +16,7 @@ if str(ROOT_DIR) not in sys.path: sys.path.insert(0, str(ROOT_DIR)) +from unilab.base.backend.xml import materialize_scene_visual_override from unilab.training import ( BackendAdapter, create_env, @@ -24,11 +25,10 @@ get_latest_run, get_log_root, parse_checkpoint_path, - render_play_mode, ) -from unilab.utils.experiment_tracking import ExperimentTracker, patch_rsl_rl_wandb_writer -from unilab.utils.rsl_rl_vec_env_wrapper import RslRlVecEnvWrapper -from unilab.utils.xml_utils import materialize_scene_visual_override +from unilab.training.experiment import ExperimentTracker, patch_rsl_rl_wandb_writer +from unilab.training.rsl_rl import RslRlVecEnvWrapper, normalize_ppo_train_cfg +from unilab.visualization import render_play_mode try: from rsl_rl.runners import OnPolicyRunner @@ -36,8 +36,6 @@ print("Could not import rsl_rl. Please ensure it is installed.") sys.exit(1) -from unilab.utils.rsl_rl_compat import convert_config_v5, is_rsl_rl_v5 - def _backend_adapter(cfg: DictConfig) -> BackendAdapter: return BackendAdapter( @@ -156,12 +154,10 @@ def play_rsl_rl(cfg: DictConfig, device: str) -> str | None: env_cfg_override=env_cfg_override, ) wrapped_env = RslRlVecEnvWrapper(env, device=device) - train_cfg = _algo_config_dict(cfg) + train_cfg = normalize_ppo_train_cfg(_algo_config_dict(cfg)) if "runner" not in train_cfg: train_cfg["runner"] = {} train_cfg["runner"]["logger"] = "none" - if is_rsl_rl_v5(): - train_cfg = cast(dict[str, Any], convert_config_v5(train_cfg)) runner = cast( Any, @@ -279,7 +275,7 @@ def main(cfg: DictConfig) -> None: ) wrapped_env = RslRlVecEnvWrapper(env, device=device) - train_cfg = _algo_config_dict(cfg) + train_cfg = normalize_ppo_train_cfg(_algo_config_dict(cfg)) if "runner" not in train_cfg: train_cfg["runner"] = {} @@ -300,9 +296,6 @@ def main(cfg: DictConfig) -> None: train_cfg["wandb_notes"] = wandb_settings["notes"] train_cfg["wandb_mode"] = wandb_settings["mode"] - if is_rsl_rl_v5(): - train_cfg = cast(dict[str, Any], convert_config_v5(train_cfg)) - runner = cast( Any, OnPolicyRunner(cast(Any, wrapped_env), train_cfg, log_dir=log_dir, device=device), diff --git a/src/unilab/algos/mlx/common/__init__.py b/src/unilab/algos/mlx/common/__init__.py index 1860b1f56..bd5a5c29b 100644 --- a/src/unilab/algos/mlx/common/__init__.py +++ b/src/unilab/algos/mlx/common/__init__.py @@ -4,16 +4,35 @@ algorithm implementations (e.g. PPO). """ -from .distributions import diag_gaussian_entropy, diag_gaussian_log_prob -from .mlp import MLP -from .normalization import EmpiricalDiscountedVariationNormalization, EmpiricalNormalization -from .rollout_storage import RolloutBuffer - -__all__ = [ - "MLP", - "EmpiricalNormalization", - "EmpiricalDiscountedVariationNormalization", - "RolloutBuffer", - "diag_gaussian_log_prob", - "diag_gaussian_entropy", -] +from __future__ import annotations + +from importlib import import_module + +_EXPORTS = { + "MLP": (".mlp", "MLP"), + "EmpiricalNormalization": (".normalization", "EmpiricalNormalization"), + "EmpiricalDiscountedVariationNormalization": ( + ".normalization", + "EmpiricalDiscountedVariationNormalization", + ), + "RolloutBuffer": (".rollout_storage", "RolloutBuffer"), + "diag_gaussian_log_prob": (".distributions", "diag_gaussian_log_prob"), + "diag_gaussian_entropy": (".distributions", "diag_gaussian_entropy"), +} + +__all__ = list(_EXPORTS) + + +def __getattr__(name: str): + try: + module_name, attr_name = _EXPORTS[name] + except KeyError as exc: + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") from exc + + value = getattr(import_module(module_name, __name__), attr_name) + globals()[name] = value + return value + + +def __dir__() -> list[str]: + return sorted(set(globals()) | set(__all__)) diff --git a/src/unilab/algos/mlx/common/rotation.py b/src/unilab/algos/mlx/common/rotation.py new file mode 100644 index 000000000..58796295f --- /dev/null +++ b/src/unilab/algos/mlx/common/rotation.py @@ -0,0 +1,38 @@ +from __future__ import annotations + +import importlib + + +def _require_mlx_core(): + """Import MLX lazily so non-MLX workflows don't crash at module import time.""" + try: + return importlib.import_module("mlx.core") + except Exception as exc: + raise RuntimeError( + "MLX backend is unavailable. Install the MLX extra to use MLX rotation helpers." + ) from exc + + +def quat_mul(q1, q2): + """Multiply two MLX quaternion batches.""" + mx = _require_mlx_core() + w1, x1, y1, z1 = q1[:, 0], q1[:, 1], q1[:, 2], q1[:, 3] + w2, x2, y2, z2 = q2[:, 0], q2[:, 1], q2[:, 2], q2[:, 3] + return mx.stack( + [ + w1 * w2 - x1 * x2 - y1 * y2 - z1 * z2, + w1 * x2 + x1 * w2 + y1 * z2 - z1 * y2, + w1 * y2 - x1 * z2 + y1 * w2 + z1 * x2, + w1 * z2 + x1 * y2 - y1 * x2 + z1 * w2, + ], + axis=1, + ) + + +def axis_angle_to_quat(axis, angle): + """Convert MLX axis-angle batches to quaternions.""" + mx = _require_mlx_core() + half_angle = angle / 2 + c = mx.cos(half_angle) + s = mx.sin(half_angle) + return mx.stack([c, axis[:, 0] * s, axis[:, 1] * s, axis[:, 2] * s], axis=1) diff --git a/src/unilab/algos/torch/appo/runner.py b/src/unilab/algos/torch/appo/runner.py index 7fd5d018f..af7b95770 100644 --- a/src/unilab/algos/torch/appo/runner.py +++ b/src/unilab/algos/torch/appo/runner.py @@ -20,8 +20,7 @@ from unilab.algos.torch.appo.learner import APPOLearner from unilab.algos.torch.appo.worker import appo_collector_fn from unilab.ipc import AsyncRunner, SharedOnPolicyStorage, SharedWeightSync -from unilab.utils.offpolicy_logger import OffPolicyLogger -from unilab.utils.rsl_rl_compat import convert_config_v3_to_v4, is_rsl_rl_v4, is_rsl_rl_v5 +from unilab.logging import OffPolicyLogger class APPORunner(AsyncRunner): @@ -87,8 +86,8 @@ def _resolve_dims(self): def _detect_dims(self): """Create a tiny env to read obs/action dims, then close it.""" from unilab.base import registry - from unilab.utils.algo_utils import ensure_registries - from unilab.utils.obs_utils import get_critic_base_dim, get_obs_dims + from unilab.base.observations import get_critic_base_dim, get_obs_dims + from unilab.base.registry import ensure_registries ensure_registries() @@ -109,10 +108,6 @@ def _detect_dims(self): def _build_learner(self): cfg = dict(self.rl_cfg) - if is_rsl_rl_v5(): - pass # appo_config is already v5-compatible (actor/critic format) - elif is_rsl_rl_v4(): - cfg = convert_config_v3_to_v4(cfg) import torch from tensordict import TensorDict diff --git a/src/unilab/algos/torch/appo/worker.py b/src/unilab/algos/torch/appo/worker.py index a9a645d05..ab5c78adb 100644 --- a/src/unilab/algos/torch/appo/worker.py +++ b/src/unilab/algos/torch/appo/worker.py @@ -15,9 +15,9 @@ import torch from rsl_rl.utils import resolve_callable -from unilab.utils.algo_utils import ensure_registries -from unilab.utils.final_observation import resolve_terminal_observation_contract -from unilab.utils.obs_utils import split_obs_dict +from unilab.base.final_observation import resolve_terminal_observation_contract +from unilab.base.observations import split_obs_dict +from unilab.base.registry import ensure_registries def compute_timeout_bootstrap_correction( @@ -78,7 +78,6 @@ def appo_collector_fn( from unilab.base import registry from unilab.ipc import SharedOnPolicyStorage, SharedWeightSync - from unilab.utils.rsl_rl_compat import convert_config_v3_to_v4, is_rsl_rl_v4, is_rsl_rl_v5 ensure_registries() @@ -107,10 +106,6 @@ def appo_collector_fn( # Build actor (stochastic MLPModel — mirrors runner._build_learner) cfg = dict(rl_cfg) - if is_rsl_rl_v5(): - pass # appo_config is already v5-compatible (actor/critic format) - elif is_rsl_rl_v4(): - cfg = convert_config_v3_to_v4(cfg) obs_example = torch.zeros((num_envs, obs_dim), device=collector_device) td_example = TensorDict({"policy": obs_example}, batch_size=num_envs) diff --git a/src/unilab/algos/torch/common/__init__.py b/src/unilab/algos/torch/common/__init__.py index b851d2614..75b190f0b 100644 --- a/src/unilab/algos/torch/common/__init__.py +++ b/src/unilab/algos/torch/common/__init__.py @@ -1,17 +1,18 @@ +from unilab.algos.torch.common.actor_factory import build_actor +from unilab.algos.torch.common.device import get_env_dims from unilab.algos.torch.common.networks import Critic, DistributionalQNetwork from unilab.algos.torch.common.normalization import EmpiricalNormalization from unilab.algos.torch.common.stability import check_nan_loss, clip_gradients, safe_tensor -from unilab.utils.algo_utils import build_actor, ensure_registries -from unilab.utils.offpolicy_logger import OffPolicyLogger +from unilab.base.registry import ensure_registries __all__ = [ "EmpiricalNormalization", "DistributionalQNetwork", "Critic", + "get_env_dims", "check_nan_loss", "clip_gradients", "safe_tensor", - "OffPolicyLogger", "ensure_registries", "build_actor", ] diff --git a/src/unilab/algos/torch/common/actor_factory.py b/src/unilab/algos/torch/common/actor_factory.py new file mode 100644 index 000000000..cd2b558b3 --- /dev/null +++ b/src/unilab/algos/torch/common/actor_factory.py @@ -0,0 +1,54 @@ +"""Actor factory helpers for torch off-policy algorithms.""" + +from __future__ import annotations + + +def build_actor( + algo_type, + obs_dim, + action_dim, + actor_hidden_dim, + use_layer_norm, + device, + num_envs=1, + actor_num_blocks: int = 2, + actor_noise_zeta_mu: float = 2.0, + actor_noise_zeta_max: int = 16, +): + """Build the correct actor model based on algorithm type.""" + if algo_type == "sac": + from unilab.algos.torch.fast_sac.learner import SACActor + + return SACActor( + obs_dim=obs_dim, + action_dim=action_dim, + hidden_dim=actor_hidden_dim, + use_layer_norm=use_layer_norm, + device=device, + ) + if algo_type == "td3": + from unilab.algos.torch.fast_td3.learner import TD3Actor + + return TD3Actor( + obs_dim=obs_dim, + n_act=action_dim, + num_envs=num_envs, + hidden_dim=actor_hidden_dim, + init_scale=0.01, + log_std_min=-0.9, + log_std_max=0.0, + device=device, + ) + if algo_type == "flashsac": + from unilab.algos.torch.flash_sac.network import FlashSACActor + + return FlashSACActor( + num_blocks=actor_num_blocks, + input_dim=obs_dim, + hidden_dim=actor_hidden_dim, + action_dim=action_dim, + noise_zeta_mu=actor_noise_zeta_mu, + noise_zeta_max=actor_noise_zeta_max, + device=device, + ) + raise ValueError(f"Unknown algo_type: {algo_type}") diff --git a/src/unilab/utils/device_utils.py b/src/unilab/algos/torch/common/device.py similarity index 67% rename from src/unilab/utils/device_utils.py rename to src/unilab/algos/torch/common/device.py index cee8d479a..1b1703ba5 100644 --- a/src/unilab/utils/device_utils.py +++ b/src/unilab/algos/torch/common/device.py @@ -1,22 +1,11 @@ -import torch - from unilab.base import registry -def get_default_device() -> str: - """Detect the best available device.""" - if torch.cuda.is_available(): - return "cuda" - if torch.backends.mps.is_available(): - return "mps" - return "cpu" - - def get_env_dims( env_name: str, sim_backend: str = "mujoco", env_cfg_override: dict | None = None ) -> tuple[int, int, int]: """Get (actor_obs_dim, action_dim, critic_obs_dim) from environment.""" - from unilab.utils.obs_utils import get_obs_dims as get_obs_dims_from_spec + from unilab.base.observations import get_obs_dims as get_obs_dims_from_spec env = registry.make( env_name, num_envs=1, sim_backend=sim_backend, env_cfg_override=env_cfg_override @@ -27,3 +16,6 @@ def get_env_dims( action_dim = action_shape[0] env.close() # type: ignore[attr-defined] return obs_dim, action_dim, critic_dim + + +__all__ = ["get_env_dims"] diff --git a/src/unilab/algos/torch/fast_sac/runner.py b/src/unilab/algos/torch/fast_sac/runner.py index 28aab2713..f07d2781d 100644 --- a/src/unilab/algos/torch/fast_sac/runner.py +++ b/src/unilab/algos/torch/fast_sac/runner.py @@ -2,9 +2,10 @@ from typing import Any +from unilab.algos.torch.common.device import get_env_dims from unilab.algos.torch.fast_sac.learner import FastSACLearner from unilab.algos.torch.offpolicy.runner import OffPolicyRunner -from unilab.utils.device_utils import get_default_device, get_env_dims +from unilab.utils.device import get_default_device class FastSACRunner(OffPolicyRunner): @@ -42,13 +43,13 @@ def __init__( world_size: int = 1, ): from unilab.base import registry - from unilab.utils.algo_utils import ensure_registries + from unilab.base.registry import ensure_registries ensure_registries() env: Any = registry.make( env_name, num_envs=1, sim_backend=sim_backend, env_cfg_override=env_cfg_override ) - from unilab.utils.obs_utils import get_obs_dims + from unilab.base.observations import get_obs_dims obs_dim, critic_obs_dim = get_obs_dims(env.obs_groups_spec) act_space_shape = env.action_space.shape diff --git a/src/unilab/algos/torch/fast_td3/runner.py b/src/unilab/algos/torch/fast_td3/runner.py index 8a6775467..f69ebdcd4 100644 --- a/src/unilab/algos/torch/fast_td3/runner.py +++ b/src/unilab/algos/torch/fast_td3/runner.py @@ -1,8 +1,9 @@ """FastTD3 runner built on top of the unified off-policy infra.""" +from unilab.algos.torch.common.device import get_env_dims from unilab.algos.torch.fast_td3.learner import FastTD3Learner from unilab.algos.torch.offpolicy.runner import OffPolicyRunner -from unilab.utils.device_utils import get_default_device, get_env_dims +from unilab.utils.device import get_default_device class FastTD3Runner(OffPolicyRunner): diff --git a/src/unilab/algos/torch/flash_sac/runner.py b/src/unilab/algos/torch/flash_sac/runner.py index a0fdcf2f4..ecff82d01 100644 --- a/src/unilab/algos/torch/flash_sac/runner.py +++ b/src/unilab/algos/torch/flash_sac/runner.py @@ -6,7 +6,7 @@ from unilab.algos.torch.flash_sac.learner import FlashSACLearner from unilab.algos.torch.offpolicy.runner import OffPolicyRunner -from unilab.utils.device_utils import get_default_device +from unilab.utils.device import get_default_device class FlashSACRunner(OffPolicyRunner): @@ -54,8 +54,8 @@ def __init__( use_compile: bool = False, ): from unilab.base import registry - from unilab.utils.algo_utils import ensure_registries - from unilab.utils.obs_utils import get_obs_dims + from unilab.base.observations import get_obs_dims + from unilab.base.registry import ensure_registries ensure_registries() env: Any = registry.make( diff --git a/src/unilab/algos/torch/offpolicy/__init__.py b/src/unilab/algos/torch/offpolicy/__init__.py index 4169e5262..96d6afd3e 100644 --- a/src/unilab/algos/torch/offpolicy/__init__.py +++ b/src/unilab/algos/torch/offpolicy/__init__.py @@ -3,5 +3,11 @@ from unilab.algos.torch.offpolicy.multi_gpu_runner import MultiGPUOffPolicyRunner from unilab.algos.torch.offpolicy.runner import OffPolicyRunner from unilab.algos.torch.offpolicy.worker import off_policy_collector_fn +from unilab.logging import OffPolicyLogger -__all__ = ["OffPolicyRunner", "MultiGPUOffPolicyRunner", "off_policy_collector_fn"] +__all__ = [ + "OffPolicyLogger", + "OffPolicyRunner", + "MultiGPUOffPolicyRunner", + "off_policy_collector_fn", +] diff --git a/src/unilab/algos/torch/offpolicy/multi_gpu_runner.py b/src/unilab/algos/torch/offpolicy/multi_gpu_runner.py index 5ef195cc4..761458fa2 100644 --- a/src/unilab/algos/torch/offpolicy/multi_gpu_runner.py +++ b/src/unilab/algos/torch/offpolicy/multi_gpu_runner.py @@ -33,7 +33,7 @@ from unilab.ipc import SharedWeightSync from unilab.ipc.async_runner import _SPAWN_CTX from unilab.ipc.replay_buffer import ReplayBuffer -from unilab.utils.offpolicy_logger import OffPolicyLogger +from unilab.logging import OffPolicyLogger def _find_free_port() -> int: diff --git a/src/unilab/algos/torch/offpolicy/runner.py b/src/unilab/algos/torch/offpolicy/runner.py index 2720c4f51..0f8869b3b 100644 --- a/src/unilab/algos/torch/offpolicy/runner.py +++ b/src/unilab/algos/torch/offpolicy/runner.py @@ -9,12 +9,13 @@ import torch +from unilab.algos.torch.common.device import get_env_dims from unilab.algos.torch.offpolicy.worker import off_policy_collector_fn from unilab.ipc import SharedObsNormStats, SharedWeightSync from unilab.ipc.async_runner import _SPAWN_CTX, AsyncRunner from unilab.ipc.replay_buffer import ReplayBuffer -from unilab.utils.device_utils import get_default_device, get_env_dims -from unilab.utils.offpolicy_logger import OffPolicyLogger +from unilab.logging import OffPolicyLogger +from unilab.utils.device import get_default_device def compute_train_start_threshold(batch_size: int, learning_starts: int, num_envs: int) -> int: diff --git a/src/unilab/algos/torch/offpolicy/worker.py b/src/unilab/algos/torch/offpolicy/worker.py index 818ff65c1..370df7d4a 100644 --- a/src/unilab/algos/torch/offpolicy/worker.py +++ b/src/unilab/algos/torch/offpolicy/worker.py @@ -11,9 +11,10 @@ import numpy as np import torch -from unilab.utils.algo_utils import build_actor, ensure_registries -from unilab.utils.final_observation import resolve_terminal_observation_contract -from unilab.utils.obs_utils import get_obs_dims, split_obs_dict +from unilab.algos.torch.common.actor_factory import build_actor +from unilab.base.final_observation import resolve_terminal_observation_contract +from unilab.base.observations import get_obs_dims, split_obs_dict +from unilab.base.registry import ensure_registries def resolve_collector_actor_dims( diff --git a/src/unilab/base/__init__.py b/src/unilab/base/__init__.py index 998dc8c51..0ce707b51 100644 --- a/src/unilab/base/__init__.py +++ b/src/unilab/base/__init__.py @@ -1 +1,31 @@ """Environment registry and base classes.""" + +from unilab.base.final_observation import ( + TerminalObservationContract, + TransitionBootstrapContract, + patch_transition_next_obs, + resolve_terminal_observation_contract, + resolve_transition_bootstrap_contract, +) +from unilab.base.observations import ( + flatten_obs_dict, + flatten_policy_obs_dict, + get_critic_base_dim, + get_obs_dims, + split_obs_dict, +) +from unilab.base.registry import ensure_registries + +__all__ = [ + "TerminalObservationContract", + "TransitionBootstrapContract", + "ensure_registries", + "flatten_obs_dict", + "flatten_policy_obs_dict", + "get_critic_base_dim", + "get_obs_dims", + "patch_transition_next_obs", + "resolve_terminal_observation_contract", + "resolve_transition_bootstrap_contract", + "split_obs_dict", +] diff --git a/src/unilab/base/backend/__init__.py b/src/unilab/base/backend/__init__.py index 50ad5b60e..c46328b60 100644 --- a/src/unilab/base/backend/__init__.py +++ b/src/unilab/base/backend/__init__.py @@ -2,6 +2,15 @@ from .base import SimBackend from .motrix_backend import MOTRIX_AVAILABLE, MotrixBackend +from .xml import ( + add_sensor, + create_discardvisual_xml, + get_named_body_ids, + inject_motrix_tracking_sensors, + inject_mujoco_tracking_sensors, + materialize_scene_visual_override, + processed_xml, +) def _load_mujoco_backend() -> Any: @@ -49,5 +58,12 @@ def __getattr__(name: str): "SimBackend", "MuJoCoBackend", "MotrixBackend", + "add_sensor", + "create_discardvisual_xml", "create_backend", + "get_named_body_ids", + "inject_motrix_tracking_sensors", + "inject_mujoco_tracking_sensors", + "materialize_scene_visual_override", + "processed_xml", ] diff --git a/src/unilab/base/backend/base.py b/src/unilab/base/backend/base.py index f36f72d77..422b8334b 100644 --- a/src/unilab/base/backend/base.py +++ b/src/unilab/base/backend/base.py @@ -97,7 +97,7 @@ def get_body_ids(self, names: Sequence[str]) -> np.ndarray: def get_motion_body_ids(self, names: Sequence[str]) -> np.ndarray: """Resolve MuJoCo-style body IDs used by motion datasets.""" - from unilab.utils.xml_utils import get_named_body_ids + from unilab.base.backend.xml import get_named_body_ids return np.asarray(get_named_body_ids(self._model_file, names), dtype=np.int32) diff --git a/src/unilab/base/backend/motrix_backend.py b/src/unilab/base/backend/motrix_backend.py index aa8cc7bd7..e19dc6c1e 100644 --- a/src/unilab/base/backend/motrix_backend.py +++ b/src/unilab/base/backend/motrix_backend.py @@ -56,7 +56,7 @@ def __init__( self._model_file = model_file if self.add_body_sensors: - from unilab.utils.xml_utils import inject_motrix_tracking_sensors + from unilab.base.backend.xml import inject_motrix_tracking_sensors tmp_path, _, valid_bnames = inject_motrix_tracking_sensors( model_file, baselink_name=base_name diff --git a/src/unilab/base/backend/mujoco_backend.py b/src/unilab/base/backend/mujoco_backend.py index fc14f8c44..a8bac0b23 100644 --- a/src/unilab/base/backend/mujoco_backend.py +++ b/src/unilab/base/backend/mujoco_backend.py @@ -24,8 +24,8 @@ ModelVariantSpec, ResetRandomizationPayload, ) +from unilab.dtype_config import get_global_dtype -from ..dtype_config import get_global_dtype from .base import BackendPlayCapabilities, SimBackend @@ -41,7 +41,7 @@ def _prepare_variant_model_xml( add_body_sensors: bool, base_name: str | None, ) -> tuple[str, list[str]]: - from unilab.utils.xml_utils import create_discardvisual_xml, inject_mujoco_tracking_sensors + from unilab.base.backend.xml import create_discardvisual_xml, inject_mujoco_tracking_sensors model_path = create_discardvisual_xml(model_file) tmp_paths = [model_path] @@ -247,7 +247,7 @@ def _load_base_model(self) -> mujoco.MjModel: return model def _prepare_model_xml(self) -> tuple[str, list[str], list[int], list[str]]: - from unilab.utils.xml_utils import create_discardvisual_xml, inject_mujoco_tracking_sensors + from unilab.base.backend.xml import create_discardvisual_xml, inject_mujoco_tracking_sensors model_path = create_discardvisual_xml(self._model_file) tmp_paths = [model_path] @@ -321,13 +321,28 @@ def _compile_model_variants( chunks = tuple( tuple(variants[idx : idx + chunk_size]) for idx in range(0, len(variants), chunk_size) ) - with ProcessPoolExecutor( - max_workers=max_workers, - mp_context=get_context("spawn"), - ) as executor: - futures = [ - executor.submit( - _compile_model_variant_chunk_to_mjb, + try: + with ProcessPoolExecutor( + max_workers=max_workers, + mp_context=get_context("spawn"), + ) as executor: + futures = [ + executor.submit( + _compile_model_variant_chunk_to_mjb, + model_file=self._model_file, + add_body_sensors=self.add_body_sensors, + base_name=self._base_name, + sim_dt=self._sim_dt, + iterations=self._iterations, + position_actuator_gains=self._position_actuator_gains, + variants=chunk, + ) + for chunk in chunks + ] + mjb_paths_nested = [future.result() for future in futures] + except PermissionError: + mjb_paths_nested = [ + _compile_model_variant_chunk_to_mjb( model_file=self._model_file, add_body_sensors=self.add_body_sensors, base_name=self._base_name, @@ -338,7 +353,6 @@ def _compile_model_variants( ) for chunk in chunks ] - mjb_paths_nested = [future.result() for future in futures] flat_paths = [path for paths in mjb_paths_nested for path in paths] try: models = tuple(mujoco.MjModel.from_binary_path(path) for path in flat_paths) diff --git a/src/unilab/utils/xml_utils.py b/src/unilab/base/backend/xml.py similarity index 100% rename from src/unilab/utils/xml_utils.py rename to src/unilab/base/backend/xml.py diff --git a/src/unilab/utils/final_observation.py b/src/unilab/base/final_observation.py similarity index 100% rename from src/unilab/utils/final_observation.py rename to src/unilab/base/final_observation.py diff --git a/src/unilab/base/np_env.py b/src/unilab/base/np_env.py index 5b03910a5..283b10cc4 100644 --- a/src/unilab/base/np_env.py +++ b/src/unilab/base/np_env.py @@ -10,8 +10,8 @@ from unilab.base.backend import SimBackend from unilab.base.base import ABEnv, EnvCfg, EnvPlayCapabilities -from unilab.base.dtype_config import get_global_dtype from unilab.dr import DomainRandomizationManager, DomainRandomizationProvider +from unilab.dtype_config import get_global_dtype if TYPE_CHECKING: from unilab.base.augmentation import SymmetryAugmentation diff --git a/src/unilab/utils/obs_utils.py b/src/unilab/base/observations.py similarity index 100% rename from src/unilab/utils/obs_utils.py rename to src/unilab/base/observations.py diff --git a/src/unilab/base/registry.py b/src/unilab/base/registry.py index 56435af34..b98d8d664 100644 --- a/src/unilab/base/registry.py +++ b/src/unilab/base/registry.py @@ -1,3 +1,6 @@ +import importlib +import logging +from collections.abc import Sequence from dataclasses import dataclass, field from typing import Any, Callable, Dict, Optional, Type, TypeVar @@ -5,6 +8,14 @@ TEnvCfg = TypeVar("TEnvCfg", bound=EnvCfg) _DEFAULT_SIM_BACKEND_ORDER: tuple[str, ...] = ("mujoco", "motrix") +_REGISTRY_MODULES_ATTR = "__unilab_registry_modules__" +_DEFAULT_REGISTRY_PACKAGES = ( + "unilab.envs.locomotion", + "unilab.envs.manipulation", + "unilab.envs.motion_tracking", +) + +logger = logging.getLogger(__name__) @dataclass @@ -187,3 +198,54 @@ def list_registered_envs() -> Dict[str, Dict[str, Any]]: "available_backends": list(meta.env_cls_dict.keys()), } return result + + +def ensure_registries( + packages: Sequence[str] | None = None, + *, + optional_packages: Sequence[str] | None = None, + fail_on_error: bool = True, +) -> None: + """Import env registry bootstrap modules.""" + package_names = list(packages) if packages is not None else list(_DEFAULT_REGISTRY_PACKAGES) + optional = set(optional_packages) if optional_packages else set() + + for package_name in package_names: + is_optional = package_name in optional + try: + package = importlib.import_module(package_name) + except ImportError as exc: + if is_optional: + logger.warning("Optional registry package not found: %s (%s)", package_name, exc) + elif fail_on_error: + raise ImportError( + f"Failed to import registry package '{package_name}'. " + f"Add to optional_packages if this is expected to be absent." + ) from exc + else: + logger.warning("Registry package not found: %s (%s)", package_name, exc) + continue + + modules = getattr(package, _REGISTRY_MODULES_ATTR, ()) + if isinstance(modules, str) or not isinstance(modules, Sequence): + raise TypeError( + f"'{package_name}.{_REGISTRY_MODULES_ATTR}' must be a sequence of module names." + ) + + for module_name in modules: + if not isinstance(module_name, str) or not module_name: + raise TypeError( + f"'{package_name}.{_REGISTRY_MODULES_ATTR}' entries must be non-empty strings." + ) + try: + importlib.import_module(module_name) + except Exception as exc: + if fail_on_error and not is_optional: + raise RuntimeError( + f"Failed to import declared registry module '{module_name}' " + f"from '{package_name}'. " + f"Fix the import error or add '{package_name}' to optional_packages." + ) from exc + logger.warning( + "Failed to import declared registry module '%s': %s", module_name, exc + ) diff --git a/src/unilab/config/__init__.py b/src/unilab/config/__init__.py deleted file mode 100644 index e69de29bb..000000000 diff --git a/src/unilab/config/locomotion_params.py b/src/unilab/config/locomotion_params.py deleted file mode 100644 index 4eb197534..000000000 --- a/src/unilab/config/locomotion_params.py +++ /dev/null @@ -1,16 +0,0 @@ -"""Locomotion task registry. - -Config factory functions (ppo_config, appo_config, offpolicy_config) have been -removed. Configurations are now managed via Hydra YAML files in conf/. -KNOWN_TASKS is kept for legacy routing utilities. -""" - -KNOWN_TASKS: frozenset[str] = frozenset( - { - "Go1JoystickFlat", - "Go2JoystickFlat", - "G1WalkFlat", - "G1WalkRough", - "G1MotionTracking", - } -) diff --git a/src/unilab/config/manipulation_params.py b/src/unilab/config/manipulation_params.py deleted file mode 100644 index 3d7098cc3..000000000 --- a/src/unilab/config/manipulation_params.py +++ /dev/null @@ -1,11 +0,0 @@ -"""Manipulation task registry. - -Config factory functions have been removed. -Configurations are now managed via Hydra YAML files in conf/. -KNOWN_TASKS is kept for legacy routing utilities. -""" - -KNOWN_TASKS: set = { - "AllegroInhandRotation", - "AllegroInhandRotationSac", -} diff --git a/src/unilab/docs/support_matrix.py b/src/unilab/docs/support_matrix.py index 7a9e4a355..36764cfa1 100644 --- a/src/unilab/docs/support_matrix.py +++ b/src/unilab/docs/support_matrix.py @@ -11,7 +11,7 @@ from omegaconf import OmegaConf from unilab.base import registry -from unilab.utils.algo_utils import ensure_registries +from unilab.base.registry import ensure_registries BEGIN_MARKER = "" END_MARKER = "" @@ -309,7 +309,7 @@ def render_support_matrix(root: Path | None = None) -> str: "", "### Source Index", "", - "- Registry bootstrap: `src/unilab/envs/**` decorators via `unilab.utils.algo_utils.ensure_registries()`.", + "- Registry bootstrap: `src/unilab/envs/**` decorators via `unilab.base.registry.ensure_registries()`.", "- Owner YAML scan: `conf/ppo/task/**`, `conf/appo/task/**`, `conf/offpolicy/task/**`.", "- Generic compose coverage: `tests/config/test_config_system.py::test_supported_task_composes`.", "- MLX-specific compose coverage only upgrades task owners listed in `tests/config/test_config_system.py::_PPO_MLX_TASKS`: " diff --git a/src/unilab/dr/__init__.py b/src/unilab/dr/__init__.py index a36a5a8b1..ff0426aa9 100644 --- a/src/unilab/dr/__init__.py +++ b/src/unilab/dr/__init__.py @@ -1,3 +1,8 @@ +"""Domain randomization package. + +Invariant: this package must not depend on unilab.base.* +""" + from .manager import DomainRandomizationManager from .provider import DomainRandomizationProvider from .types import ( diff --git a/src/unilab/dr/dr_utils.py b/src/unilab/dr/dr_utils.py index e2707eb5e..1df84eb6c 100644 --- a/src/unilab/dr/dr_utils.py +++ b/src/unilab/dr/dr_utils.py @@ -4,12 +4,12 @@ import numpy as np -from unilab.base.dtype_config import get_global_dtype from unilab.dr.types import ( DomainRandomizationCapabilities, IntervalRandomizationPlan, ResetRandomizationPayload, ) +from unilab.dtype_config import get_global_dtype def zero_actions(num_reset: int, num_action: int) -> np.ndarray: diff --git a/src/unilab/base/dtype_config.py b/src/unilab/dtype_config.py similarity index 100% rename from src/unilab/base/dtype_config.py rename to src/unilab/dtype_config.py diff --git a/src/unilab/envs/common/__init__.py b/src/unilab/envs/common/__init__.py new file mode 100644 index 000000000..a05e8faba --- /dev/null +++ b/src/unilab/envs/common/__init__.py @@ -0,0 +1 @@ +"""Shared environment-owned numeric helpers.""" diff --git a/src/unilab/envs/common/math.py b/src/unilab/envs/common/math.py new file mode 100644 index 000000000..e96d7f105 --- /dev/null +++ b/src/unilab/envs/common/math.py @@ -0,0 +1,13 @@ +from __future__ import annotations + +import numpy as np + + +def np_sample_uniform( + lower: float | np.ndarray, + upper: float | np.ndarray, + size: tuple[int, ...], + dtype=np.float32, +) -> np.ndarray: + """Sample uniformly from [lower, upper] with output dtype.""" + return np.random.uniform(lower, upper, size).astype(dtype) diff --git a/src/unilab/utils/math_utils.py b/src/unilab/envs/common/rotation.py similarity index 86% rename from src/unilab/utils/math_utils.py rename to src/unilab/envs/common/rotation.py index 429df0927..28aa7cd55 100644 --- a/src/unilab/utils/math_utils.py +++ b/src/unilab/envs/common/rotation.py @@ -1,49 +1,8 @@ from __future__ import annotations -import importlib - import numpy as np -def _require_mlx_core(): - """Import MLX lazily so non-MLX workflows don't crash at module import time.""" - try: - return importlib.import_module("mlx.core") - except Exception as exc: - raise RuntimeError( - "MLX backend is unavailable. Use NumPy helpers (np_quat_mul/np_yaw_to_quat) in non-MLX paths." - ) from exc - - -def quat_mul(q1, q2): - """ - Multiply two quaternions. - """ - mx = _require_mlx_core() - w1, x1, y1, z1 = q1[:, 0], q1[:, 1], q1[:, 2], q1[:, 3] - w2, x2, y2, z2 = q2[:, 0], q2[:, 1], q2[:, 2], q2[:, 3] - return mx.stack( - [ - w1 * w2 - x1 * x2 - y1 * y2 - z1 * z2, - w1 * x2 + x1 * w2 + y1 * z2 - z1 * y2, - w1 * y2 - x1 * z2 + y1 * w2 + z1 * x2, - w1 * z2 + x1 * y2 - y1 * x2 + z1 * w2, - ], - axis=1, - ) - - -def axis_angle_to_quat(axis, angle): - """ - Convert axis-angle to quaternion. - """ - mx = _require_mlx_core() - half_angle = angle / 2 - c = mx.cos(half_angle) - s = mx.sin(half_angle) - return mx.stack([c, axis[:, 0] * s, axis[:, 1] * s, axis[:, 2] * s], axis=1) - - def np_quat_mul(q1: np.ndarray, q2: np.ndarray) -> np.ndarray: """Multiply quaternions in NumPy, supports (N, 4) and (4,) inputs.""" q1_was_1d = q1.ndim == 1 @@ -334,13 +293,3 @@ def np_subtract_frame_transforms( rel_pos = np_quat_apply_inverse(quat1, pos2 - pos1) rel_quat = np_quat_mul(np_quat_inv(quat1), quat2) return rel_pos, rel_quat - - -def np_sample_uniform( - lower: float | np.ndarray, - upper: float | np.ndarray, - size: tuple[int, ...], - dtype=np.float32, -) -> np.ndarray: - """Sample uniformly from [lower, upper] with output dtype.""" - return np.random.uniform(lower, upper, size).astype(dtype) diff --git a/src/unilab/envs/locomotion/common/base.py b/src/unilab/envs/locomotion/common/base.py index 0d8718d6e..7bc24c8ab 100644 --- a/src/unilab/envs/locomotion/common/base.py +++ b/src/unilab/envs/locomotion/common/base.py @@ -8,8 +8,8 @@ from unilab.base.backend import SimBackend from unilab.base.base import EnvCfg -from unilab.base.dtype_config import get_global_dtype from unilab.base.np_env import NpEnv, NpEnvState +from unilab.dtype_config import get_global_dtype @dataclass diff --git a/src/unilab/envs/locomotion/common/commands.py b/src/unilab/envs/locomotion/common/commands.py index 13316bab9..d12b55f68 100644 --- a/src/unilab/envs/locomotion/common/commands.py +++ b/src/unilab/envs/locomotion/common/commands.py @@ -4,7 +4,7 @@ import numpy as np -from unilab.base.dtype_config import get_global_dtype +from unilab.dtype_config import get_global_dtype @dataclass diff --git a/src/unilab/envs/locomotion/common/dr_provider.py b/src/unilab/envs/locomotion/common/dr_provider.py index dc17caa20..db22dffbe 100644 --- a/src/unilab/envs/locomotion/common/dr_provider.py +++ b/src/unilab/envs/locomotion/common/dr_provider.py @@ -11,7 +11,6 @@ import numpy as np -from unilab.base.dtype_config import get_global_dtype from unilab.dr import ( DomainRandomizationCapabilities, DomainRandomizationProvider, @@ -25,7 +24,8 @@ validate_interval_push_support, zero_actions, ) -from unilab.utils.math_utils import np_quat_mul, np_yaw_to_quat +from unilab.dtype_config import get_global_dtype +from unilab.envs.common.rotation import np_quat_mul, np_yaw_to_quat class LocomotionDRProvider(DomainRandomizationProvider): diff --git a/src/unilab/envs/locomotion/common/rewards.py b/src/unilab/envs/locomotion/common/rewards.py index 9abad4357..1f4d38800 100644 --- a/src/unilab/envs/locomotion/common/rewards.py +++ b/src/unilab/envs/locomotion/common/rewards.py @@ -13,7 +13,7 @@ import numpy as np -from unilab.base.dtype_config import get_global_dtype +from unilab.dtype_config import get_global_dtype @dataclass diff --git a/src/unilab/envs/locomotion/g1/joystick.py b/src/unilab/envs/locomotion/g1/joystick.py index 40ebe7706..d9361535f 100644 --- a/src/unilab/envs/locomotion/g1/joystick.py +++ b/src/unilab/envs/locomotion/g1/joystick.py @@ -14,8 +14,8 @@ from unilab.base.augmentation import SymmetryObsLayout from unilab.base.backend import create_backend from unilab.base.curriculum import EpisodeLengthTracker, PenaltyCurriculum -from unilab.base.dtype_config import get_global_dtype from unilab.base.np_env import NpEnvState +from unilab.dtype_config import get_global_dtype from unilab.envs.locomotion.common import rewards from unilab.envs.locomotion.common.commands import Commands from unilab.envs.locomotion.common.domain_rand import DomainRandConfig diff --git a/src/unilab/envs/locomotion/go1/joystick.py b/src/unilab/envs/locomotion/go1/joystick.py index 06882f55d..2af270868 100644 --- a/src/unilab/envs/locomotion/go1/joystick.py +++ b/src/unilab/envs/locomotion/go1/joystick.py @@ -9,15 +9,15 @@ from unilab.assets import ASSETS_ROOT_PATH from unilab.base import registry from unilab.base.backend import create_backend -from unilab.base.dtype_config import get_global_dtype from unilab.base.np_env import NpEnvState +from unilab.dtype_config import get_global_dtype +from unilab.envs.common.rotation import np_quat_mul, np_yaw_to_quat from unilab.envs.locomotion.common import rewards from unilab.envs.locomotion.common.commands import Commands from unilab.envs.locomotion.common.domain_rand import DomainRandConfig from unilab.envs.locomotion.common.dr_provider import LocomotionDRProvider from unilab.envs.locomotion.common.rewards import RewardContext from unilab.envs.locomotion.go1.base import Go1BaseCfg, Go1BaseEnv -from unilab.utils.math_utils import np_quat_mul, np_yaw_to_quat @dataclass diff --git a/src/unilab/envs/locomotion/go2/handstand.py b/src/unilab/envs/locomotion/go2/handstand.py index 4ce200ac0..fec4a4695 100644 --- a/src/unilab/envs/locomotion/go2/handstand.py +++ b/src/unilab/envs/locomotion/go2/handstand.py @@ -9,8 +9,8 @@ from unilab.assets import ASSETS_ROOT_PATH from unilab.base import registry from unilab.base.backend import create_backend -from unilab.base.dtype_config import get_global_dtype from unilab.base.np_env import NpEnvState +from unilab.dtype_config import get_global_dtype from unilab.envs.locomotion.common import rewards from unilab.envs.locomotion.common.commands import Commands from unilab.envs.locomotion.common.domain_rand import DomainRandConfig @@ -156,7 +156,6 @@ def _init_reward_functions(self): } def update_state(self, state: NpEnvState) -> NpEnvState: - linvel = self.get_local_linvel() gyro = self.get_gyro() gravity = self._backend.get_sensor_data("upvector") diff --git a/src/unilab/envs/locomotion/go2/joystick.py b/src/unilab/envs/locomotion/go2/joystick.py index cec7db73c..9ef9c936d 100644 --- a/src/unilab/envs/locomotion/go2/joystick.py +++ b/src/unilab/envs/locomotion/go2/joystick.py @@ -9,8 +9,8 @@ from unilab.assets import ASSETS_ROOT_PATH from unilab.base import registry from unilab.base.backend import create_backend -from unilab.base.dtype_config import get_global_dtype from unilab.base.np_env import NpEnvState +from unilab.dtype_config import get_global_dtype from unilab.envs.locomotion.common import rewards from unilab.envs.locomotion.common.commands import Commands from unilab.envs.locomotion.common.domain_rand import DomainRandConfig diff --git a/src/unilab/envs/manipulation/inhand_rot_allegro/base.py b/src/unilab/envs/manipulation/inhand_rot_allegro/base.py index 81fb9ff1e..4133e1c3c 100644 --- a/src/unilab/envs/manipulation/inhand_rot_allegro/base.py +++ b/src/unilab/envs/manipulation/inhand_rot_allegro/base.py @@ -7,8 +7,8 @@ from unilab.base.backend import SimBackend from unilab.base.base import EnvCfg -from unilab.base.dtype_config import get_global_dtype from unilab.base.np_env import NpEnv, NpEnvState +from unilab.dtype_config import get_global_dtype @dataclass diff --git a/src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py b/src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py index 71686bb91..376ef078e 100644 --- a/src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py +++ b/src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py @@ -11,7 +11,6 @@ from unilab.assets import ASSETS_ROOT_PATH from unilab.base import registry from unilab.base.backend import create_backend -from unilab.base.dtype_config import get_global_dtype from unilab.base.np_env import NpEnvState from unilab.dr import ( DomainRandomizationCapabilities, @@ -26,8 +25,9 @@ validate_interval_push_support, zero_actions, ) +from unilab.dtype_config import get_global_dtype +from unilab.envs.common.rotation import np_quat_conjugate, np_quat_mul, np_quat_to_axis_angle from unilab.envs.manipulation.inhand_rot_allegro.base import AllegroBaseCfg, AllegroBaseEnv -from unilab.utils.math_utils import np_quat_conjugate, np_quat_mul, np_quat_to_axis_angle def normalize_rotation_axis(rotation_axis: tuple[float, float, float]) -> np.ndarray: diff --git a/src/unilab/envs/manipulation/sharpa_inhand/base.py b/src/unilab/envs/manipulation/sharpa_inhand/base.py index 5f7eb9d67..b0bb49ccd 100644 --- a/src/unilab/envs/manipulation/sharpa_inhand/base.py +++ b/src/unilab/envs/manipulation/sharpa_inhand/base.py @@ -10,9 +10,9 @@ from unilab.assets import ASSETS_ROOT_PATH from unilab.base.backend import SimBackend from unilab.base.base import EnvCfg -from unilab.base.dtype_config import get_global_dtype from unilab.base.np_env import NpEnv, NpEnvState -from unilab.utils.math_utils import np_quat_apply, np_quat_mul +from unilab.dtype_config import get_global_dtype +from unilab.envs.common.rotation import np_quat_apply, np_quat_mul DEFAULT_ACTUATED_JOINT_NAMES: list[str] = [ "right_thumb_CMC_FE", diff --git a/src/unilab/envs/manipulation/sharpa_inhand/grasp_gen.py b/src/unilab/envs/manipulation/sharpa_inhand/grasp_gen.py index 6536f058d..4fcdc98b2 100644 --- a/src/unilab/envs/manipulation/sharpa_inhand/grasp_gen.py +++ b/src/unilab/envs/manipulation/sharpa_inhand/grasp_gen.py @@ -9,6 +9,7 @@ from unilab.base.np_env import NpEnvState from unilab.dr import ResetPlan from unilab.dr.dr_utils import build_common_reset_randomization +from unilab.envs.common.rotation import np_quat_error_magnitude from unilab.envs.manipulation.sharpa_inhand.base import ( SOURCE_DEFAULT_HAND_JOINT_POS_DEG, resolve_grasp_cache_file, @@ -19,7 +20,6 @@ SharpaInhandRotationDRProvider, SharpaInhandRotationEnv, ) -from unilab.utils.math_utils import np_quat_error_magnitude @dataclass diff --git a/src/unilab/envs/manipulation/sharpa_inhand/rotation.py b/src/unilab/envs/manipulation/sharpa_inhand/rotation.py index 1ac5d695d..781e1ecf5 100644 --- a/src/unilab/envs/manipulation/sharpa_inhand/rotation.py +++ b/src/unilab/envs/manipulation/sharpa_inhand/rotation.py @@ -7,7 +7,6 @@ from unilab.base import registry from unilab.base.backend import create_backend -from unilab.base.dtype_config import get_global_dtype from unilab.base.np_env import NpEnvState from unilab.dr import ( DomainRandomizationCapabilities, @@ -18,6 +17,8 @@ ResetPlan, ) from unilab.dr.dr_utils import build_common_reset_randomization, validate_common_reset_randomization +from unilab.dtype_config import get_global_dtype +from unilab.envs.common.rotation import np_quat_conjugate, np_quat_mul, np_quat_to_axis_angle from unilab.envs.manipulation.sharpa_inhand.base import ( SharpaInhandBaseCfg, SharpaInhandBaseEnv, @@ -26,7 +27,6 @@ resolve_grasp_cache_file, sample_bucketed_grasp_cache, ) -from unilab.utils.math_utils import np_quat_conjugate, np_quat_mul, np_quat_to_axis_angle @dataclass diff --git a/src/unilab/envs/motion_tracking/g1/tracking.py b/src/unilab/envs/motion_tracking/g1/tracking.py index 679abe7d2..afed3712c 100644 --- a/src/unilab/envs/motion_tracking/g1/tracking.py +++ b/src/unilab/envs/motion_tracking/g1/tracking.py @@ -10,7 +10,6 @@ from unilab.assets import ASSETS_ROOT_PATH from unilab.base import registry from unilab.base.backend import create_backend -from unilab.base.dtype_config import get_global_dtype from unilab.base.np_env import NpEnvState from unilab.dr import ( DomainRandomizationCapabilities, @@ -25,18 +24,19 @@ validate_interval_push_support, zero_actions, ) -from unilab.envs.locomotion.g1.base import G1BaseCfg, G1BaseEnv -from unilab.utils.math_utils import ( +from unilab.dtype_config import get_global_dtype +from unilab.envs.common.math import np_sample_uniform +from unilab.envs.common.rotation import ( np_matrix_from_quat, np_quat_apply, np_quat_error_magnitude, np_quat_from_euler_xyz, np_quat_inv, np_quat_mul, - np_sample_uniform, np_subtract_frame_transforms, np_yaw_quat, ) +from unilab.envs.locomotion.g1.base import G1BaseCfg, G1BaseEnv from .motion_loader import MotionLoader, MotionSampler diff --git a/src/unilab/envs/motion_tracking/g1/tracking_sac.py b/src/unilab/envs/motion_tracking/g1/tracking_sac.py index 747dcbf76..5cd8bd2f3 100644 --- a/src/unilab/envs/motion_tracking/g1/tracking_sac.py +++ b/src/unilab/envs/motion_tracking/g1/tracking_sac.py @@ -14,7 +14,7 @@ import numpy as np from unilab.base import registry -from unilab.base.dtype_config import get_global_dtype +from unilab.dtype_config import get_global_dtype from .tracking import G1MotionTrackingCfg, G1MotionTrackingEnv diff --git a/src/unilab/ipc/replay_buffer.py b/src/unilab/ipc/replay_buffer.py index 35d7d7c02..88078e024 100644 --- a/src/unilab/ipc/replay_buffer.py +++ b/src/unilab/ipc/replay_buffer.py @@ -4,7 +4,7 @@ import torch -from unilab.algos.torch.common.shared_buffer import SharedBufferBase +from unilab.ipc.shared_buffer import SharedBufferBase class ReplayBuffer(SharedBufferBase): diff --git a/src/unilab/algos/torch/common/shared_buffer.py b/src/unilab/ipc/shared_buffer.py similarity index 100% rename from src/unilab/algos/torch/common/shared_buffer.py rename to src/unilab/ipc/shared_buffer.py diff --git a/src/unilab/logging/__init__.py b/src/unilab/logging/__init__.py new file mode 100644 index 000000000..3c262c1d9 --- /dev/null +++ b/src/unilab/logging/__init__.py @@ -0,0 +1,11 @@ +"""Rich-based training loggers shared across algorithm and training layers.""" + +from unilab.logging.common import BaseTrainingLogger +from unilab.logging.offpolicy import OffPolicyLogger +from unilab.logging.onpolicy import OnPolicyLogger + +__all__ = [ + "BaseTrainingLogger", + "OffPolicyLogger", + "OnPolicyLogger", +] diff --git a/src/unilab/utils/logging_common.py b/src/unilab/logging/common.py similarity index 100% rename from src/unilab/utils/logging_common.py rename to src/unilab/logging/common.py diff --git a/src/unilab/utils/offpolicy_logger.py b/src/unilab/logging/offpolicy.py similarity index 62% rename from src/unilab/utils/offpolicy_logger.py rename to src/unilab/logging/offpolicy.py index c160dae36..c13c9360e 100644 --- a/src/unilab/utils/offpolicy_logger.py +++ b/src/unilab/logging/offpolicy.py @@ -1,33 +1,4 @@ -"""Rich-based training logger for off-policy RL algorithms (SAC, TD3, etc). - -Usage: - from unilab.utils.offpolicy_logger import OffPolicyLogger - - logger = OffPolicyLogger( - algo_name="FastSAC", - max_iterations=1500, - num_envs=4096, - log_dir="logs/run_01", # for tensorboard - log_backend="tensorboard", # "tensorboard", "wandb", or "none" - ) - - logger.start() # Begin Live display - - logger.log_buffer_fill(cur, total) # During warmup/buffer fill - logger.log_collector(step, buf, rew) # Collector progress (from subprocess) - - logger.log_step( # Each training iteration - iteration=100, - metrics={"qf_loss": 5.1, "actor_loss": -0.3, "alpha": 0.001}, - reward=8.5, - reward_components={"track_lin_vel": 1.2, "action_rate": -0.05}, - collect_time=0.03, - train_time=0.15, - ) - - logger.log_save(path) # Checkpoint saved - logger.finish() # End Live display -""" +"""Rich-based training logger for off-policy RL algorithms (SAC, TD3, etc).""" from __future__ import annotations @@ -40,21 +11,11 @@ from rich.panel import Panel from rich.table import Table -from unilab.utils.logging_common import BaseTrainingLogger, _fmt_number, _fmt_time, _load_wandb +from unilab.logging.common import BaseTrainingLogger, _fmt_number, _fmt_time, _load_wandb class OffPolicyLogger(BaseTrainingLogger): - """Rich logger for off-policy RL algorithms (SAC, TD3, etc). - - Features: - - Real-time Live table with training metrics - - Loss tracking (any key-value pairs) - - Reward tracking (mean + per-component breakdown) - - Timing: collect/train per step, total elapsed, ETA - - Buffer fill progress bar - - Checkpoint save notifications - - TensorBoard / W&B backend logging - """ + """Rich logger for off-policy RL algorithms (SAC, TD3, etc).""" def __init__( self, @@ -66,7 +27,7 @@ def __init__( action_dim: int = 0, refresh_per_second: int = 4, log_dir: str = "", - log_backend: str = "tensorboard", # "tensorboard", "wandb", "none" + log_backend: str = "tensorboard", wandb_project: str = "unilab", wandb_entity: str | None = None, wandb_name: str = "", @@ -99,7 +60,6 @@ def __init__( ) self.obs_dim = obs_dim self.action_dim = action_dim - self._total_steps: int = 0 self._buffer_size: int = 0 self._buffer_target: int = 0 @@ -115,8 +75,6 @@ def __init__( self._replay_queue_max: int = 0 self._status: str = "Initializing..." - # ---- Lifecycle ---- - def _format_tensorboard_message(self, tb_dir: str) -> str: return f"[dim]TensorBoard logging to: {tb_dir}[/]" @@ -124,20 +82,15 @@ def _format_wandb_message(self, project: str, name: str) -> str: return f"[dim]W&B logging to project: {project}, run: {name}[/]" def start(self, *, status: str = "Warming up..."): - """Begin the Live display.""" super().start(status=status) def finish(self, *, title: str = "Training Summary", extra_summary: str = ""): - """Stop the Live display and print a summary.""" super().finish( title=title, extra_summary=f" Total env steps: [yellow]{self._total_steps:,}[/]\n{extra_summary}", ) - # ---- Logging API ---- - def log_buffer_fill(self, current: int, target: int): - """Update buffer fill progress.""" self._buffer_size = current self._buffer_target = target pct = current / max(target, 1) * 100 @@ -145,30 +98,24 @@ def log_buffer_fill(self, current: int, target: int): self._refresh() def update_collector_timing(self, timing_ms: dict[str, float]): - """Update collector-side environment timing (milliseconds).""" self._collector_timing.update(timing_ms) def update_done_rates(self, timeout_rate: float, terminated_rate: float): - """Update timeout/terminated ratio among completed episodes in collector window.""" self._timeout_rate = float(timeout_rate) self._terminated_rate = float(terminated_rate) def update_buffer_utilization(self, utilization: float): - """Update buffer fill ratio (0.0–1.0). Displayed in the timing panel.""" self._buffer_utilization = float(utilization) def update_replay_queue(self, current_len: int, max_size: int): - """Update replay queue occupancy (APPO-specific).""" self._replay_queue_len = current_len self._replay_queue_max = max_size def set_collection_sync(self, enabled: bool, env_steps_per_sync: int = 0): - """Set collection/training synchronization status for display.""" self._sync_collection = enabled self._env_steps_per_sync = env_steps_per_sync def log_collector(self, total_steps: int, buffer_size: int, mean_reward: float = 0.0): - """Update collector progress (called periodically from metrics queue drain).""" self._total_steps = total_steps self._buffer_size = buffer_size if mean_reward != 0: @@ -186,24 +133,20 @@ def log_step( wait_time: float = 0.0, extra_info: dict | None = None, ): - """Log one training iteration.""" + del extra_info self._iteration = iteration self._collect_time = collect_time self._train_time = train_time self._wait_time = wait_time self._iter_times.append(collect_time + train_time) - if metrics: self._latest_metrics.update(metrics) if reward is not None: self._reward_history.append(reward) if reward_components: self._latest_reward_components = reward_components - self._status = "Training" self._refresh() - - # ---- Write to backend ---- self._backend_log_step( iteration, metrics, reward, reward_components, collect_time, train_time ) @@ -217,123 +160,88 @@ def _backend_log_step( collect_time: float, train_time: float, ): - """Write metrics to TensorBoard / W&B.""" global_step = self._total_steps if self._total_steps > 0 else iteration - elapsed = time.time() - self._start_time if self._start_time else 0 - # ---- TensorBoard ---- if self._tb_writer: - w = self._tb_writer - - # train/ — model outputs (losses, alpha, etc.) + writer = self._tb_writer if metrics: - for k, v in metrics.items(): - w.add_scalar(f"train/{k}", v, global_step) - - # reward/ — reward signals + for key, value in metrics.items(): + writer.add_scalar(f"train/{key}", value, global_step) if reward is not None: - w.add_scalar("reward/mean", reward, global_step) + writer.add_scalar("reward/mean", reward, global_step) if reward_components: - for k, v in reward_components.items(): - w.add_scalar(f"reward/{k}", v, global_step) - - # episode/ — per-episode statistics + for key, value in reward_components.items(): + writer.add_scalar(f"reward/{key}", value, global_step) if self._mean_ep_length > 0: - w.add_scalar("episode/length", self._mean_ep_length, global_step) - w.add_scalar("episode/timeout_rate", self._timeout_rate, global_step) - w.add_scalar("episode/terminated_rate", self._terminated_rate, global_step) - - # timing/ — learner-side and collector-side timing - w.add_scalar("timing/learner_wait_ms", self._wait_time * 1000, global_step) - w.add_scalar("timing/learner_collect_ms", collect_time * 1000, global_step) - w.add_scalar("timing/learner_train_ms", train_time * 1000, global_step) - for key, val in self._collector_timing.items(): - w.add_scalar(f"timing/collector_{key}", val, global_step) - - # perf/ — throughput and efficiency + writer.add_scalar("episode/length", self._mean_ep_length, global_step) + writer.add_scalar("episode/timeout_rate", self._timeout_rate, global_step) + writer.add_scalar("episode/terminated_rate", self._terminated_rate, global_step) + writer.add_scalar("timing/learner_wait_ms", self._wait_time * 1000, global_step) + writer.add_scalar("timing/learner_collect_ms", collect_time * 1000, global_step) + writer.add_scalar("timing/learner_train_ms", train_time * 1000, global_step) + for key, value in self._collector_timing.items(): + writer.add_scalar(f"timing/collector_{key}", value, global_step) if elapsed > 0 and self._total_steps > 0: - w.add_scalar("perf/steps_per_sec", self._total_steps / elapsed, global_step) - w.add_scalar( + writer.add_scalar("perf/steps_per_sec", self._total_steps / elapsed, global_step) + writer.add_scalar( "perf/iter_ms", (self._collect_time + self._train_time) * 1000, global_step ) - w.add_scalar( + writer.add_scalar( "perf/collect_train_ratio", self._collect_time / max(self._train_time, 1e-6), global_step, ) - # ---- W&B ---- if self._wandb_run: wandb = _load_wandb() if wandb is None: return - log_dict: dict[str, Any] = {"iteration": iteration} if metrics: - for k, v in metrics.items(): - log_dict[f"train/{k}"] = v - - # reward/ + for key, value in metrics.items(): + log_dict[f"train/{key}"] = value if reward is not None: log_dict["reward/mean"] = reward if reward_components: - for k, v in reward_components.items(): - log_dict[f"reward/{k}"] = v - - # episode/ + for key, value in reward_components.items(): + log_dict[f"reward/{key}"] = value if self._mean_ep_length > 0: log_dict["episode/length"] = self._mean_ep_length log_dict["episode/timeout_rate"] = self._timeout_rate log_dict["episode/terminated_rate"] = self._terminated_rate - - # timing/ log_dict["timing/learner_wait_ms"] = self._wait_time * 1000 log_dict["timing/learner_collect_ms"] = collect_time * 1000 log_dict["timing/learner_train_ms"] = train_time * 1000 - for key, val in self._collector_timing.items(): - log_dict[f"timing/collector_{key}"] = val - - # perf/ + for key, value in self._collector_timing.items(): + log_dict[f"timing/collector_{key}"] = value if elapsed > 0 and self._total_steps > 0: log_dict["perf/steps_per_sec"] = self._total_steps / elapsed log_dict["perf/iter_ms"] = (self._collect_time + self._train_time) * 1000 log_dict["perf/collect_train_ratio"] = self._collect_time / max(self._train_time, 1e-6) - wandb.log(log_dict, step=global_step) def log_status(self, status: str): - """Set a custom status message.""" self._status = status self._refresh() - # ---- Display Building ---- - def _build_display(self) -> Panel: - """Build the full rich display panel.""" header_panel = self._build_header(include_status=True) - - # Body: side-by-side tables left = self._build_metrics_table() right = self._build_reward_table() bottom = self._build_timing_table() - grid = Table.grid(expand=True) grid.add_column(ratio=1) grid.add_column(ratio=1) grid.add_row(left, right) - - main_group = Group(header_panel, grid, bottom) - return Panel( - main_group, + Group(header_panel, grid, bottom), title="[bold] 🚀 UniLab Off-Policy Training [/]", border_style="bright_blue", padding=(0, 1), ) def _build_metrics_table(self) -> Table: - """Build the losses/metrics table.""" table = Table( title="[bold]Losses & Metrics[/]", box=box.SIMPLE_HEAVY, @@ -344,34 +252,24 @@ def _build_metrics_table(self) -> Table: ) table.add_column("Metric", style="white", ratio=2) table.add_column("Value", style="yellow", justify="right", ratio=1) - if not self._latest_metrics: table.add_row("[dim]Waiting for data...[/]", "") else: - # Sort: losses first, then other metrics - loss_keys = sorted([k for k in self._latest_metrics if "loss" in k.lower()]) - other_keys = sorted([k for k in self._latest_metrics if "loss" not in k.lower()]) - - for k in loss_keys: - v = self._latest_metrics[k] - name = k.replace("_", " ").title() - val_str = _fmt_number(v) - style = "red" if v > 10 else "yellow" - table.add_row(f"{name}", f"[{style}]{val_str}[/]") - - for k in other_keys: - v = self._latest_metrics[k] - name = k.replace("_", " ").title() - table.add_row(f" {name}", _fmt_number(v)) - + loss_keys = sorted([key for key in self._latest_metrics if "loss" in key.lower()]) + other_keys = sorted([key for key in self._latest_metrics if "loss" not in key.lower()]) + for key in loss_keys: + value = self._latest_metrics[key] + style = "red" if value > 10 else "yellow" + table.add_row(key.replace("_", " ").title(), f"[{style}]{_fmt_number(value)}[/]") + for key in other_keys: + value = self._latest_metrics[key] + table.add_row(f" {key.replace('_', ' ').title()}", _fmt_number(value)) return table def _build_reward_table(self) -> Table: - """Build the reward breakdown table.""" return self._build_reward_table_common(wait_message="[dim]Waiting for data...[/]") def _build_timing_table(self) -> Table: - """Build the timing info table.""" table = Table( title="[bold]Timing & System[/]", box=box.SIMPLE_HEAVY, @@ -386,15 +284,7 @@ def _build_timing_table(self) -> Table: table.add_column("Value", style="yellow", justify="right", ratio=1, no_wrap=True) elapsed = time.time() - self._start_time if self._start_time else 0 - - table.add_row( - "Elapsed", - _fmt_time(elapsed), - "Buffer", - f"{self._buffer_size:,}", - ) - - # Wait time with color coding + table.add_row("Elapsed", _fmt_time(elapsed), "Buffer", f"{self._buffer_size:,}") wait_ms = self._wait_time * 1000 wait_color = "red" if wait_ms > 1.0 else "yellow" table.add_row( @@ -410,39 +300,32 @@ def _build_timing_table(self) -> Table: "", ) timing_items = list(self._collector_timing.items()) - for i in range(0, len(timing_items), 2): - left_key, left_val = timing_items[i] - if i + 1 < len(timing_items): - right_key, right_val = timing_items[i + 1] + for index in range(0, len(timing_items), 2): + left_key, left_value = timing_items[index] + if index + 1 < len(timing_items): + right_key, right_value = timing_items[index + 1] table.add_row( f"[dim]collector[/] {left_key}", - f"{left_val:.1f}ms", + f"{left_value:.1f}ms", f"[dim]collector[/] {right_key}", - f"{right_val:.1f}ms", + f"{right_value:.1f}ms", ) else: - table.add_row( - f"[dim]collector[/] {left_key}", - f"{left_val:.1f}ms", - "", - "", - ) + table.add_row(f"[dim]collector[/] {left_key}", f"{left_value:.1f}ms", "", "") table.add_row( "Timeout Rate", f"{self._timeout_rate * 100:.1f}%", "Terminated Rate", f"{self._terminated_rate * 100:.1f}%", ) - - util = self._buffer_utilization - if util >= 1.5: - util_str = f"[bold red]{util:.2f} (collector >> learner)[/]" - elif util >= 1.0: - util_str = f"[yellow]{util:.2f}[/]" + utilization = self._buffer_utilization + if utilization >= 1.5: + utilization_str = f"[bold red]{utilization:.2f} (collector >> learner)[/]" + elif utilization >= 1.0: + utilization_str = f"[yellow]{utilization:.2f}[/]" else: - util_str = f"[green]{util:.2f}[/]" - table.add_row("Write/Read", util_str, "", "") - + utilization_str = f"[green]{utilization:.2f}[/]" + table.add_row("Write/Read", utilization_str, "", "") table.add_row( "Envs", f"{self.num_envs:,}", @@ -451,19 +334,14 @@ def _build_timing_table(self) -> Table: if self._sync_collection else "✗", ) - if self._replay_queue_max > 0: - rq_color = "green" if self._replay_queue_len < self._replay_queue_max else "yellow" + replay_color = "green" if self._replay_queue_len < self._replay_queue_max else "yellow" table.add_row( "Replay Queue", - f"[{rq_color}]{self._replay_queue_len}/{self._replay_queue_max}[/]", + f"[{replay_color}]{self._replay_queue_len}/{self._replay_queue_max}[/]", "", "", ) - - # Steps per second if elapsed > 0 and self._total_steps > 0: - sps = self._total_steps / elapsed - table.add_row("Steps/s", f"{sps:,.0f}", "", "") - + table.add_row("Steps/s", f"{self._total_steps / elapsed:,.0f}", "", "") return table diff --git a/src/unilab/utils/onpolicy_logger.py b/src/unilab/logging/onpolicy.py similarity index 98% rename from src/unilab/utils/onpolicy_logger.py rename to src/unilab/logging/onpolicy.py index f6ddd4b6a..e72c54530 100644 --- a/src/unilab/utils/onpolicy_logger.py +++ b/src/unilab/logging/onpolicy.py @@ -8,7 +8,7 @@ from rich.panel import Panel from rich.table import Table -from unilab.utils.logging_common import BaseTrainingLogger, _fmt_number, _fmt_time, _load_wandb +from unilab.logging.common import BaseTrainingLogger, _fmt_number, _fmt_time, _load_wandb class OnPolicyLogger(BaseTrainingLogger): diff --git a/src/unilab/config/structured_configs.py b/src/unilab/structured_configs.py similarity index 100% rename from src/unilab/config/structured_configs.py rename to src/unilab/structured_configs.py diff --git a/src/unilab/training/__init__.py b/src/unilab/training/__init__.py index 1726b7fcc..9da3921db 100644 --- a/src/unilab/training/__init__.py +++ b/src/unilab/training/__init__.py @@ -5,20 +5,25 @@ assert_offpolicy_task_choice_matches_algo, create_env, ensure_registries, - get_entrypoint_log_root, get_hydra_runtime_choice, + setup_logger, +) +from unilab.training.experiment import ExperimentTracker +from unilab.training.monitoring import HardwareMonitor +from unilab.training.run import ( + get_entrypoint_log_root, get_latest_checkpoint, get_latest_run, get_log_root, parse_checkpoint_path, - render_play_mode, resolve_checkpoint_path, resolve_task_checkpoint_path, - setup_logger, ) __all__ = [ "BackendAdapter", + "ExperimentTracker", + "HardwareMonitor", "assert_offpolicy_task_choice_matches_algo", "create_env", "ensure_registries", @@ -28,7 +33,6 @@ "get_latest_run", "get_log_root", "parse_checkpoint_path", - "render_play_mode", "resolve_checkpoint_path", "resolve_task_checkpoint_path", "setup_logger", diff --git a/src/unilab/training/backend_adapter.py b/src/unilab/training/backend_adapter.py index 18008750f..c915cf92a 100644 --- a/src/unilab/training/backend_adapter.py +++ b/src/unilab/training/backend_adapter.py @@ -7,8 +7,8 @@ from omegaconf import DictConfig, OmegaConf -from unilab.utils.reward_utils import extract_reward_config -from unilab.utils.xml_utils import materialize_scene_visual_override +from unilab.base.backend.xml import materialize_scene_visual_override +from unilab.training.reward import extract_reward_config class BackendAdapter: diff --git a/src/unilab/training/common.py b/src/unilab/training/common.py index 6aeddde1f..ff68b27a5 100644 --- a/src/unilab/training/common.py +++ b/src/unilab/training/common.py @@ -3,18 +3,13 @@ from __future__ import annotations import logging -import tempfile -import time from pathlib import Path -from typing import Any, Callable, TypeVar +from typing import Any -import numpy as np from hydra.core.hydra_config import HydraConfig from omegaconf import DictConfig, OmegaConf -from unilab.utils.algo_utils import ensure_registries as _ensure_registries - -ObsT = TypeVar("ObsT") +from unilab.base.registry import ensure_registries as _ensure_registries def ensure_registries() -> None: @@ -67,165 +62,6 @@ def assert_offpolicy_task_choice_matches_algo( ) -def get_log_root(root_dir: str | Path, cfg: DictConfig) -> Path: - """Resolve the algorithm log root, honoring optional training.log_root overrides.""" - configured_root = OmegaConf.select(cfg, "training.log_root") - if configured_root: - log_root = Path(str(configured_root)) - return log_root if log_root.is_absolute() else Path(root_dir) / log_root - return Path(root_dir) / "logs" / str(OmegaConf.select(cfg, "algo.algo_log_name")) - - -def get_entrypoint_log_root( - root_dir: str | Path, - *, - algo_log_name: str, - log_root: str | Path | None = None, -) -> Path: - """Resolve the log root for non-Hydra entrypoints using training helper semantics.""" - if log_root is not None: - configured_root = Path(log_root) - return ( - configured_root if configured_root.is_absolute() else Path(root_dir) / configured_root - ) - return Path(root_dir) / "logs" / algo_log_name - - -def get_latest_run(log_dir: str | Path) -> Path | None: - """Return the lexicographically latest run directory under a task log root.""" - base_dir = Path(log_dir) - if not base_dir.exists(): - return None - runs = sorted(path for path in base_dir.iterdir() if path.is_dir()) - return runs[-1] if runs else None - - -def get_latest_checkpoint(run_dir: str | Path, *, suffix: str = ".pt") -> Path | None: - """Return the latest model checkpoint inside a run directory.""" - run_path = Path(run_dir) - if not run_path.exists(): - return None - - def _iteration(path: Path) -> int: - stem_parts = path.stem.split("_", 1) - if len(stem_parts) != 2: - return -1 - try: - return int(stem_parts[1]) - except ValueError: - return -1 - - model_files = [ - path - for path in run_path.iterdir() - if path.is_file() and path.name.startswith("model_") and path.suffix == suffix - ] - if not model_files: - return None - return max(model_files, key=_iteration) - - -def resolve_checkpoint_path( - base_log_dir: str | Path, - load_run: str, - *, - suffix: str = ".pt", -) -> tuple[Path | None, Path | None]: - """Resolve a latest or explicit checkpoint path from a task log root.""" - base_dir = Path(base_log_dir) - if load_run == "-1": - run_dir = get_latest_run(base_dir) - if run_dir is None: - return None, None - checkpoint = get_latest_checkpoint(run_dir, suffix=suffix) - return (checkpoint, run_dir) if checkpoint is not None else (None, None) - - candidate = Path(load_run) - if not candidate.exists(): - candidate = base_dir / load_run - if candidate.is_file(): - return candidate, candidate.parent - if candidate.is_dir(): - checkpoint = get_latest_checkpoint(candidate, suffix=suffix) - return (checkpoint, candidate) if checkpoint is not None else (None, None) - return None, None - - -def parse_checkpoint_path( - cfg: DictConfig, - *, - root_dir: str | Path, - load_run: str | None = None, - task_name: str | None = None, - checkpoint: str | int | None = None, - suffix: str = ".pt", -) -> tuple[Path | None, Path | None]: - """Resolve a checkpoint path from Hydra config and repository root.""" - selected_task = task_name or str(OmegaConf.select(cfg, "training.task_name")) - selected_run = load_run or str(OmegaConf.select(cfg, "algo.load_run", default="-1")) - selected_checkpoint = checkpoint - if selected_checkpoint is None: - selected_checkpoint = OmegaConf.select(cfg, "algo.checkpoint", default=-1) - if selected_checkpoint in (None, "", -1, "-1"): - selected_checkpoint = None - - return resolve_task_checkpoint_path( - root_dir, - task_name=selected_task, - load_run=selected_run, - algo_log_name=str(OmegaConf.select(cfg, "algo.algo_log_name")), - checkpoint=str(selected_checkpoint) if selected_checkpoint is not None else None, - suffix=suffix, - log_root=OmegaConf.select(cfg, "training.log_root"), - ) - - -def resolve_task_checkpoint_path( - root_dir: str | Path, - *, - task_name: str, - load_run: str, - algo_log_name: str, - checkpoint: str | None = None, - suffix: str = ".pt", - log_root: str | Path | None = None, -) -> tuple[Path | None, Path | None]: - """Resolve checkpoint paths for auxiliary entrypoints through shared training semantics.""" - task_log_root = ( - get_entrypoint_log_root( - root_dir, - algo_log_name=algo_log_name, - log_root=log_root, - ) - / task_name - ) - - run_dir: Path | None - if load_run == "-1": - run_dir = get_latest_run(task_log_root) - else: - candidate = Path(load_run) - if not candidate.exists(): - candidate = task_log_root / load_run - if candidate.is_file(): - return candidate, candidate.parent - run_dir = candidate if candidate.is_dir() else None - - if run_dir is None: - return None, None - - checkpoint_path: Path | None - if checkpoint is not None: - checkpoint_name = ( - f"model_{checkpoint}{suffix}" if str(checkpoint).isdigit() else str(checkpoint) - ) - checkpoint_path = run_dir / checkpoint_name - return (checkpoint_path, run_dir) if checkpoint_path.exists() else (None, run_dir) - - checkpoint_path = get_latest_checkpoint(run_dir, suffix=suffix) - return (checkpoint_path, run_dir) if checkpoint_path is not None else (None, run_dir) - - def setup_logger( log_dir: str | Path, algo_name: str, @@ -277,194 +113,3 @@ def create_env( sim_backend=sim_backend or str(OmegaConf.select(cfg, "training.sim_backend")), env_cfg_override=env_cfg_override, ) - - -def render_play_mode( - env, - *, - sim_backend: str, - initialize: Callable[[], ObsT], - step: Callable[[ObsT], ObsT], - num_steps: int | None, - output_video: str | Path | None = None, - render_spacing: float | None = None, - frame_state_getter: Callable[[], np.ndarray] | None = None, - camera_kwargs: dict[str, Any] | None = None, -) -> str | None: - """Render interactive Motrix play or MuJoCo video generation through shared callbacks.""" - if sim_backend == "motrix": - env.init_play_renderer(render_spacing=render_spacing) - - obs = initialize() - last_render_time = time.perf_counter() - render_dt = 1.0 / 60.0 - steps_run = 0 - - while num_steps is None or steps_run < num_steps: - obs = step(obs) - current_time = time.perf_counter() - elapsed = current_time - last_render_time - if elapsed < render_dt: - time.sleep(render_dt - elapsed) - last_render_time = time.perf_counter() - env.render_play_frame() - steps_run += 1 - return None - - if num_steps is None: - raise ValueError("MuJoCo play rendering requires a finite num_steps value.") - if output_video is None: - raise ValueError("MuJoCo play rendering requires an output_video path.") - if frame_state_getter is None: - frame_state_getter = env.get_physics_state_snapshot - assert frame_state_getter is not None - - obs = initialize() - state_list = [] - for _ in range(num_steps): - obs = step(obs) - state_list.append(np.asarray(frame_state_getter(), dtype=np.float32).copy()) - - from unilab.utils import render_many - - cam_kw = dict(camera_kwargs or {}) - use_tracking = bool(cam_kw.pop("cam_tracking", False)) - tracking_env_idx = int(cam_kw.pop("cam_tracking_env_idx", 0)) - tracking_extra_envs = int(cam_kw.pop("cam_tracking_extra_envs", 2)) - effective_spacing = ( - float(render_spacing) - if render_spacing is not None - else float(getattr(env.cfg, "render_spacing", 1.0)) - ) - with tempfile.TemporaryDirectory(prefix="unilab-playback-models-") as tmp_dir: - model_files = _resolve_render_play_model_files( - env, - num_envs=state_list[0].shape[0], - tmp_dir=tmp_dir, - ) - - if use_tracking: - frames = render_many.render_states_get_frames_tracking( - state_list, - model_files, - width=1280, - height=720, - tracking_env_idx=tracking_env_idx, - max_extra_envs=tracking_extra_envs, - cam_distance=cam_kw.get("cam_distance", 2.0), - cam_elevation=cam_kw.get("cam_elevation", -20), - cam_azimuth=cam_kw.get("cam_azimuth", 90), - render_spacing=effective_spacing, - ) - else: - frames = render_many.render_states_get_frames( - state_list, - model_files, - width=1280, - height=720, - camera_id=-1, - render_spacing=effective_spacing, - **cam_kw, - ) - - import mediapy as media - - media.write_video(str(output_video), frames, fps=int(1.0 / env.cfg.ctrl_dt)) - return str(output_video) - - -def _resolve_render_play_model_files( - env: Any, - *, - num_envs: int, - tmp_dir: str | Path, -) -> str | list[str]: - """Resolve visual MuJoCo model files for offline play/video export. - - Args: - env: Environment exposing the playback-model contract. - num_envs: Number of envs present in the playback batch. - tmp_dir: Temporary directory used for serialized visual playback models. - - Returns: - A single visual model path when all envs share one model, otherwise a - per-env list of visual model file paths aligned with the env dimension. - """ - visual_model_file = str(env.cfg.model_file) - if not hasattr(env, "get_playback_model"): - return visual_model_file - - first_model = env.get_playback_model(0) - if isinstance(first_model, (str, Path)): - return str(first_model) - - import mujoco as _mujoco - - mujoco: Any = _mujoco - - visual_base = mujoco.MjModel.from_xml_path(visual_model_file) - tmp_root = Path(tmp_dir) - path_by_model_id: dict[int, str] = {} - model_files: list[str] = [] - for env_idx in range(num_envs): - playback_model = env.get_playback_model(env_idx) - if isinstance(playback_model, (str, Path)): - model_files.append(str(playback_model)) - continue - key = id(playback_model) - saved = path_by_model_id.get(key) - if saved is None: - saved = _materialize_visual_playback_model( - visual_model_file=visual_model_file, - visual_base_model=visual_base, - playback_model=playback_model, - output_path=tmp_root / f"model_{len(path_by_model_id)}.mjb", - ) - path_by_model_id[key] = saved - model_files.append(saved) - - if len(set(model_files)) == 1: - return model_files[0] - return model_files - - -def _materialize_visual_playback_model( - *, - visual_model_file: str, - visual_base_model: Any, - playback_model: Any, - output_path: str | Path, -) -> str: - """Compile a visual MuJoCo model using geom sizes from a playback model. - - Args: - visual_model_file: Source visual scene XML used for offline rendering. - visual_base_model: Compiled MuJoCo model from ``visual_model_file``. - playback_model: Backend playback model whose geom sizes should drive the - rendered geometry. - output_path: Destination ``.mjb`` path. - - Returns: - The saved ``.mjb`` path as a string. - """ - import mujoco as _mujoco - - mujoco: Any = _mujoco - - spec = mujoco.MjSpec.from_file(visual_model_file) - for geom_id in range(visual_base_model.ngeom): - geom_name = mujoco.mj_id2name(visual_base_model, mujoco.mjtObj.mjOBJ_GEOM, geom_id) - if not geom_name: - continue - playback_geom_id = mujoco.mj_name2id(playback_model, mujoco.mjtObj.mjOBJ_GEOM, geom_name) - if playback_geom_id < 0: - continue - geom = spec.geom(geom_name) - if geom is None: - continue - geom.size = list(np.asarray(playback_model.geom_size[playback_geom_id], dtype=np.float64)) - - visual_model = spec.compile() - output = Path(output_path) - mujoco.mj_saveModel(visual_model, str(output)) - return str(output) diff --git a/src/unilab/utils/experiment_tracking.py b/src/unilab/training/experiment.py similarity index 96% rename from src/unilab/utils/experiment_tracking.py rename to src/unilab/training/experiment.py index e40216dce..f809ca2dc 100644 --- a/src/unilab/utils/experiment_tracking.py +++ b/src/unilab/training/experiment.py @@ -5,6 +5,7 @@ import dataclasses import getpass import importlib +import importlib.util import json import os import socket @@ -58,7 +59,15 @@ def _json_safe(value: Any) -> Any: def get_device_info_dict() -> dict[str, str]: try: - module = importlib.import_module("benchmark.core.device_info") + module_path = Path(__file__).resolve().parents[4] / "benchmark" / "core" / "device_info.py" + spec = importlib.util.spec_from_file_location( + "unilab_benchmark_device_info", + module_path, + ) + if spec is None or spec.loader is None: + raise ImportError(f"Unable to load device info module from {module_path}") + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) getter = getattr(module, "get_device_info_dict") return dict(getter()) except Exception: diff --git a/src/unilab/utils/hardware_monitor.py b/src/unilab/training/monitoring.py similarity index 100% rename from src/unilab/utils/hardware_monitor.py rename to src/unilab/training/monitoring.py diff --git a/src/unilab/utils/reward_utils.py b/src/unilab/training/reward.py similarity index 100% rename from src/unilab/utils/reward_utils.py rename to src/unilab/training/reward.py diff --git a/src/unilab/training/rsl_rl.py b/src/unilab/training/rsl_rl.py new file mode 100644 index 000000000..2270fd500 --- /dev/null +++ b/src/unilab/training/rsl_rl.py @@ -0,0 +1,221 @@ +"""RSL-RL-specific training helpers.""" + +from __future__ import annotations + +from copy import deepcopy +from typing import Any + +import numpy as np +import torch +from tensordict import TensorDict + +from unilab.base.final_observation import resolve_terminal_observation_contract +from unilab.utils.tensor import to_numpy, to_torch + + +def get_policy_obs_dims(obs_groups_spec: dict[str, int]) -> tuple[int, int]: + """Return ``(actor_obs_dim, flat_policy_obs_dim)`` for RSL-RL policies.""" + actor_obs_dim = int(obs_groups_spec.get("obs", 0)) + flat_policy_obs_dim = int( + sum(dim for group_name, dim in obs_groups_spec.items() if group_name != "critic") + ) + return actor_obs_dim, flat_policy_obs_dim or actor_obs_dim + + +def normalize_ppo_train_cfg(train_cfg: dict[str, Any]) -> dict[str, Any]: + """Map UniLab PPO owner config to the current RSL-RL schema.""" + normalized = deepcopy(train_cfg) + algorithm_cfg = normalized.get("algorithm") + if isinstance(algorithm_cfg, dict): + for key in ( + "target_kl_stop", + "adaptive_kl_beta", + "adaptive_lr_growth", + "adaptive_lr_decay", + "adaptive_lr_update_interval", + "metrics_interval", + "finite_check_interval", + "enable_compile", + "warmup_strict_iters", + "warmup_metrics_interval", + "warmup_finite_check_interval", + "disable_finite_checks", + ): + algorithm_cfg.pop(key, None) + + if "actor" in normalized and "critic" in normalized: + return normalized + + policy_cfg = normalized.pop("policy", None) + if not isinstance(policy_cfg, dict): + return normalized + + actor_hidden_dims = policy_cfg.get("actor_hidden_dims", [512, 256, 128]) + critic_hidden_dims = policy_cfg.get("critic_hidden_dims", actor_hidden_dims) + activation = policy_cfg.get("activation", "elu") + init_noise_std = float(policy_cfg.get("init_noise_std", 1.0)) + + normalized["actor"] = { + "class_name": "rsl_rl.models.MLPModel", + "hidden_dims": actor_hidden_dims, + "activation": activation, + "distribution_cfg": { + "class_name": "rsl_rl.modules.distribution.GaussianDistribution", + "init_std": init_noise_std, + "std_type": "scalar", + }, + } + normalized["critic"] = { + "class_name": "rsl_rl.models.MLPModel", + "hidden_dims": critic_hidden_dims, + "activation": activation, + } + + obs_groups = normalized.get("obs_groups") + if isinstance(obs_groups, dict) and "actor" not in obs_groups and "default" in obs_groups: + default_groups = obs_groups.pop("default") + if isinstance(default_groups, list) and default_groups: + obs_groups["actor"] = list(default_groups) + + return normalized + + +class RslRlVecEnvWrapper: + """Adapter from UniLab's env contract to the RSL-RL VecEnv contract.""" + + def __init__( + self, + env: Any, + device: str = "cpu", + policy_obs_mode: str = "flat", + ) -> None: + if policy_obs_mode == "auto": + policy_obs_mode = "flat" + if policy_obs_mode not in {"actor", "flat"}: + raise ValueError( + f"Unsupported policy_obs_mode={policy_obs_mode!r}; expected 'actor' or 'flat'." + ) + + self.env = env + self.cfg = env.cfg + self.device = device + self.policy_obs_mode = policy_obs_mode + self.num_envs = env.num_envs + self.observation_space = env.observation_space + self.action_space = env.action_space + + self._actor_obs_dim, self._flat_obs_dim = get_policy_obs_dims(env.obs_groups_spec) + self.num_obs = ( + self._actor_obs_dim if self.policy_obs_mode == "actor" else self._flat_obs_dim + ) + self.num_privileged_obs = int(env.obs_groups_spec.get("critic", self.num_obs)) + action_shape = env.action_space.shape + if action_shape is None: + raise ValueError("env.action_space.shape must be defined") + self.num_actions = int(action_shape[0]) + + self.episode_returns = torch.zeros(self.num_envs, device=device) + self.episode_lengths = torch.zeros(self.num_envs, device=device) + self.episode_length_buf = self.episode_lengths + self.max_episode_length = np.ceil(env.cfg.max_episode_seconds / env.cfg.ctrl_dt) + self.reset() + + def _policy_obs(self, obs: dict[str, Any]) -> torch.Tensor: + if self.policy_obs_mode == "actor": + return to_torch(obs["obs"], self.device) + + policy_groups = [ + to_numpy(value) for group_name, value in obs.items() if group_name != "critic" + ] + if not policy_groups: + raise KeyError("Observation dict must contain at least one non-critic group") + if len(policy_groups) == 1: + return to_torch(policy_groups[0], self.device) + return to_torch(np.concatenate(policy_groups, axis=1), self.device) + + def _obs_to_tensordict(self, obs: dict[str, Any]) -> TensorDict: + actor_obs = to_torch(obs["obs"], self.device) + td_dict: dict[str, torch.Tensor] = { + "actor": actor_obs, + "policy": self._policy_obs(obs), + } + if "critic" in obs: + td_dict["critic"] = to_torch(obs["critic"], self.device) + return TensorDict(td_dict, batch_size=self.num_envs, device=self.device) + + def _resolve_done(self, state: Any) -> torch.Tensor: + if hasattr(state, "done"): + return to_torch(state.done, self.device).bool() + terminated = np.asarray(getattr(state, "terminated"), dtype=bool).ravel() + truncated = np.asarray(getattr(state, "truncated"), dtype=bool).ravel() + return to_torch(np.logical_or(terminated, truncated), self.device).bool() + + def _resolve_final_observation(self, state: Any) -> dict[str, Any] | None: + final_observation = getattr(state, "final_observation", None) + if isinstance(final_observation, dict): + return final_observation + info = getattr(state, "info", None) + if isinstance(info, dict): + final_observation = info.get("final_observation") + if isinstance(final_observation, dict): + return final_observation + return None + + def step( + self, actions: torch.Tensor | np.ndarray + ) -> tuple[TensorDict, torch.Tensor, torch.Tensor, dict]: + actions_np = to_numpy(actions) + state = self.env.step(actions_np) + rewards = to_torch(state.reward, self.device) + dones = self._resolve_done(state) + + self.episode_returns += rewards + self.episode_lengths += 1 + + infos: dict[str, torch.Tensor | TensorDict | dict[str, Any]] = {} + done_idx = torch.nonzero(dones).flatten() + if len(done_idx) > 0: + truncated = getattr(state, "truncated", None) + if truncated is not None: + infos["time_outs"] = to_torch(truncated, self.device).bool() + + final_observation = self._resolve_final_observation(state) + terminal_contract = resolve_terminal_observation_contract( + next_obs_batch_size=self.num_envs, + final_observation=final_observation, + done=to_numpy(dones), + info=getattr(state, "info", None), + truncated=to_numpy(infos["time_outs"]) if "time_outs" in infos else None, + ) + if np.any(terminal_contract.timeout_terminal_mask) and final_observation is not None: + infos["time_out_bootstrap_obs"] = self._obs_to_tensordict(final_observation) + + self.episode_returns[done_idx] = 0 + self.episode_lengths[done_idx] = 0 + + if hasattr(state, "info") and "log" in state.info: + infos["log"] = state.info["log"] + + return self._obs_to_tensordict(state.obs), rewards, dones, infos + + def reset(self) -> tuple[TensorDict, dict[str, Any]]: + if self.env.state is None: + self.env.init_state() + + env_indices = np.arange(self.num_envs, dtype=np.int32) + obs_out, info = self.env.reset(env_indices) + self.episode_returns[:] = 0 + self.episode_lengths[:] = 0 + return self._obs_to_tensordict(obs_out), info + + def get_observations(self) -> TensorDict: + assert self.env.state is not None + return self._obs_to_tensordict(self.env.state.obs) + + def get_privileged_observations(self) -> torch.Tensor: + assert self.env.state is not None + obs = self.env.state.obs + return to_torch(obs.get("critic", obs["obs"]), self.device) + + def close(self) -> None: + self.env.close() diff --git a/src/unilab/training/run.py b/src/unilab/training/run.py new file mode 100644 index 000000000..e624b4e20 --- /dev/null +++ b/src/unilab/training/run.py @@ -0,0 +1,166 @@ +"""Run directory and checkpoint resolution helpers.""" + +from __future__ import annotations + +from pathlib import Path + +from omegaconf import DictConfig, OmegaConf + + +def get_log_root(root_dir: str | Path, cfg: DictConfig) -> Path: + """Resolve the algorithm log root, honoring optional training.log_root overrides.""" + configured_root = OmegaConf.select(cfg, "training.log_root") + if configured_root: + log_root = Path(str(configured_root)) + return log_root if log_root.is_absolute() else Path(root_dir) / log_root + return Path(root_dir) / "logs" / str(OmegaConf.select(cfg, "algo.algo_log_name")) + + +def get_entrypoint_log_root( + root_dir: str | Path, + *, + algo_log_name: str, + log_root: str | Path | None = None, +) -> Path: + """Resolve the log root for non-Hydra entrypoints using training helper semantics.""" + if log_root is not None: + configured_root = Path(log_root) + return ( + configured_root if configured_root.is_absolute() else Path(root_dir) / configured_root + ) + return Path(root_dir) / "logs" / algo_log_name + + +def get_latest_run(log_dir: str | Path) -> Path | None: + """Return the lexicographically latest run directory under a task log root.""" + base_dir = Path(log_dir) + if not base_dir.exists(): + return None + runs = sorted(path for path in base_dir.iterdir() if path.is_dir()) + return runs[-1] if runs else None + + +def get_latest_checkpoint(run_dir: str | Path, *, suffix: str = ".pt") -> Path | None: + """Return the latest model checkpoint inside a run directory.""" + run_path = Path(run_dir) + if not run_path.exists(): + return None + + def _iteration(path: Path) -> int: + stem_parts = path.stem.split("_", 1) + if len(stem_parts) != 2: + return -1 + try: + return int(stem_parts[1]) + except ValueError: + return -1 + + model_files = [ + path + for path in run_path.iterdir() + if path.is_file() and path.name.startswith("model_") and path.suffix == suffix + ] + if not model_files: + return None + return max(model_files, key=_iteration) + + +def resolve_checkpoint_path( + base_log_dir: str | Path, + load_run: str, + *, + suffix: str = ".pt", +) -> tuple[Path | None, Path | None]: + """Resolve a latest or explicit checkpoint path from a task log root.""" + base_dir = Path(base_log_dir) + if load_run == "-1": + run_dir = get_latest_run(base_dir) + if run_dir is None: + return None, None + checkpoint = get_latest_checkpoint(run_dir, suffix=suffix) + return (checkpoint, run_dir) if checkpoint is not None else (None, None) + + candidate = Path(load_run) + if not candidate.exists(): + candidate = base_dir / load_run + if candidate.is_file(): + return candidate, candidate.parent + if candidate.is_dir(): + checkpoint = get_latest_checkpoint(candidate, suffix=suffix) + return (checkpoint, candidate) if checkpoint is not None else (None, None) + return None, None + + +def parse_checkpoint_path( + cfg: DictConfig, + *, + root_dir: str | Path, + load_run: str | None = None, + task_name: str | None = None, + checkpoint: str | int | None = None, + suffix: str = ".pt", +) -> tuple[Path | None, Path | None]: + """Resolve a checkpoint path from Hydra config and repository root.""" + selected_task = task_name or str(OmegaConf.select(cfg, "training.task_name")) + selected_run = load_run or str(OmegaConf.select(cfg, "algo.load_run", default="-1")) + selected_checkpoint = checkpoint + if selected_checkpoint is None: + selected_checkpoint = OmegaConf.select(cfg, "algo.checkpoint", default=-1) + if selected_checkpoint in (None, "", -1, "-1"): + selected_checkpoint = None + + return resolve_task_checkpoint_path( + root_dir, + task_name=selected_task, + load_run=selected_run, + algo_log_name=str(OmegaConf.select(cfg, "algo.algo_log_name")), + checkpoint=str(selected_checkpoint) if selected_checkpoint is not None else None, + suffix=suffix, + log_root=OmegaConf.select(cfg, "training.log_root"), + ) + + +def resolve_task_checkpoint_path( + root_dir: str | Path, + *, + task_name: str, + load_run: str, + algo_log_name: str, + checkpoint: str | None = None, + suffix: str = ".pt", + log_root: str | Path | None = None, +) -> tuple[Path | None, Path | None]: + """Resolve checkpoint paths for auxiliary entrypoints through shared training semantics.""" + task_log_root = ( + get_entrypoint_log_root( + root_dir, + algo_log_name=algo_log_name, + log_root=log_root, + ) + / task_name + ) + + run_dir: Path | None + if load_run == "-1": + run_dir = get_latest_run(task_log_root) + else: + candidate = Path(load_run) + if not candidate.exists(): + candidate = task_log_root / load_run + if candidate.is_file(): + return candidate, candidate.parent + run_dir = candidate if candidate.is_dir() else None + + if run_dir is None: + return None, None + + checkpoint_path: Path | None + if checkpoint is not None: + checkpoint_name = ( + f"model_{checkpoint}{suffix}" if str(checkpoint).isdigit() else str(checkpoint) + ) + checkpoint_path = run_dir / checkpoint_name + return (checkpoint_path, run_dir) if checkpoint_path.exists() else (None, run_dir) + + checkpoint_path = get_latest_checkpoint(run_dir, suffix=suffix) + return (checkpoint_path, run_dir) if checkpoint_path is not None else (None, run_dir) diff --git a/src/unilab/utils/__init__.py b/src/unilab/utils/__init__.py index 0eec1a8d7..987a4066b 100644 --- a/src/unilab/utils/__init__.py +++ b/src/unilab/utils/__init__.py @@ -1,16 +1,4 @@ -# Utility modules for UniLab -from unilab.utils.algo_utils import build_actor, ensure_registries -from unilab.utils.offpolicy_logger import OffPolicyLogger -from unilab.utils.onpolicy_logger import OnPolicyLogger -from unilab.utils.rsl_rl_vec_env_wrapper import RslRlVecEnvWrapper -from unilab.utils.torch_utils import to_numpy, to_torch +from unilab.utils.device import get_default_device +from unilab.utils.tensor import to_numpy, to_torch -__all__ = [ - "to_torch", - "to_numpy", - "OffPolicyLogger", - "OnPolicyLogger", - "ensure_registries", - "build_actor", - "RslRlVecEnvWrapper", -] +__all__ = ["get_default_device", "to_numpy", "to_torch"] diff --git a/src/unilab/utils/algo_utils.py b/src/unilab/utils/algo_utils.py deleted file mode 100644 index 767f02f3a..000000000 --- a/src/unilab/utils/algo_utils.py +++ /dev/null @@ -1,135 +0,0 @@ -"""Common utilities for RL algorithms.""" - -from __future__ import annotations - -import importlib -import logging -from typing import Sequence - -logger = logging.getLogger(__name__) - -# Attribute name for package-level registry bootstrap contracts. -_REGISTRY_MODULES_ATTR = "__unilab_registry_modules__" - -# Default packages to import for env registration bootstrap. -_DEFAULT_REGISTRY_PACKAGES = ( - "unilab.envs.locomotion", - "unilab.envs.manipulation", - "unilab.envs.motion_tracking", -) - - -def ensure_registries( - packages: Sequence[str] | None = None, - *, - optional_packages: Sequence[str] | None = None, - fail_on_error: bool = True, -) -> None: - """Import env registry bootstrap modules. - - Args: - packages: Package or module names to import for env registration. - Package-level registry modules should declare - ``__unilab_registry_modules__`` as explicit bootstrap targets. - Defaults to standard unilab env registry packages. - optional_packages: List of optional package names that may not be present. - Import failures for these are logged as warnings, not raised. - fail_on_error: If True (default), raise exceptions for non-optional packages. - If False, log warnings instead of raising. - - Raises: - ImportError: If a non-optional package fails to import and fail_on_error is True. - RuntimeError: If a declared registry module fails to import and fail_on_error is True. - TypeError: If ``__unilab_registry_modules__`` has an invalid format. - """ - pkgs = list(packages) if packages is not None else list(_DEFAULT_REGISTRY_PACKAGES) - optional = set(optional_packages) if optional_packages else set() - - for pkg_name in pkgs: - is_optional = pkg_name in optional - try: - package = importlib.import_module(pkg_name) - except ImportError as e: - if is_optional: - logger.warning("Optional registry package not found: %s (%s)", pkg_name, e) - elif fail_on_error: - raise ImportError( - f"Failed to import registry package '{pkg_name}'. " - f"Add to optional_packages if this is expected to be absent." - ) from e - else: - logger.warning("Registry package not found: %s (%s)", pkg_name, e) - continue - - modules = getattr(package, _REGISTRY_MODULES_ATTR, ()) - if isinstance(modules, str) or not isinstance(modules, Sequence): - raise TypeError( - f"'{pkg_name}.{_REGISTRY_MODULES_ATTR}' must be a sequence of module names." - ) - - for name in modules: - if not isinstance(name, str) or not name: - raise TypeError( - f"'{pkg_name}.{_REGISTRY_MODULES_ATTR}' entries must be non-empty strings." - ) - try: - importlib.import_module(name) - except Exception as e: - if fail_on_error and not is_optional: - raise RuntimeError( - f"Failed to import declared registry module '{name}' from '{pkg_name}'. " - f"Fix the import error or add '{pkg_name}' to optional_packages." - ) from e - logger.warning("Failed to import declared registry module '%s': %s", name, e) - - -def build_actor( - algo_type, - obs_dim, - action_dim, - actor_hidden_dim, - use_layer_norm, - device, - num_envs=1, - actor_num_blocks: int = 2, - actor_noise_zeta_mu: float = 2.0, - actor_noise_zeta_max: int = 16, -): - """Build the correct actor model based on algorithm type.""" - if algo_type == "sac": - from unilab.algos.torch.fast_sac.learner import SACActor - - return SACActor( - obs_dim=obs_dim, - action_dim=action_dim, - hidden_dim=actor_hidden_dim, - use_layer_norm=use_layer_norm, - device=device, - ) - elif algo_type == "td3": - from unilab.algos.torch.fast_td3.learner import TD3Actor - - return TD3Actor( - obs_dim=obs_dim, - n_act=action_dim, - num_envs=num_envs, - hidden_dim=actor_hidden_dim, - init_scale=0.01, - log_std_min=-0.9, - log_std_max=0.0, - device=device, - ) - elif algo_type == "flashsac": - from unilab.algos.torch.flash_sac.network import FlashSACActor - - return FlashSACActor( - num_blocks=actor_num_blocks, - input_dim=obs_dim, - hidden_dim=actor_hidden_dim, - action_dim=action_dim, - noise_zeta_mu=actor_noise_zeta_mu, - noise_zeta_max=actor_noise_zeta_max, - device=device, - ) - else: - raise ValueError(f"Unknown algo_type: {algo_type}") diff --git a/src/unilab/utils/device.py b/src/unilab/utils/device.py new file mode 100644 index 000000000..1118429bd --- /dev/null +++ b/src/unilab/utils/device.py @@ -0,0 +1,12 @@ +from __future__ import annotations + +import torch + + +def get_default_device() -> str: + """Detect the best available device.""" + if torch.cuda.is_available(): + return "cuda" + if torch.backends.mps.is_available(): + return "mps" + return "cpu" diff --git a/src/unilab/utils/rsl_rl_compat.py b/src/unilab/utils/rsl_rl_compat.py deleted file mode 100644 index fb17346e8..000000000 --- a/src/unilab/utils/rsl_rl_compat.py +++ /dev/null @@ -1,230 +0,0 @@ -""" -Compatibility utilities for supporting both rsl_rl 3.x and 4.x. - -rsl_rl 4.0 introduced breaking API changes: - - Config format: single `policy` dict → separate `actor`/`critic` dicts - - `construct_algorithm` moved from Runner to Algorithm (PPO) - - `rnd_cfg` must exist in algorithm config (can be None) - - `empirical_normalization` deprecated - -This module provides runtime version detection and config conversion so that -the codebase can work with both versions without code duplication. -""" - -import importlib.metadata -from copy import deepcopy - -from packaging.version import Version - - -def get_rsl_rl_version() -> str: - """Get the installed rsl_rl version string.""" - try: - return importlib.metadata.version("rsl-rl-lib") - except importlib.metadata.PackageNotFoundError: - # Fallback: try the old package name - try: - return importlib.metadata.version("rsl-rl") - except importlib.metadata.PackageNotFoundError: - raise ImportError("rsl_rl is not installed. Install via: pip install rsl-rl-lib") - - -def is_rsl_rl_v4() -> bool: - """Check if the installed rsl_rl version is 4.x or above.""" - version_str = get_rsl_rl_version() - print(version_str) - return bool(Version(version_str) >= Version("4.0.0")) - - -def is_rsl_rl_v5() -> bool: - """Check if the installed rsl_rl version is 4.x or above.""" - version_str = get_rsl_rl_version() - print(version_str) - return bool(Version(version_str) >= Version("5.0.0")) - - -def convert_config_v3_to_v4(cfg: dict) -> dict: - """Convert rsl_rl 3.x config format to 4.x format. - - 3.x uses a single `policy` dict with class_name="ActorCritic": - policy: - class_name: "ActorCritic" - actor_hidden_dims: [...] - critic_hidden_dims: [...] - activation: "elu" - init_noise_std: 1.0 - - 4.x uses separate `actor` and `critic` dicts with class_name="MLPModel": - actor: - class_name: "MLPModel" - hidden_dims: [...] - activation: "elu" - init_noise_std: 1.0 - critic: - class_name: "MLPModel" - hidden_dims: [...] - activation: "elu" - """ - cfg = deepcopy(cfg) - - # Remove deprecated fields, but capture value for migration first - empirical_normalization = cfg.pop("empirical_normalization", False) - cfg.pop("runner_class_name", None) - - # Convert policy → actor + critic - if "policy" in cfg: - policy = cfg.pop("policy") - cfg["actor"] = { - "class_name": "MLPModel", - "hidden_dims": policy.get("actor_hidden_dims", [256, 256, 256]), - "activation": policy.get("activation", "elu"), - "obs_normalization": empirical_normalization, - "distribution_cfg": { - "class_name": "rsl_rl.modules.distribution.GaussianDistribution", - "init_std": policy.get("init_noise_std", 1.0), - "std_type": policy.get("noise_std_type", "scalar"), - }, - } - cfg["critic"] = { - "class_name": "MLPModel", - "hidden_dims": policy.get("critic_hidden_dims", [256, 256, 256]), - "activation": policy.get("activation", "elu"), - "obs_normalization": empirical_normalization, - } - - # 4.x requires rnd_cfg in algorithm config (can be None) - if "algorithm" in cfg: - cfg["algorithm"].setdefault("rnd_cfg", None) - # Remove class_name for algorithm - it's popped by construct_algorithm - # but needs to stay for the runner to resolve it, so leave it alone - - # Strip mlx_ppo-only parameters that rsl_rl PPO does not accept - _MLX_PPO_ONLY_KEYS = { - "adaptive_kl_beta", - "adaptive_lr_decay", - "adaptive_lr_growth", - "adaptive_lr_update_interval", - "target_kl_stop", - "metrics_interval", - "finite_check_interval", - "enable_compile", - "warmup_strict_iters", - "warmup_metrics_interval", - "warmup_finite_check_interval", - "disable_finite_checks", - } - for key in _MLX_PPO_ONLY_KEYS: - cfg["algorithm"].pop(key, None) - - # 4.x requires obs_groups - obs_groups = cfg.get("obs_groups", {}) - if "default" in obs_groups: - if "actor" not in obs_groups: - obs_groups["actor"] = obs_groups["default"] - if "critic" not in obs_groups: - obs_groups["critic"] = obs_groups["default"] - else: - # Fallback if no groups defined at all (from V3 policy) - if "actor" not in obs_groups: - obs_groups["actor"] = ["policy"] - if "critic" not in obs_groups: - obs_groups["critic"] = ["policy"] - - cfg["obs_groups"] = obs_groups - - # 4.x requires multi_gpu config - if "multi_gpu" not in cfg: - cfg["multi_gpu"] = None - - return cfg - - -def convert_config_v5(cfg: dict) -> dict: - """Convert config format to 5.x format. - - - 5.x uses separate `actor` and `critic` dicts with class_name="MLPModel": - actor: - class_name: "MLPModel" - hidden_dims: [...] - activation: "elu" - distribution_cfg: { 'class_name': - 'init_std': - 'std_type': - } - critic: - class_name: "MLPModel" - hidden_dims: [...] - activation: "elu" - """ - cfg = deepcopy(cfg) - - # Remove deprecated fields, but capture value for migration first - empirical_normalization = cfg.pop("empirical_normalization", False) - cfg.pop("runner_class_name", None) - - # Convert policy → actor + critic - if "policy" in cfg: - policy = cfg.pop("policy") - cfg["actor"] = { - "class_name": "MLPModel", - "hidden_dims": policy.get("actor_hidden_dims", [256, 256, 256]), - "activation": policy.get("activation", "elu"), - # "init_noise_std": policy.get("init_noise_std", 1.0), - # "noise_std_type": policy.get("noise_std_type", "scalar"), - # "stochastic": True, # Required: MLPModel needs this to create output distribution - "obs_normalization": empirical_normalization, - "distribution_cfg": { - "class_name": "GaussianDistribution", - "init_std": policy.get("init_noise_std", 1.0), - "std_type": policy.get("noise_std_type", "scalar"), - }, - } - cfg["critic"] = { - "class_name": "MLPModel", - "hidden_dims": policy.get("critic_hidden_dims", [256, 256, 256]), - "activation": policy.get("activation", "elu"), - "obs_normalization": empirical_normalization, - } - - # 4.x requires rnd_cfg in algorithm config (can be None) - if "algorithm" in cfg: - cfg["algorithm"].setdefault("rnd_cfg", None) - # Remove class_name for algorithm - it's popped by construct_algorithm - # but needs to stay for the runner to resolve it, so leave it alone - - # Strip mlx_ppo-only parameters that rsl_rl PPO does not accept - _MLX_PPO_ONLY_KEYS = { - "adaptive_kl_beta", - "adaptive_lr_decay", - "adaptive_lr_growth", - "adaptive_lr_update_interval", - "target_kl_stop", - "metrics_interval", - "finite_check_interval", - "enable_compile", - "warmup_strict_iters", - "warmup_metrics_interval", - "warmup_finite_check_interval", - "disable_finite_checks", - } - for key in _MLX_PPO_ONLY_KEYS: - cfg["algorithm"].pop(key, None) - - # 4.x requires obs_groups - obs_groups = cfg.get("obs_groups", {}) - if "default" in obs_groups: - if "actor" not in obs_groups: - obs_groups["actor"] = obs_groups["default"] - if "critic" not in obs_groups: - obs_groups["critic"] = obs_groups["default"] - else: - # Fallback if no groups defined at all (from V3 policy) - if "actor" not in obs_groups: - obs_groups["actor"] = ["policy"] - if "critic" not in obs_groups: - obs_groups["critic"] = ["policy"] - - cfg["obs_groups"] = obs_groups - - return cfg diff --git a/src/unilab/utils/rsl_rl_vec_env_wrapper.py b/src/unilab/utils/rsl_rl_vec_env_wrapper.py deleted file mode 100644 index 6748a6cd0..000000000 --- a/src/unilab/utils/rsl_rl_vec_env_wrapper.py +++ /dev/null @@ -1,166 +0,0 @@ -"""Shared RSL-RL vectorized environment wrapper. - -This module provides a unified RslRlVecEnvWrapper that aligns with the current -env contract (obs, info) reset format and is used by both training and play scripts. -""" - -import numpy as np -import torch -from tensordict import TensorDict - -from unilab.utils.obs_utils import flatten_policy_obs_dict -from unilab.utils.torch_utils import to_torch - - -class RslRlVecEnvWrapper: - """Wrapper to adapt NpEnv to RSL-RL OnPolicyRunner interface. - - This wrapper aligns with the current env contract: - - reset() returns (obs_dict, info_dict) - - step() returns state with .obs, .reward, .done, .truncated, .info attributes - - Args: - env: The environment to wrap (must follow NpEnv contract). - device: Device to place tensors on ("cuda", "mps", or "cpu"). - policy_obs_mode: Observation mode for policy ("flat" or "actor"). - "flat" uses flattened obs dict, "actor" uses only the "obs" key. - Default is "flat" for backward compatibility with training scripts. - """ - - def __init__(self, env, device: str = "cuda", policy_obs_mode: str = "flat"): - self.env = env - self.cfg = env.cfg - self.device = device - self.policy_obs_mode = policy_obs_mode - self.num_envs = env.num_envs - self.observation_space = env.observation_space - self.action_space = env.action_space - - # Compute observation dimensions - self._actor_obs_dim = int(env.obs_groups_spec.get("obs", sum(env.obs_groups_spec.values()))) - self._critic_obs_dim = int(env.obs_groups_spec.get("critic", self._actor_obs_dim)) - self._flat_obs_dim = self._actor_obs_dim - self.num_obs = self._flat_obs_dim if policy_obs_mode == "flat" else self._actor_obs_dim - # Legacy RSL-RL field name; semantically this is the critic-path observation dim. - self.num_privileged_obs = self._critic_obs_dim - self.num_actions = env.action_space.shape[0] - - # Episode tracking - self.episode_returns = torch.zeros(self.num_envs, device=device) - self.episode_lengths = torch.zeros(self.num_envs, device=device) - self.episode_length_buf = self.episode_lengths - self.max_episode_length = int(env.cfg.max_episode_seconds / env.cfg.ctrl_dt) - - # Initialize - self.reset() - - def _obs_to_tensordict(self, obs: dict[str, np.ndarray]) -> TensorDict: - """Convert observation dict to TensorDict for RSL-RL. - - Args: - obs: Observation dictionary with mandatory "obs" and optional "critic". - - Returns: - TensorDict with "policy" and "actor" keys, plus optional "critic". - """ - actor = to_torch(obs["obs"], self.device) - - if self.policy_obs_mode == "actor": - policy = actor - else: - policy = to_torch(flatten_policy_obs_dict(obs), self.device) - - td: dict[str, torch.Tensor] = {"policy": policy, "actor": actor} - - if "critic" in obs: - td["critic"] = to_torch(obs["critic"], self.device) - - return TensorDict(td, batch_size=self.num_envs, device=self.device) - - def step(self, actions): - """Execute one step in the environment. - - Args: - actions: Actions to execute (torch.Tensor or numpy array). - - Returns: - Tuple of (obs_tensordict, rewards, dones, infos). - """ - # Convert actions to numpy - if isinstance(actions, torch.Tensor): - actions_np = actions.detach().cpu().numpy() - else: - actions_np = actions - - # Step the environment - state = self.env.step(actions_np) - - # Convert outputs to torch tensors - rewards = to_torch(state.reward, self.device) - dones = to_torch(state.done, self.device).bool() - - # Update episode statistics - self.episode_returns += rewards - self.episode_lengths += 1 - - # Build info dict - infos = {} - done_indices = torch.nonzero(dones).flatten() - if len(done_indices) > 0: - if hasattr(state, "truncated"): - infos["time_outs"] = to_torch(state.truncated, self.device).bool() - final_observation = getattr(state, "final_observation", None) - if ( - final_observation is None - and hasattr(state, "info") - and isinstance(state.info, dict) - ): - final_observation = state.info.get("final_observation") - if torch.any(infos["time_outs"]) and isinstance(final_observation, dict): - infos["time_out_bootstrap_obs"] = self._obs_to_tensordict(final_observation) - self.episode_returns[done_indices] = 0 - self.episode_lengths[done_indices] = 0 - - if hasattr(state, "info") and "log" in state.info: - infos["log"] = state.info["log"] - - obs_dict = self._obs_to_tensordict(state.obs) - return obs_dict, rewards, dones, infos - - def reset(self): - """Reset the environment. - - Returns: - Tuple of (obs_tensordict, empty_info_dict). - """ - # Ensure state is initialized - if self.env.state is None: - self.env.init_state() - - # Reset all environments - env_indices = np.arange(self.num_envs, dtype=np.int32) - obs_out, _ = self.env.reset(env_indices) - - # Reset episode statistics - self.episode_returns[:] = 0 - self.episode_lengths[:] = 0 - - return self._obs_to_tensordict(obs_out), {} - - def get_observations(self): - """Get current observations without stepping. - - Returns: - TensorDict with current observations. - """ - return self._obs_to_tensordict(self.env.state.obs) - - def get_privileged_observations(self): - """Get current critic observations via the legacy RSL-RL hook name. - - Returns: - Torch tensor with critic-path observations (or actor obs if unavailable). - """ - obs = self.env.state.obs - critic_base = obs.get("critic", obs["obs"]) - return to_torch(critic_base, self.device) diff --git a/src/unilab/utils/run_utils.py b/src/unilab/utils/run_utils.py deleted file mode 100644 index 34817f2bc..000000000 --- a/src/unilab/utils/run_utils.py +++ /dev/null @@ -1,35 +0,0 @@ -import os - - -def get_latest_run(log_dir: str) -> str | None: - """Find the latest run in the log directory that contains a model. - - Args: - log_dir: Path to the base log directory (e.g., logs/fast_sac_Go2LocoFlatTerrain) - - Returns: - Path to the latest run directory containing a model, or None if none found. - """ - if not os.path.exists(log_dir): - return None - runs = sorted( - [ - d - for d in os.listdir(log_dir) - if os.path.isdir(os.path.join(log_dir, d)) - and d != "git" - and d[0].isdigit() # skip non-timestamp dirs (e.g. "appo-...", "play_temp") - ] - ) - - # Iterate backwards to find first run with models - for run_id in reversed(runs): - run_path = os.path.join(log_dir, run_id) - # Check if any .pt file exists - if any(f.endswith(".pt") for f in os.listdir(run_path)): - return run_path - - # Fallback to just the latest directory if no models found - if len(runs) > 0: - return os.path.join(log_dir, runs[-1]) - return None diff --git a/src/unilab/utils/torch_utils.py b/src/unilab/utils/tensor.py similarity index 100% rename from src/unilab/utils/torch_utils.py rename to src/unilab/utils/tensor.py diff --git a/src/unilab/visualization/__init__.py b/src/unilab/visualization/__init__.py new file mode 100644 index 000000000..7e93c76ce --- /dev/null +++ b/src/unilab/visualization/__init__.py @@ -0,0 +1,5 @@ +"""Visualization and playback helpers.""" + +from unilab.visualization.playback import render_play_mode + +__all__ = ["render_play_mode"] diff --git a/src/unilab/visualization/playback.py b/src/unilab/visualization/playback.py new file mode 100644 index 000000000..5b09281ec --- /dev/null +++ b/src/unilab/visualization/playback.py @@ -0,0 +1,182 @@ +"""Playback rendering helpers for interactive and offline visualization.""" + +from __future__ import annotations + +import tempfile +import time +from pathlib import Path +from typing import Any, Callable, TypeVar + +import numpy as np + +ObsT = TypeVar("ObsT") + + +def render_play_mode( + env, + *, + sim_backend: str, + initialize: Callable[[], ObsT], + step: Callable[[ObsT], ObsT], + num_steps: int | None, + output_video: str | Path | None = None, + render_spacing: float | None = None, + frame_state_getter: Callable[[], np.ndarray] | None = None, + camera_kwargs: dict[str, Any] | None = None, +) -> str | None: + """Render interactive Motrix play or MuJoCo video generation through shared callbacks.""" + if sim_backend == "motrix": + env.init_play_renderer(render_spacing=render_spacing) + + obs = initialize() + last_render_time = time.perf_counter() + render_dt = 1.0 / 60.0 + steps_run = 0 + + while num_steps is None or steps_run < num_steps: + obs = step(obs) + current_time = time.perf_counter() + elapsed = current_time - last_render_time + if elapsed < render_dt: + time.sleep(render_dt - elapsed) + last_render_time = time.perf_counter() + env.render_play_frame() + steps_run += 1 + return None + + if num_steps is None: + raise ValueError("MuJoCo play rendering requires a finite num_steps value.") + if output_video is None: + raise ValueError("MuJoCo play rendering requires an output_video path.") + if frame_state_getter is None: + frame_state_getter = env.get_physics_state_snapshot + assert frame_state_getter is not None + + obs = initialize() + state_list = [] + for _ in range(num_steps): + obs = step(obs) + state_list.append(np.asarray(frame_state_getter(), dtype=np.float32).copy()) + + from unilab.visualization import render_many + + cam_kw = dict(camera_kwargs or {}) + use_tracking = bool(cam_kw.pop("cam_tracking", False)) + tracking_env_idx = int(cam_kw.pop("cam_tracking_env_idx", 0)) + tracking_extra_envs = int(cam_kw.pop("cam_tracking_extra_envs", 2)) + effective_spacing = ( + float(render_spacing) + if render_spacing is not None + else float(getattr(env.cfg, "render_spacing", 1.0)) + ) + with tempfile.TemporaryDirectory(prefix="unilab-playback-models-") as tmp_dir: + model_files = _resolve_render_play_model_files( + env, + num_envs=state_list[0].shape[0], + tmp_dir=tmp_dir, + ) + + if use_tracking: + frames = render_many.render_states_get_frames_tracking( + state_list, + model_files, + width=1280, + height=720, + tracking_env_idx=tracking_env_idx, + max_extra_envs=tracking_extra_envs, + cam_distance=cam_kw.get("cam_distance", 2.0), + cam_elevation=cam_kw.get("cam_elevation", -20), + cam_azimuth=cam_kw.get("cam_azimuth", 90), + render_spacing=effective_spacing, + ) + else: + frames = render_many.render_states_get_frames( + state_list, + model_files, + width=1280, + height=720, + camera_id=-1, + render_spacing=effective_spacing, + **cam_kw, + ) + + import mediapy as media + + media.write_video(str(output_video), frames, fps=int(1.0 / env.cfg.ctrl_dt)) + return str(output_video) + + +def _resolve_render_play_model_files( + env: Any, + *, + num_envs: int, + tmp_dir: str | Path, +) -> str | list[str]: + """Resolve visual MuJoCo model files for offline play/video export.""" + visual_model_file = str(env.cfg.model_file) + if not hasattr(env, "get_playback_model"): + return visual_model_file + + first_model = env.get_playback_model(0) + if isinstance(first_model, (str, Path)): + return str(first_model) + + import mujoco as _mujoco + + mujoco: Any = _mujoco + + visual_base = mujoco.MjModel.from_xml_path(visual_model_file) + tmp_root = Path(tmp_dir) + path_by_model_id: dict[int, str] = {} + model_files: list[str] = [] + for env_idx in range(num_envs): + playback_model = env.get_playback_model(env_idx) + if isinstance(playback_model, (str, Path)): + model_files.append(str(playback_model)) + continue + key = id(playback_model) + saved = path_by_model_id.get(key) + if saved is None: + saved = _materialize_visual_playback_model( + visual_model_file=visual_model_file, + visual_base_model=visual_base, + playback_model=playback_model, + output_path=tmp_root / f"model_{len(path_by_model_id)}.mjb", + ) + path_by_model_id[key] = saved + model_files.append(saved) + + if len(set(model_files)) == 1: + return model_files[0] + return model_files + + +def _materialize_visual_playback_model( + *, + visual_model_file: str, + visual_base_model: Any, + playback_model: Any, + output_path: str | Path, +) -> str: + """Compile a visual MuJoCo model using geom sizes from a playback model.""" + import mujoco as _mujoco + + mujoco: Any = _mujoco + + spec = mujoco.MjSpec.from_file(visual_model_file) + for geom_id in range(visual_base_model.ngeom): + geom_name = mujoco.mj_id2name(visual_base_model, mujoco.mjtObj.mjOBJ_GEOM, geom_id) + if not geom_name: + continue + playback_geom_id = mujoco.mj_name2id(playback_model, mujoco.mjtObj.mjOBJ_GEOM, geom_name) + if playback_geom_id < 0: + continue + geom = spec.geom(geom_name) + if geom is None: + continue + geom.size = list(np.asarray(playback_model.geom_size[playback_geom_id], dtype=np.float64)) + + visual_model = spec.compile() + output = Path(output_path) + mujoco.mj_saveModel(visual_model, str(output)) + return str(output) diff --git a/src/unilab/utils/render_many.py b/src/unilab/visualization/render_many.py similarity index 100% rename from src/unilab/utils/render_many.py rename to src/unilab/visualization/render_many.py diff --git a/src/unilab/utils/viser_scene.py b/src/unilab/visualization/viser_scene.py similarity index 99% rename from src/unilab/utils/viser_scene.py rename to src/unilab/visualization/viser_scene.py index f6fbaf858..fdb64da61 100644 --- a/src/unilab/utils/viser_scene.py +++ b/src/unilab/visualization/viser_scene.py @@ -7,7 +7,7 @@ Usage (from ``scripts/play_viser.py``):: - from unilab.utils.viser_scene import MujocoViserScene, VISER_AVAILABLE + from unilab.visualization.viser_scene import MujocoViserScene, VISER_AVAILABLE """ from __future__ import annotations diff --git a/tests/algos/test_appo_runner.py b/tests/algos/test_appo_runner.py index e3e0d0c25..ca35227d0 100644 --- a/tests/algos/test_appo_runner.py +++ b/tests/algos/test_appo_runner.py @@ -13,7 +13,7 @@ pytest.importorskip("mujoco") from unilab.algos.torch.appo.runner import APPORunner -from unilab.config.structured_configs import APPOConfig +from unilab.structured_configs import APPOConfig @pytest.mark.slow diff --git a/tests/algos/test_fast_sac_symmetry_contract.py b/tests/algos/test_fast_sac_symmetry_contract.py index 0dfb1ad76..b757a4661 100644 --- a/tests/algos/test_fast_sac_symmetry_contract.py +++ b/tests/algos/test_fast_sac_symmetry_contract.py @@ -45,12 +45,11 @@ def close(self): def test_fast_sac_runner_uses_env_owned_symmetry_contract(monkeypatch: pytest.MonkeyPatch): from unilab.algos.torch.fast_sac.runner import FastSACRunner from unilab.base import registry - from unilab.utils import algo_utils augmentation = _FakeSymmetryAugmentation() fake_env = _FakeEnv(augmentation) - monkeypatch.setattr(algo_utils, "ensure_registries", lambda: None) + monkeypatch.setattr(registry, "ensure_registries", lambda: None) monkeypatch.setattr(registry, "make", lambda *args, **kwargs: fake_env) runner = FastSACRunner( @@ -75,7 +74,6 @@ def test_fast_sac_runner_uses_env_owned_symmetry_contract(monkeypatch: pytest.Mo def test_fast_sac_runner_skips_symmetry_builder_when_disabled(monkeypatch: pytest.MonkeyPatch): from unilab.algos.torch.fast_sac.runner import FastSACRunner from unilab.base import registry - from unilab.utils import algo_utils fake_env = _FakeEnv(_FakeSymmetryAugmentation()) @@ -84,7 +82,7 @@ def _unexpected_builder(*args, **kwargs): fake_env.build_symmetry_augmentation = _unexpected_builder # type: ignore[method-assign] - monkeypatch.setattr(algo_utils, "ensure_registries", lambda: None) + monkeypatch.setattr(registry, "ensure_registries", lambda: None) monkeypatch.setattr(registry, "make", lambda *args, **kwargs: fake_env) runner = FastSACRunner( diff --git a/tests/algos/test_mlx_ppo.py b/tests/algos/test_mlx_ppo.py index fa1c48dda..1bbfe19a8 100644 --- a/tests/algos/test_mlx_ppo.py +++ b/tests/algos/test_mlx_ppo.py @@ -328,9 +328,9 @@ def test_mlx_ppo_one_iteration_real_env(default_go2_reward_config): _mujoco = pytest.importorskip("mujoco") from unilab.base import registry - from unilab.config.structured_configs import PPOConfig as PPOStructuredConfig - from unilab.utils.algo_utils import ensure_registries - from unilab.utils.obs_utils import flatten_obs_dict + from unilab.base.observations import flatten_obs_dict + from unilab.base.registry import ensure_registries + from unilab.structured_configs import PPOConfig as PPOStructuredConfig ensure_registries() diff --git a/tests/algos/test_offpolicy_runner.py b/tests/algos/test_offpolicy_runner.py index ebaadf6e1..a1a63e8d8 100644 --- a/tests/algos/test_offpolicy_runner.py +++ b/tests/algos/test_offpolicy_runner.py @@ -15,7 +15,7 @@ from unilab.algos.torch.fast_sac.learner import FastSACLearner from unilab.algos.torch.fast_td3.learner import FastTD3Learner from unilab.algos.torch.offpolicy.runner import OffPolicyRunner -from unilab.config.structured_configs import SACConfig +from unilab.structured_configs import SACConfig def _make_sac_runner(env_name: str) -> OffPolicyRunner: diff --git a/tests/algos/test_rsl_rl_runner.py b/tests/algos/test_rsl_rl_runner.py index 30d966d19..28f9873d0 100644 --- a/tests/algos/test_rsl_rl_runner.py +++ b/tests/algos/test_rsl_rl_runner.py @@ -20,10 +20,10 @@ from tensordict import TensorDict from unilab.base import registry -from unilab.config.structured_configs import PPOConfig -from unilab.utils.algo_utils import ensure_registries -from unilab.utils.rsl_rl_compat import convert_config_v5, is_rsl_rl_v5 -from unilab.utils.torch_utils import to_torch +from unilab.base.registry import ensure_registries +from unilab.structured_configs import PPOConfig +from unilab.training.rsl_rl import normalize_ppo_train_cfg +from unilab.utils.tensor import to_torch ensure_registries() @@ -154,8 +154,7 @@ def test_rsl_rl_ppo_one_iteration( } train_cfg["algorithm"]["num_learning_epochs"] = 1 train_cfg["algorithm"]["num_mini_batches"] = 2 - if is_rsl_rl_v5(): - train_cfg = convert_config_v5(train_cfg) + train_cfg = normalize_ppo_train_cfg(train_cfg) with tempfile.TemporaryDirectory() as tmpdir: runner = OnPolicyRunner(cast(Any, wrapped), train_cfg, log_dir=tmpdir, device="cpu") diff --git a/tests/base/test_reward_override.py b/tests/base/test_reward_override.py index 3ad5a7342..cd9b17efa 100644 --- a/tests/base/test_reward_override.py +++ b/tests/base/test_reward_override.py @@ -5,7 +5,7 @@ import pytest from unilab.base import registry -from unilab.utils.algo_utils import ensure_registries +from unilab.base.registry import ensure_registries def test_reward_override_go1(): diff --git a/tests/base/test_sim_backend.py b/tests/base/test_sim_backend.py index 2bef1fe7e..50c3acf50 100644 --- a/tests/base/test_sim_backend.py +++ b/tests/base/test_sim_backend.py @@ -13,6 +13,7 @@ import pytest from unilab.assets import ASSETS_ROOT_PATH +from unilab.base.backend.xml import get_named_body_ids from unilab.dr import ( GeomSizeOverride, InitRandomizationPlan, @@ -20,7 +21,6 @@ ModelVariantSpec, ResetRandomizationPayload, ) -from unilab.utils.xml_utils import get_named_body_ids # --------------------------------------------------------------------------- diff --git a/tests/config/test_locomotion_params.py b/tests/config/test_locomotion_params.py index c27505c85..d6566514d 100644 --- a/tests/config/test_locomotion_params.py +++ b/tests/config/test_locomotion_params.py @@ -16,7 +16,7 @@ def test_sac_config_defaults(): - from unilab.config.structured_configs import SACAlgoParams, SACConfig + from unilab.structured_configs import SACAlgoParams, SACConfig cfg = SACConfig() assert cfg.algo == "sac" @@ -28,7 +28,7 @@ def test_sac_config_defaults(): def test_td3_config_defaults(): - from unilab.config.structured_configs import TD3Config + from unilab.structured_configs import TD3Config cfg = TD3Config() assert cfg.algo == "td3" @@ -38,7 +38,7 @@ def test_td3_config_defaults(): def test_flashsac_config_defaults(): - from unilab.config.structured_configs import FlashSACAlgoParams, FlashSACConfig + from unilab.structured_configs import FlashSACAlgoParams, FlashSACConfig cfg = FlashSACConfig() assert cfg.algo == "flashsac" @@ -53,7 +53,7 @@ def test_flashsac_config_defaults(): def test_ppo_config_defaults(): - from unilab.config.structured_configs import PPOConfig + from unilab.structured_configs import PPOConfig cfg = PPOConfig() assert cfg.algo == "ppo" @@ -64,7 +64,7 @@ def test_ppo_config_defaults(): def test_appo_config_defaults(): - from unilab.config.structured_configs import APPOConfig + from unilab.structured_configs import APPOConfig cfg = APPOConfig() assert cfg.algo == "appo" @@ -73,7 +73,7 @@ def test_appo_config_defaults(): def test_base_config_to_dict(): - from unilab.config.structured_configs import SACConfig + from unilab.structured_configs import SACConfig cfg = SACConfig() d = cfg.to_dict() diff --git a/tests/config/test_reward_injection.py b/tests/config/test_reward_injection.py index 0135528c6..397093136 100644 --- a/tests/config/test_reward_injection.py +++ b/tests/config/test_reward_injection.py @@ -28,7 +28,7 @@ def test_reward_config_loading_g1_motrix(): def test_resolve_reward_dict_reads_task_reward(): """Task-backend configs should expose the final reward mapping directly.""" - from unilab.utils.reward_utils import resolve_reward_dict + from unilab.training.reward import resolve_reward_dict with initialize(config_path="../../conf/ppo", version_base="1.3"): cfg = compose( @@ -45,7 +45,7 @@ def test_resolve_reward_dict_reads_task_reward(): def test_reward_config_conversion(): """Test reward config converts to dataclasses via registry.""" from unilab.base import registry - from unilab.utils.algo_utils import ensure_registries + from unilab.base.registry import ensure_registries ensure_registries() diff --git a/tests/envs/locomotion/g1/test_issue175_regression.py b/tests/envs/locomotion/g1/test_issue175_regression.py index 29a8d7324..fb8263f2b 100644 --- a/tests/envs/locomotion/g1/test_issue175_regression.py +++ b/tests/envs/locomotion/g1/test_issue175_regression.py @@ -10,8 +10,8 @@ from omegaconf import OmegaConf from unilab.base import registry +from unilab.base.registry import ensure_registries from unilab.training.backend_adapter import BackendAdapter -from unilab.utils.algo_utils import ensure_registries ROOT_DIR = Path(__file__).parents[4] CONF_DIR = ROOT_DIR / "conf" diff --git a/tests/envs/locomotion/g1/test_symmetry_contract.py b/tests/envs/locomotion/g1/test_symmetry_contract.py index 857a95d29..b0666d917 100644 --- a/tests/envs/locomotion/g1/test_symmetry_contract.py +++ b/tests/envs/locomotion/g1/test_symmetry_contract.py @@ -6,8 +6,8 @@ import torch from unilab.base import registry +from unilab.base.registry import ensure_registries from unilab.envs.locomotion.g1.joystick import G1WalkRewardConfig -from unilab.utils.algo_utils import ensure_registries pytest.importorskip("mujoco", reason="mujoco is required for G1 symmetry contract tests") diff --git a/tests/envs/test_allegro_domain_randomization.py b/tests/envs/test_allegro_domain_randomization.py index 489676f3f..7afaad129 100644 --- a/tests/envs/test_allegro_domain_randomization.py +++ b/tests/envs/test_allegro_domain_randomization.py @@ -14,7 +14,7 @@ "mujoco.batch_env not available (platform/libstdc++ issue)", allow_module_level=True ) -from unilab.utils.algo_utils import ensure_registries +from unilab.base.registry import ensure_registries @pytest.mark.slow diff --git a/tests/envs/test_env_configs.py b/tests/envs/test_env_configs.py index 6139a8814..58bfda01b 100644 --- a/tests/envs/test_env_configs.py +++ b/tests/envs/test_env_configs.py @@ -17,7 +17,7 @@ import numpy as np import pytest -from unilab.utils.algo_utils import ensure_registries +from unilab.base.registry import ensure_registries def _require_mujoco_runtime() -> None: @@ -52,7 +52,7 @@ def blocked_import(name, globals=None, locals=None, fromlist=(), level=0): from unilab.base.backend import create_backend from unilab.envs.manipulation.inhand_rot_allegro.rotation import AllegroRotationCfg from unilab.envs.motion_tracking.g1.tracking import G1MotionTrackingCfg - from unilab.utils.algo_utils import ensure_registries + from unilab.base.registry import ensure_registries ensure_registries() assert callable(create_backend) diff --git a/tests/envs/test_go1_domain_randomization.py b/tests/envs/test_go1_domain_randomization.py index 1cb92a02d..0237eb2c6 100644 --- a/tests/envs/test_go1_domain_randomization.py +++ b/tests/envs/test_go1_domain_randomization.py @@ -14,7 +14,7 @@ "mujoco.batch_env not available (platform/libstdc++ issue)", allow_module_level=True ) -from unilab.utils.algo_utils import ensure_registries +from unilab.base.registry import ensure_registries @pytest.mark.slow diff --git a/tests/integration/test_appo_rsl_reward.py b/tests/integration/test_appo_rsl_reward.py index b44c0fb1e..cecefdbc3 100644 --- a/tests/integration/test_appo_rsl_reward.py +++ b/tests/integration/test_appo_rsl_reward.py @@ -6,7 +6,7 @@ def test_appo_reward_override(): """Test APPO with reward override.""" from unilab.base import registry - from unilab.utils.algo_utils import ensure_registries + from unilab.base.registry import ensure_registries ensure_registries() @@ -30,7 +30,7 @@ def test_appo_reward_override(): def test_rsl_rl_reward_override(): """Test RSL-RL with reward override.""" from unilab.base import registry - from unilab.utils.algo_utils import ensure_registries + from unilab.base.registry import ensure_registries ensure_registries() diff --git a/tests/integration/test_reward_injection_integration.py b/tests/integration/test_reward_injection_integration.py index 57dd8c49e..1d76bd69d 100644 --- a/tests/integration/test_reward_injection_integration.py +++ b/tests/integration/test_reward_injection_integration.py @@ -41,8 +41,8 @@ def test_reward_injection_in_training(): def test_reward_override_propagation(): """Test reward override propagates through multiprocess collector.""" from unilab.base import registry + from unilab.base.registry import ensure_registries from unilab.envs.locomotion.go1.joystick import RewardConfig - from unilab.utils.algo_utils import ensure_registries ensure_registries() @@ -89,7 +89,7 @@ def test_reward_override_propagation(): def test_backward_compatibility_no_reward_config(): """Test env requires reward config - should fail without it.""" from unilab.base import registry - from unilab.utils.algo_utils import ensure_registries + from unilab.base.registry import ensure_registries ensure_registries() @@ -105,8 +105,8 @@ def test_backward_compatibility_no_reward_config(): def test_zero_scale_skips_computation(): """Test that reward functions with scale=0 are skipped.""" from unilab.base import registry + from unilab.base.registry import ensure_registries from unilab.envs.locomotion.go1.joystick import RewardConfig - from unilab.utils.algo_utils import ensure_registries ensure_registries() diff --git a/tests/scripts/test_mujoco_only_tooling_markers.py b/tests/scripts/test_mujoco_only_tooling_markers.py index e5169733c..9f471bc18 100644 --- a/tests/scripts/test_mujoco_only_tooling_markers.py +++ b/tests/scripts/test_mujoco_only_tooling_markers.py @@ -11,7 +11,7 @@ def test_mujoco_only_tooling_is_marked_explicitly(): root / "scripts" / "motion" / "replay_npz.py", root / "scripts" / "motion" / "bones_seed_csv_to_npz.py", root / "scripts" / "motion" / "replay_bones_seed_csv.py", - root / "src" / "unilab" / "utils" / "render_many.py", + root / "src" / "unilab" / "visualization" / "render_many.py", root / "src" / "unilab" / "envs" / "locomotion" / "g1" / "symmetry.py", ] diff --git a/tests/scripts/test_train_scripts.py b/tests/scripts/test_train_scripts.py index 5f77bf96f..c01b5feb6 100644 --- a/tests/scripts/test_train_scripts.py +++ b/tests/scripts/test_train_scripts.py @@ -580,7 +580,7 @@ def step(self, actions): def test_g1_motion_tracking_appo_reward_extraction_prefers_backend_specific_reward(): - from unilab.utils.reward_utils import extract_reward_config + from unilab.training.reward import extract_reward_config cfg = _appo_cfg(["task=g1_motion_tracking/motrix"]) @@ -1043,7 +1043,7 @@ def _play_interactive(): def test_play_wrapper_imports_shared_implementation(): """Verify play_interactive.py uses shared RslRlVecEnvWrapper.""" - from unilab.utils.rsl_rl_vec_env_wrapper import RslRlVecEnvWrapper as SharedWrapper + from unilab.training.rsl_rl import RslRlVecEnvWrapper as SharedWrapper mod = _play_interactive() # The wrapper class in play_interactive should be the shared one @@ -1055,7 +1055,7 @@ def test_play_wrapper_uses_current_reset_contract(): import numpy as np from tensordict import TensorDict - from unilab.utils.rsl_rl_vec_env_wrapper import RslRlVecEnvWrapper + from unilab.training.rsl_rl import RslRlVecEnvWrapper # Create a fake environment that returns (obs, info) tuple class FakeEnv: @@ -1090,7 +1090,7 @@ def test_play_wrapper_policy_obs_mode_actor(): """Verify wrapper supports policy_obs_mode='actor'.""" import numpy as np - from unilab.utils.rsl_rl_vec_env_wrapper import RslRlVecEnvWrapper + from unilab.training.rsl_rl import RslRlVecEnvWrapper class FakeEnv: def __init__(self): @@ -1128,7 +1128,7 @@ def reset(self, env_indices): def test_play_wrapper_flat_policy_excludes_critic_only_group(): import numpy as np - from unilab.utils.rsl_rl_vec_env_wrapper import RslRlVecEnvWrapper + from unilab.training.rsl_rl import RslRlVecEnvWrapper class FakeEnv: def __init__(self): @@ -1172,7 +1172,7 @@ def reset(self, env_indices): def test_play_wrapper_step_exports_timeout_bootstrap_obs(): import torch - from unilab.utils.rsl_rl_vec_env_wrapper import RslRlVecEnvWrapper + from unilab.training.rsl_rl import RslRlVecEnvWrapper class FakeEnv: def __init__(self): @@ -1556,8 +1556,6 @@ def sync(self): monkeypatch.setattr(mod, "RslRlVecEnvWrapper", FakeWrapper) monkeypatch.setattr(mod, "OnPolicyRunner", FakeRunner) monkeypatch.setattr(mod, "PPOConfig", lambda: types.SimpleNamespace(to_dict=lambda: {})) - monkeypatch.setattr(mod, "is_rsl_rl_v4", lambda: False) - monkeypatch.setattr(mod, "convert_config_v3_to_v4", lambda cfg: cfg) monkeypatch.setattr(mod.mujoco, "MjData", lambda model: object()) monkeypatch.setattr(mod.mujoco, "mj_setState", lambda *args, **kwargs: None) monkeypatch.setattr(mod.mujoco, "mj_forward", lambda *args, **kwargs: None) diff --git a/tests/training/test_training_helpers.py b/tests/training/test_training_helpers.py index d6bdc8e0f..2519f88fe 100644 --- a/tests/training/test_training_helpers.py +++ b/tests/training/test_training_helpers.py @@ -14,9 +14,9 @@ get_latest_checkpoint, get_latest_run, parse_checkpoint_path, - render_play_mode, resolve_task_checkpoint_path, ) +from unilab.visualization.playback import render_play_mode _ROOT_DIR = Path(__file__).resolve().parents[2] _CONF_DIR = _ROOT_DIR / "conf" @@ -262,7 +262,7 @@ def _render_states_get_frames(state_list, model_file, **kwargs): monkeypatch.setitem(sys.modules, "mediapy", fake_media) monkeypatch.setattr( - "unilab.utils.render_many.render_states_get_frames", + "unilab.visualization.render_many.render_states_get_frames", _render_states_get_frames, ) @@ -354,7 +354,7 @@ def _render_states_get_frames(state_list, model_file, **kwargs): monkeypatch.setitem(sys.modules, "mediapy", fake_media) monkeypatch.setattr( - "unilab.utils.render_many.render_states_get_frames", + "unilab.visualization.render_many.render_states_get_frames", _render_states_get_frames, ) diff --git a/tests/utils/test_algo_utils.py b/tests/utils/test_algo_utils.py index 418cca16c..9d80d5b52 100644 --- a/tests/utils/test_algo_utils.py +++ b/tests/utils/test_algo_utils.py @@ -1,4 +1,4 @@ -"""Tests for unilab.utils.algo_utils.""" +"""Tests for registry bootstrap and torch actor factory helpers.""" from __future__ import annotations @@ -8,7 +8,8 @@ import pytest -from unilab.utils.algo_utils import build_actor, ensure_registries +from unilab.algos.torch.common.actor_factory import build_actor +from unilab.base.registry import ensure_registries class TestEnsureRegistries: @@ -65,7 +66,7 @@ def test_empty_packages_list(self) -> None: def test_non_package_module_import_is_supported(self) -> None: """A plain module path should be accepted without package scanning.""" - ensure_registries(["unilab.utils.algo_utils"]) + ensure_registries(["unilab.base.registry"]) def test_declared_registry_module_failure_raises(self, tmp_path, monkeypatch) -> None: """Required packages should fail fast when declared registry modules fail to import.""" diff --git a/tests/utils/test_experiment_tracking.py b/tests/utils/test_experiment_tracking.py index 84b04eda1..009f0b7b3 100644 --- a/tests/utils/test_experiment_tracking.py +++ b/tests/utils/test_experiment_tracking.py @@ -4,9 +4,8 @@ import sys from pathlib import Path -from unilab.utils.experiment_tracking import ExperimentTracker, build_wandb_settings -from unilab.utils.offpolicy_logger import OffPolicyLogger -from unilab.utils.onpolicy_logger import OnPolicyLogger +from unilab.logging import OffPolicyLogger, OnPolicyLogger +from unilab.training.experiment import ExperimentTracker, build_wandb_settings class _FakeConfig(dict): diff --git a/tests/utils/test_final_observation.py b/tests/utils/test_final_observation.py index 5f2c6180c..261e5c4f9 100644 --- a/tests/utils/test_final_observation.py +++ b/tests/utils/test_final_observation.py @@ -2,7 +2,7 @@ import numpy as np -from unilab.utils.final_observation import ( +from unilab.base.final_observation import ( patch_transition_next_obs, resolve_terminal_observation_contract, resolve_transition_bootstrap_contract, diff --git a/tests/utils/test_math_utils.py b/tests/utils/test_math_utils.py index 19b2af6ca..0092f2ce4 100644 --- a/tests/utils/test_math_utils.py +++ b/tests/utils/test_math_utils.py @@ -1,10 +1,10 @@ -"""Tests for quaternion helpers in unilab.utils.math_utils.""" +"""Tests for quaternion helpers in unilab.envs.common.rotation.""" from __future__ import annotations import numpy as np -from unilab.utils.math_utils import ( +from unilab.envs.common.rotation import ( np_quat_angular_velocity, np_quat_ensure_continuity, np_quat_error_magnitude, diff --git a/tests/utils/test_obs_utils.py b/tests/utils/test_obs_utils.py index 1a2158c27..f08fbc7b6 100644 --- a/tests/utils/test_obs_utils.py +++ b/tests/utils/test_obs_utils.py @@ -1,11 +1,11 @@ -"""Tests for unilab.utils.obs_utils.""" +"""Tests for unilab.base.observations.""" from __future__ import annotations import numpy as np import pytest -from unilab.utils.obs_utils import flatten_obs_dict, get_obs_dims, split_obs_dict +from unilab.base.observations import flatten_obs_dict, get_obs_dims, split_obs_dict # --------------------------------------------------------------------------- # flatten_obs_dict — basic behaviour diff --git a/tests/utils/test_render_many.py b/tests/utils/test_render_many.py index e9715361c..832df5e80 100644 --- a/tests/utils/test_render_many.py +++ b/tests/utils/test_render_many.py @@ -1,4 +1,4 @@ -"""Tests for MuJoCo GL backend resolution in unilab.utils.render_many.""" +"""Tests for MuJoCo GL backend resolution in unilab.visualization.render_many.""" from __future__ import annotations @@ -18,8 +18,8 @@ def _reload_render_many(monkeypatch): monkeypatch.setitem(sys.modules, "mujoco", types.SimpleNamespace()) - sys.modules.pop("unilab.utils.render_many", None) - return importlib.import_module("unilab.utils.render_many") + sys.modules.pop("unilab.visualization.render_many", None) + return importlib.import_module("unilab.visualization.render_many") def test_resolve_gl_backend_uses_egl_when_probe_succeeds(monkeypatch) -> None: diff --git a/tests/utils/test_torch_utils.py b/tests/utils/test_torch_utils.py index e4ce27fed..61dbc0c66 100644 --- a/tests/utils/test_torch_utils.py +++ b/tests/utils/test_torch_utils.py @@ -1,4 +1,4 @@ -"""Tests for unilab.utils.torch_utils.""" +"""Tests for unilab.utils.tensor.""" from __future__ import annotations @@ -6,7 +6,7 @@ import pytest import torch -from unilab.utils.torch_utils import to_numpy, to_torch +from unilab.utils.tensor import to_numpy, to_torch class TestToTorch: diff --git a/tests/utils/test_utils_package_policy.py b/tests/utils/test_utils_package_policy.py new file mode 100644 index 000000000..31f6ce018 --- /dev/null +++ b/tests/utils/test_utils_package_policy.py @@ -0,0 +1,66 @@ +import importlib +from pathlib import Path + +import unilab.utils + +ALLOWED_UTILS_API = {"get_default_device", "to_numpy", "to_torch"} +ALLOWED_UTILS_MODULES = {"__init__", "device", "tensor"} +REMOVED_UTILS_SHIMS = { + "algo_utils", + "device_utils", + "experiment_tracking", + "final_observation", + "hardware_monitor", + "logging_common", + "math_utils", + "obs_utils", + "offpolicy_logger", + "onpolicy_logger", + "render_many", + "reward_utils", + "rsl_rl_compat", + "rsl_rl_vec_env_wrapper", + "run_utils", + "torch_utils", + "viser_scene", + "xml_utils", +} +REMOVED_OWNER_ALIASES = { + "unilab.algos.torch.offpolicy.logging", + "unilab.algos.torch.common.tensor", +} + + +def test_utils_api_is_whitelisted() -> None: + assert set(unilab.utils.__all__) == ALLOWED_UTILS_API + + +def test_utils_directory_is_whitelisted() -> None: + modules = {path.stem for path in Path("src/unilab/utils").glob("*.py")} + assert modules == ALLOWED_UTILS_MODULES + + +def test_repo_has_no_package_level_utils_imports() -> None: + current_file = Path(__file__).resolve() + for root in (Path("src"), Path("tests"), Path("scripts"), Path("benchmark")): + for path in root.rglob("*.py"): + if path.resolve() == current_file: + continue + assert "from unilab.utils import" not in path.read_text(encoding="utf-8"), path + + +def test_removed_utils_shims_are_not_importable() -> None: + for module_name in sorted(f"unilab.utils.{name}" for name in REMOVED_UTILS_SHIMS): + assert importlib.util.find_spec(module_name) is None, module_name + + +def test_removed_owner_aliases_are_not_importable() -> None: + for module_name in sorted(REMOVED_OWNER_ALIASES): + assert importlib.util.find_spec(module_name) is None, module_name + + +def test_algos_torch_common_no_longer_reexports_utils_primitives() -> None: + common = importlib.import_module("unilab.algos.torch.common") + assert "get_default_device" not in common.__all__ + assert "to_numpy" not in common.__all__ + assert "to_torch" not in common.__all__ diff --git a/tests/utils/test_viser_scene.py b/tests/utils/test_viser_scene.py index e9159f051..d3c18d010 100644 --- a/tests/utils/test_viser_scene.py +++ b/tests/utils/test_viser_scene.py @@ -6,7 +6,11 @@ import numpy as np import pytest -from unilab.utils.viser_scene import VISER_AVAILABLE, MujocoViserScene, build_visible_env_indices +from unilab.visualization.viser_scene import ( + VISER_AVAILABLE, + MujocoViserScene, + build_visible_env_indices, +) class _FakeHandle: