From f83e8c9a166914efcf80b390de8cf0d80b97024f Mon Sep 17 00:00:00 2001 From: tatp-yf Date: Fri, 24 Apr 2026 00:50:02 +0800 Subject: [PATCH 1/9] refactor: split utils ownership by owner layer --- .pre-commit-config.yaml | 1 + AGENTS.md | 4 + CONTRIBUTING.md | 2 + .../ADR-0004-registry-bootstrap-contract.md | 2 +- docs/developers/zh_CN/CONTRIBUTING.md | 2 + docs/users/zh_CN/02-simulation-backends.md | 2 +- pyproject.toml | 6 +- scripts/motion/bones_seed_csv_to_npz.py | 4 +- scripts/motion/csv_to_npz.py | 4 +- scripts/play_interactive.py | 10 +- scripts/play_viser.py | 20 +- scripts/train_appo.py | 6 +- scripts/train_mlx_ppo.py | 82 ++- scripts/train_offpolicy.py | 13 +- scripts/train_rsl_rl.py | 8 +- src/unilab/algos/mlx/common/__init__.py | 45 +- src/unilab/algos/mlx/common/rotation.py | 38 ++ src/unilab/algos/torch/appo/runner.py | 8 +- src/unilab/algos/torch/appo/worker.py | 8 +- src/unilab/algos/torch/common/__init__.py | 12 +- .../algos/torch/common/actor_factory.py | 54 ++ src/unilab/algos/torch/common/device.py | 22 + src/unilab/algos/torch/common/tensor.py | 3 + src/unilab/algos/torch/fast_sac/runner.py | 7 +- src/unilab/algos/torch/fast_td3/runner.py | 3 +- src/unilab/algos/torch/flash_sac/runner.py | 6 +- src/unilab/algos/torch/offpolicy/__init__.py | 8 +- src/unilab/algos/torch/offpolicy/logging.py | 14 + .../algos/torch/offpolicy/multi_gpu_runner.py | 2 +- src/unilab/algos/torch/offpolicy/runner.py | 5 +- src/unilab/algos/torch/offpolicy/worker.py | 7 +- src/unilab/base/__init__.py | 30 + src/unilab/base/backend/__init__.py | 16 + src/unilab/base/backend/base.py | 2 +- src/unilab/base/backend/motrix_backend.py | 2 +- src/unilab/base/backend/mujoco_backend.py | 34 +- src/unilab/base/backend/xml.py | 307 +++++++++ src/unilab/base/final_observation.py | 160 +++++ src/unilab/base/observations.py | 37 ++ src/unilab/base/registry.py | 62 ++ src/unilab/config/__init__.py | 3 + src/unilab/config/reward.py | 54 ++ src/unilab/docs/support_matrix.py | 4 +- src/unilab/envs/common/__init__.py | 1 + src/unilab/envs/common/math.py | 13 + src/unilab/envs/common/rotation.py | 295 +++++++++ .../envs/locomotion/common/dr_provider.py | 2 +- src/unilab/envs/locomotion/go1/joystick.py | 2 +- .../inhand_rot_allegro/rotation.py | 2 +- .../envs/manipulation/sharpa_inhand/base.py | 2 +- .../manipulation/sharpa_inhand/grasp_gen.py | 2 +- .../manipulation/sharpa_inhand/rotation.py | 2 +- .../envs/motion_tracking/g1/tracking.py | 6 +- src/unilab/training/__init__.py | 14 +- src/unilab/training/backend_adapter.py | 4 +- src/unilab/training/common.py | 359 +---------- src/unilab/training/logging/__init__.py | 23 + src/unilab/training/logging/common.py | 316 +++++++++ src/unilab/training/logging/experiment.py | 425 ++++++++++++ src/unilab/training/logging/offpolicy.py | 347 ++++++++++ src/unilab/training/logging/onpolicy.py | 196 ++++++ src/unilab/training/monitoring.py | 60 ++ src/unilab/training/run.py | 166 +++++ src/unilab/utils/__init__.py | 18 +- src/unilab/utils/algo_utils.py | 138 +--- src/unilab/utils/device.py | 12 + src/unilab/utils/device_utils.py | 37 +- src/unilab/utils/experiment_tracking.py | 440 +------------ src/unilab/utils/final_observation.py | 182 +----- src/unilab/utils/hardware_monitor.py | 66 +- src/unilab/utils/logging_common.py | 320 +-------- src/unilab/utils/math_utils.py | 394 ++--------- src/unilab/utils/obs_utils.py | 59 +- src/unilab/utils/offpolicy_logger.py | 473 +------------- src/unilab/utils/onpolicy_logger.py | 200 +----- src/unilab/utils/render_many.py | 610 +----------------- src/unilab/utils/reward_utils.py | 60 +- src/unilab/utils/rsl_rl_compat.py | 236 +------ src/unilab/utils/rsl_rl_vec_env_wrapper.py | 172 +---- src/unilab/utils/run_utils.py | 57 +- src/unilab/utils/tensor.py | 33 + src/unilab/utils/torch_utils.py | 37 +- src/unilab/utils/viser_scene.py | 289 +-------- src/unilab/utils/xml_utils.py | 311 +-------- src/unilab/visualization/__init__.py | 5 + src/unilab/visualization/playback.py | 182 ++++++ src/unilab/visualization/render_many.py | 606 +++++++++++++++++ src/unilab/visualization/viser_scene.py | 281 ++++++++ .../algos/test_fast_sac_symmetry_contract.py | 6 +- tests/algos/test_mlx_ppo.py | 4 +- tests/algos/test_rsl_rl_runner.py | 6 +- tests/base/test_reward_override.py | 2 +- tests/base/test_sim_backend.py | 2 +- tests/config/test_reward_injection.py | 4 +- .../locomotion/g1/test_issue175_regression.py | 2 +- .../locomotion/g1/test_symmetry_contract.py | 2 +- .../envs/test_allegro_domain_randomization.py | 2 +- tests/envs/test_env_configs.py | 4 +- tests/envs/test_go1_domain_randomization.py | 2 +- tests/integration/test_appo_rsl_reward.py | 4 +- .../test_reward_injection_integration.py | 6 +- .../test_mujoco_only_tooling_markers.py | 2 +- tests/scripts/test_train_scripts.py | 12 +- tests/training/test_training_helpers.py | 4 +- tests/utils/test_algo_utils.py | 7 +- tests/utils/test_experiment_tracking.py | 6 +- tests/utils/test_final_observation.py | 2 +- tests/utils/test_math_utils.py | 4 +- tests/utils/test_obs_utils.py | 4 +- tests/utils/test_render_many.py | 6 +- tests/utils/test_torch_utils.py | 4 +- tests/utils/test_utils_package_policy.py | 38 ++ tests/utils/test_viser_scene.py | 6 +- 113 files changed, 4369 insertions(+), 4346 deletions(-) create mode 100644 src/unilab/algos/mlx/common/rotation.py create mode 100644 src/unilab/algos/torch/common/actor_factory.py create mode 100644 src/unilab/algos/torch/common/device.py create mode 100644 src/unilab/algos/torch/common/tensor.py create mode 100644 src/unilab/algos/torch/offpolicy/logging.py create mode 100644 src/unilab/base/backend/xml.py create mode 100644 src/unilab/base/final_observation.py create mode 100644 src/unilab/base/observations.py create mode 100644 src/unilab/config/reward.py create mode 100644 src/unilab/envs/common/__init__.py create mode 100644 src/unilab/envs/common/math.py create mode 100644 src/unilab/envs/common/rotation.py create mode 100644 src/unilab/training/logging/__init__.py create mode 100644 src/unilab/training/logging/common.py create mode 100644 src/unilab/training/logging/experiment.py create mode 100644 src/unilab/training/logging/offpolicy.py create mode 100644 src/unilab/training/logging/onpolicy.py create mode 100644 src/unilab/training/monitoring.py create mode 100644 src/unilab/training/run.py create mode 100644 src/unilab/utils/device.py create mode 100644 src/unilab/utils/tensor.py create mode 100644 src/unilab/visualization/__init__.py create mode 100644 src/unilab/visualization/playback.py create mode 100644 src/unilab/visualization/render_many.py create mode 100644 src/unilab/visualization/viser_scene.py create mode 100644 tests/utils/test_utils_package_policy.py 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..450a8030a 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -32,6 +32,10 @@ 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` +- 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/config/structured_configs.py` - async runner: `src/unilab/ipc/async_runner.py` 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/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/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..94ecfe0a6 100644 --- a/scripts/play_interactive.py +++ b/scripts/play_interactive.py @@ -53,15 +53,15 @@ ensure_registries() -from unilab.base import registry -from unilab.config.structured_configs import PPOConfig as _StructuredPPOConfig -from unilab.utils.rsl_rl_compat import ( +from unilab.algos.torch.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.algos.torch.rsl_rl.vec_env_wrapper import RslRlVecEnvWrapper +from unilab.base import registry +from unilab.config.structured_configs import PPOConfig as _StructuredPPOConfig PPOConfig = _StructuredPPOConfig @@ -122,8 +122,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, diff --git a/scripts/play_viser.py b/scripts/play_viser.py index 081899a0d..8c6e27ac7 100644 --- a/scripts/play_viser.py +++ b/scripts/play_viser.py @@ -47,19 +47,23 @@ if str(ROOT_DIR) not in sys.path: sys.path.insert(0, str(ROOT_DIR)) -from unilab.training import ( - ensure_registries, - get_entrypoint_log_root, -) -from unilab.utils.render_many import get_grid_offsets -from unilab.utils.rsl_rl_compat import ( +from unilab.algos.torch.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.utils.viser_scene import VISER_AVAILABLE, MujocoViserScene, build_visible_env_indices +from unilab.algos.torch.rsl_rl.vec_env_wrapper import RslRlVecEnvWrapper +from unilab.training import ( + ensure_registries, + get_entrypoint_log_root, +) +from unilab.visualization.render_many import get_grid_offsets +from unilab.visualization.viser_scene import ( + VISER_AVAILABLE, + MujocoViserScene, + build_visible_env_indices, +) ensure_registries() diff --git a/scripts/train_appo.py b/scripts/train_appo.py index 3aad4db56..1825d0fdb 100644 --- a/scripts/train_appo.py +++ b/scripts/train_appo.py @@ -22,7 +22,7 @@ get_log_root, render_play_mode, ) -from unilab.utils.experiment_tracking import ExperimentTracker +from unilab.training.logging.experiment import ExperimentTracker def build_appo_runner_kwargs( @@ -113,7 +113,7 @@ 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 + from unilab.algos.torch.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" @@ -136,7 +136,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 diff --git a/scripts/train_mlx_ppo.py b/scripts/train_mlx_ppo.py index 568048026..b0fc3e67d 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,19 +14,20 @@ 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.training import ( BackendAdapter, create_env, @@ -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.logging.experiment import ExperimentTracker +from unilab.training.logging.onpolicy import OnPolicyLogger 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..200919cb4 100644 --- a/scripts/train_offpolicy.py +++ b/scripts/train_offpolicy.py @@ -25,7 +25,7 @@ from unilab.training import ( resolve_checkpoint_path as resolve_checkpoint_path_common, ) -from unilab.utils.experiment_tracking import ExperimentTracker +from unilab.training.logging.experiment import ExperimentTracker 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..1dd1f7847 100644 --- a/scripts/train_rsl_rl.py +++ b/scripts/train_rsl_rl.py @@ -16,6 +16,8 @@ if str(ROOT_DIR) not in sys.path: sys.path.insert(0, str(ROOT_DIR)) +from unilab.algos.torch.rsl_rl.vec_env_wrapper import RslRlVecEnvWrapper +from unilab.base.backend.xml import materialize_scene_visual_override from unilab.training import ( BackendAdapter, create_env, @@ -26,9 +28,7 @@ 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.logging.experiment import ExperimentTracker, patch_rsl_rl_wandb_writer try: from rsl_rl.runners import OnPolicyRunner @@ -36,7 +36,7 @@ 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 +from unilab.algos.torch.rsl_rl.compat import convert_config_v5, is_rsl_rl_v5 def _backend_adapter(cfg: DictConfig) -> BackendAdapter: 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..84fc62013 100644 --- a/src/unilab/algos/torch/appo/runner.py +++ b/src/unilab/algos/torch/appo/runner.py @@ -19,9 +19,9 @@ from unilab.algos.torch.appo.learner import APPOLearner from unilab.algos.torch.appo.worker import appo_collector_fn +from unilab.algos.torch.rsl_rl.compat import convert_config_v3_to_v4, is_rsl_rl_v4, is_rsl_rl_v5 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.training.logging.offpolicy import OffPolicyLogger class APPORunner(AsyncRunner): @@ -87,8 +87,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() diff --git a/src/unilab/algos/torch/appo/worker.py b/src/unilab/algos/torch/appo/worker.py index a9a645d05..36f8c402f 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( @@ -76,9 +76,9 @@ def appo_collector_fn( from tensordict import TensorDict + from unilab.algos.torch.rsl_rl.compat import convert_config_v3_to_v4, is_rsl_rl_v4, is_rsl_rl_v5 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() diff --git a/src/unilab/algos/torch/common/__init__.py b/src/unilab/algos/torch/common/__init__.py index b851d2614..c4ca017fa 100644 --- a/src/unilab/algos/torch/common/__init__.py +++ b/src/unilab/algos/torch/common/__init__.py @@ -1,17 +1,23 @@ +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 +from unilab.utils.device import get_default_device +from unilab.utils.tensor import to_numpy, to_torch __all__ = [ "EmpiricalNormalization", "DistributionalQNetwork", "Critic", + "get_default_device", + "get_env_dims", + "to_numpy", + "to_torch", "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/algos/torch/common/device.py b/src/unilab/algos/torch/common/device.py new file mode 100644 index 000000000..99f61379f --- /dev/null +++ b/src/unilab/algos/torch/common/device.py @@ -0,0 +1,22 @@ +from unilab.base import registry +from unilab.utils.device import get_default_device + + +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.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 + ) + obs_dim, critic_dim = get_obs_dims_from_spec(env.obs_groups_spec) + action_shape = env.action_space.shape + assert action_shape is not None + action_dim = action_shape[0] + env.close() # type: ignore[attr-defined] + return obs_dim, action_dim, critic_dim + + +__all__ = ["get_default_device", "get_env_dims"] diff --git a/src/unilab/algos/torch/common/tensor.py b/src/unilab/algos/torch/common/tensor.py new file mode 100644 index 000000000..ce2f5e258 --- /dev/null +++ b/src/unilab/algos/torch/common/tensor.py @@ -0,0 +1,3 @@ +from unilab.utils.tensor import to_numpy, to_torch + +__all__ = ["to_numpy", "to_torch"] 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..69068d011 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.training.logging.offpolicy 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/logging.py b/src/unilab/algos/torch/offpolicy/logging.py new file mode 100644 index 000000000..8ccd0e92f --- /dev/null +++ b/src/unilab/algos/torch/offpolicy/logging.py @@ -0,0 +1,14 @@ +from __future__ import annotations + +import warnings + +from unilab.training.logging.offpolicy import OffPolicyLogger + +warnings.warn( + "`unilab.algos.torch.offpolicy.logging` is deprecated and will be removed in 0.2.0; " + "use `unilab.training.logging.offpolicy` instead.", + DeprecationWarning, + stacklevel=2, +) + +__all__ = ["OffPolicyLogger"] diff --git a/src/unilab/algos/torch/offpolicy/multi_gpu_runner.py b/src/unilab/algos/torch/offpolicy/multi_gpu_runner.py index 5ef195cc4..cc2f25c07 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.training.logging.offpolicy 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..cf548c668 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.training.logging.offpolicy 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..1b537878d 100644 --- a/src/unilab/base/backend/mujoco_backend.py +++ b/src/unilab/base/backend/mujoco_backend.py @@ -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/base/backend/xml.py b/src/unilab/base/backend/xml.py new file mode 100644 index 000000000..d85697c70 --- /dev/null +++ b/src/unilab/base/backend/xml.py @@ -0,0 +1,307 @@ +from __future__ import annotations + +import os +import tempfile +import xml.etree.ElementTree as ET +from collections.abc import Iterator, Sequence +from pathlib import Path + + +def _enable_discardvisual(root: ET.Element) -> None: + compiler_tag = root.find("compiler") + if compiler_tag is None: + compiler_tag = ET.Element("compiler") + root.insert(0, compiler_tag) + compiler_tag.set("discardvisual", "true") + + +def create_discardvisual_xml(model_file: str) -> str: + tree = ET.parse(model_file) + _enable_discardvisual(tree.getroot()) + return _write_temp_xml(tree, model_file) + + +def _iter_expanded_children( + parent: ET.Element, base_dir: Path +) -> Iterator[tuple[ET.Element, Path]]: + for child in parent: + if child.tag != "include": + yield child, base_dir + continue + + include_file = child.get("file") + if not include_file: + raise ValueError(f"Invalid without file attribute in {base_dir}") + include_path = (base_dir / include_file).resolve() + include_root = ET.parse(include_path).getroot() + yield from _iter_expanded_children(include_root, include_path.parent) + + +def _iter_named_bodies(root: ET.Element, base_dir: Path) -> Iterator[str]: + for child, child_base_dir in _iter_expanded_children(root, base_dir): + if child.tag == "body": + body_name = child.get("name") + if body_name: + yield body_name + yield from _iter_named_bodies(child, child_base_dir) + + +def _get_named_bodies(model_file: str) -> tuple[list[int], list[str]]: + model_path = Path(model_file).resolve() + names = list(_iter_named_bodies(ET.parse(model_path).getroot(), model_path.parent)) + ids = list(range(1, len(names) + 1)) + return ids, names + + +def get_named_body_ids(model_file: str, names: Sequence[str]) -> list[int]: + """Resolve MuJoCo-style body ids from XML without importing mujoco.""" + body_ids, body_names = _get_named_bodies(model_file) + body_id_by_name = dict(zip(body_names, body_ids, strict=True)) + missing = [name for name in names if name not in body_id_by_name] + if missing: + missing_str = ", ".join(missing) + raise ValueError(f"Bodies not found in XML '{model_file}': {missing_str}") + return [body_id_by_name[name] for name in names] + + +def _add_w_sensors(sensor_tag: ET.Element, valid_bnames: list[str]) -> None: + for bname in valid_bnames: + ET.SubElement( + sensor_tag, "framepos", name=f"track_pos_w_{bname}", objtype="xbody", objname=bname + ) + for bname in valid_bnames: + ET.SubElement( + sensor_tag, "framequat", name=f"track_quat_w_{bname}", objtype="xbody", objname=bname + ) + for bname in valid_bnames: + ET.SubElement( + sensor_tag, + "framelinvel", + name=f"track_linvel_w_{bname}", + objtype="xbody", + objname=bname, + ) + for bname in valid_bnames: + ET.SubElement( + sensor_tag, + "frameangvel", + name=f"track_angvel_w_{bname}", + objtype="xbody", + objname=bname, + ) + + +def _add_b_sensors(sensor_tag: ET.Element, valid_bnames: list[str], baselink_name: str) -> None: + for bname in valid_bnames: + ET.SubElement( + sensor_tag, + "framepos", + name=f"track_pos_b_{bname}", + objtype="xbody", + objname=bname, + reftype="xbody", + refname=baselink_name, + ) + for bname in valid_bnames: + ET.SubElement( + sensor_tag, + "framequat", + name=f"track_quat_b_{bname}", + objtype="xbody", + objname=bname, + reftype="xbody", + refname=baselink_name, + ) + for bname in valid_bnames: + ET.SubElement( + sensor_tag, + "framelinvel", + name=f"track_linvel_b_{bname}", + objtype="xbody", + objname=bname, + reftype="xbody", + refname=baselink_name, + ) + for bname in valid_bnames: + ET.SubElement( + sensor_tag, + "frameangvel", + name=f"track_angvel_b_{bname}", + objtype="xbody", + objname=bname, + reftype="xbody", + refname=baselink_name, + ) + + +def _write_temp_xml(tree: ET.ElementTree[ET.Element], model_file: str) -> str: # type: ignore[type-arg] + fd, output_path = tempfile.mkstemp( + suffix=".xml", dir=os.path.dirname(os.path.abspath(model_file)) + ) + os.close(fd) + tree.write(output_path) + return output_path + + +def _format_values(values: list[float] | tuple[float, ...]) -> str: + return " ".join(str(float(value)) for value in values) + + +def materialize_scene_visual_override( + source_model_file: str, + *, + ground_texture_file: str | None = None, + ground_texrepeat: list[float] | tuple[float, float] | None = None, + skybox_rgb1: list[float] | tuple[float, float, float] | None = None, + skybox_rgb2: list[float] | tuple[float, float, float] | None = None, +) -> str: + """Create a temporary scene XML with visual-only overrides applied.""" + tree = ET.parse(source_model_file) + root = tree.getroot() + asset_tag = root.find("asset") + if asset_tag is None: + raise ValueError(f"Scene '{source_model_file}' is missing an tag.") + + if skybox_rgb1 is not None or skybox_rgb2 is not None: + skybox = asset_tag.find("./texture[@type='skybox']") + if skybox is None: + raise ValueError(f"Scene '{source_model_file}' is missing a skybox texture.") + if skybox_rgb1 is not None: + skybox.set("rgb1", _format_values(tuple(skybox_rgb1))) + if skybox_rgb2 is not None: + skybox.set("rgb2", _format_values(tuple(skybox_rgb2))) + + if ground_texture_file is not None: + ground_texture = asset_tag.find("./texture[@name='groundplane']") + if ground_texture is None: + raise ValueError(f"Scene '{source_model_file}' is missing the groundplane texture.") + for attr in ("builtin", "mark", "rgb1", "rgb2", "markrgb", "width", "height"): + ground_texture.attrib.pop(attr, None) + ground_texture.set("file", str(Path(ground_texture_file))) + + if ground_texrepeat is not None: + ground_material = asset_tag.find("./material[@name='groundplane']") + if ground_material is None: + raise ValueError(f"Scene '{source_model_file}' is missing the groundplane material.") + ground_material.set("texrepeat", _format_values(tuple(ground_texrepeat))) + + return _write_temp_xml(tree, source_model_file) + + +def inject_mujoco_tracking_sensors( + model_file: str, + baselink_name: str | None = None, +) -> tuple[str, list, list]: + """为 MuJoCo 后端注入 tracking sensors。 + + 注入所有 body 的世界系 (_w) sensors;若指定 baselink_name, + 同时注入相对 baselink 坐标系的 (_b) sensors。 + + Returns: + (tmp_xml_path, tracked_body_ids, valid_bnames) + """ + tracked_body_ids, valid_bnames = _get_named_bodies(model_file) + + tree = ET.parse(model_file) + root = tree.getroot() + sensor_tag = root.find("sensor") + if sensor_tag is None: + sensor_tag = ET.SubElement(root, "sensor") + + _add_w_sensors(sensor_tag, valid_bnames) + if baselink_name and baselink_name in valid_bnames: + _add_b_sensors(sensor_tag, valid_bnames, baselink_name) + + return _write_temp_xml(tree, model_file), tracked_body_ids, valid_bnames + + +def inject_motrix_tracking_sensors(model_file: str, baselink_name: str) -> tuple[str, list, list]: + """为 MotrixSim 后端注入 tracking sensors。 + + 只注入相对 baselink 坐标系的 (_b) sensors。 + 世界系 (_w) 数据由 motrixsim body API 直接提供,无需 sensor 注入。 + + Returns: + (tmp_xml_path, tracked_body_ids, valid_bnames) + """ + tracked_body_ids, valid_bnames = _get_named_bodies(model_file) + + tree = ET.parse(model_file) + root = tree.getroot() + sensor_tag = root.find("sensor") + if sensor_tag is None: + sensor_tag = ET.SubElement(root, "sensor") + + _add_b_sensors(sensor_tag, valid_bnames, baselink_name) + + return _write_temp_xml(tree, model_file), tracked_body_ids, valid_bnames + + +def processed_xml(xml_path): + xml_dir = os.path.dirname(os.path.abspath(xml_path)) + + tree = ET.parse(xml_path) + root = tree.getroot() + + compiler = root.find("compiler") + if compiler is not None: + meshdir = compiler.get("meshdir") + if meshdir: + abs_meshdir = os.path.normpath(os.path.join(xml_dir, meshdir)) + compiler.set("meshdir", abs_meshdir) + + bodys = root.findall(".//body") + + geom_names = [] + for body in bodys: + body_name = body.get("name", "unnamed_body") + geoms = body.findall("geom") + + if geoms: + filtered_geoms = [] + for geom in geoms: + geom_class = geom.get("class") + if geom_class != "visual": + filtered_geoms.append(geom) + + if filtered_geoms: + i = 0 + for geom in filtered_geoms: + geom_name = geom.get("name", "unnamed_geom") + if geom_name == "unnamed_geom": + new_name = f"{body_name}_geom{i}" + i += 1 + geom.set("name", new_name) + geom_name = new_name + geom_names.append(geom_name) + + new_xml_string = ET.tostring(root, encoding="unicode") + return new_xml_string, geom_names + + +def add_sensor(root, sensor_type, name, **kwargs): + """ + 在 MuJoCo XML 的 sensor 节点下添加传感器的通用函数。 + + 参数: + - root: XML 的根节点 + - sensor_type: 传感器标签名 (如 'gyro', 'contact', 'framepos') + - name: 传感器的 name 属性 + - **kwargs: 其他任意属性 (如 site='imu', geom1='floor' 等) + """ + # 1. 查找或创建 标签 + sensor_element = root.find("sensor") + if sensor_element is None: + sensor_element = ET.SubElement(root, "sensor") + + # 2. 创建具体的传感器子节点 + sensor = ET.SubElement(sensor_element, sensor_type) + + # 3. 设置必选的 name 属性 + sensor.set("name", name) + + # 4. 循环设置其他传入的属性 + for key, value in kwargs.items(): + sensor.set(key, str(value)) + + return sensor diff --git a/src/unilab/base/final_observation.py b/src/unilab/base/final_observation.py new file mode 100644 index 000000000..f18fc1ba6 --- /dev/null +++ b/src/unilab/base/final_observation.py @@ -0,0 +1,160 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +import numpy as np + + +@dataclass(frozen=True) +class TransitionBootstrapContract: + actor_next_obs: np.ndarray + transition_next_obs: np.ndarray + terminal_mask: np.ndarray + timeout_terminal_mask: np.ndarray + actor_next_critic: np.ndarray | None = None + transition_next_critic: np.ndarray | None = None + + +@dataclass(frozen=True) +class TerminalObservationContract: + terminal_obs: np.ndarray | None + terminal_mask: np.ndarray + timeout_terminal_mask: np.ndarray + terminal_critic: np.ndarray | None = None + + +def patch_transition_next_obs( + next_obs: np.ndarray, + final_observation: dict[str, Any] | None = None, + done: np.ndarray | None = None, + info: dict[str, Any] | None = None, + next_critic: np.ndarray | None = None, +) -> tuple[np.ndarray, np.ndarray | None, np.ndarray]: + """Patch transition next obs with final_observation without mutating actor inputs.""" + terminal_contract = resolve_terminal_observation_contract( + next_obs_batch_size=next_obs.shape[0], + final_observation=final_observation, + done=done, + info=info, + ) + if not np.any(terminal_contract.terminal_mask) or terminal_contract.terminal_obs is None: + return ( + next_obs, + next_critic, + np.zeros((next_obs.shape[0],), dtype=bool), + ) + + transition_next_obs = next_obs.copy() + transition_next_obs[terminal_contract.terminal_mask] = np.asarray( + terminal_contract.terminal_obs, dtype=next_obs.dtype + )[terminal_contract.terminal_mask] + + transition_next_critic = next_critic + if next_critic is not None and terminal_contract.terminal_critic is not None: + transition_next_critic = next_critic.copy() + transition_next_critic[terminal_contract.terminal_mask] = np.asarray( + terminal_contract.terminal_critic, dtype=next_critic.dtype + )[terminal_contract.terminal_mask] + + return ( + transition_next_obs, + transition_next_critic, + terminal_contract.terminal_mask, + ) + + +def resolve_transition_bootstrap_contract( + next_obs: np.ndarray, + info: dict[str, Any] | None = None, + final_observation: dict[str, Any] | None = None, + done: np.ndarray | None = None, + truncated: np.ndarray | None = None, + next_critic: np.ndarray | None = None, +) -> TransitionBootstrapContract: + """Resolve actor/storage observations and timeout bootstrap masks for a step.""" + ( + transition_next_obs, + transition_next_critic, + terminal_mask, + ) = patch_transition_next_obs( + next_obs, + final_observation=final_observation, + done=done, + info=info, + next_critic=next_critic, + ) + timeout_terminal_mask = terminal_mask + if truncated is not None: + timeout_terminal_mask = np.logical_and( + terminal_mask, np.asarray(truncated, dtype=bool).ravel() + ) + return TransitionBootstrapContract( + actor_next_obs=next_obs, + transition_next_obs=transition_next_obs, + terminal_mask=terminal_mask, + timeout_terminal_mask=timeout_terminal_mask, + actor_next_critic=next_critic, + transition_next_critic=transition_next_critic, + ) + + +def resolve_terminal_observation_contract( + next_obs_batch_size: int, + final_observation: dict[str, Any] | None = None, + done: np.ndarray | None = None, + info: dict[str, Any] | None = None, + truncated: np.ndarray | None = None, +) -> TerminalObservationContract: + """Resolve terminal observation facts without constructing patched next obs.""" + terminal_mask = _resolve_terminal_mask(next_obs_batch_size, done, info) + resolved_final_observation = _resolve_final_observation(final_observation, info) + + terminal_obs: np.ndarray | None = None + terminal_critic: np.ndarray | None = None + if np.any(terminal_mask) and isinstance(resolved_final_observation, dict): + terminal_obs = resolved_final_observation.get("obs") + terminal_critic = resolved_final_observation.get("critic") + + timeout_terminal_mask = terminal_mask + if truncated is not None: + timeout_terminal_mask = np.logical_and( + terminal_mask, np.asarray(truncated, dtype=bool).ravel() + ) + + return TerminalObservationContract( + terminal_obs=terminal_obs, + terminal_mask=terminal_mask, + timeout_terminal_mask=timeout_terminal_mask, + terminal_critic=terminal_critic, + ) + + +def _resolve_final_observation( + final_observation: dict[str, Any] | None, + info: dict[str, Any] | None, +) -> dict[str, Any] | None: + if isinstance(final_observation, dict): + return final_observation + if isinstance(info, dict): + final_obs = info.get("final_observation") + if isinstance(final_obs, dict): + return final_obs + return None + + +def _resolve_terminal_mask( + next_obs_batch_size: int, + done: np.ndarray | None, + info: dict[str, Any] | None, +) -> np.ndarray: + if done is not None: + done_mask = np.asarray(done, dtype=bool).ravel() + if done_mask.shape == (next_obs_batch_size,): + return done_mask + return np.zeros((next_obs_batch_size,), dtype=bool) + if isinstance(info, dict): + terminal_mask = np.asarray(info.get("_final_observation"), dtype=bool) + if terminal_mask.shape == (next_obs_batch_size,): + return terminal_mask + return np.zeros((next_obs_batch_size,), dtype=bool) diff --git a/src/unilab/base/observations.py b/src/unilab/base/observations.py new file mode 100644 index 000000000..3b17f9cfa --- /dev/null +++ b/src/unilab/base/observations.py @@ -0,0 +1,37 @@ +from __future__ import annotations + +import numpy as np + + +def flatten_obs_dict(obs: dict[str, np.ndarray]) -> np.ndarray: + """Concatenate obs groups in insertion order -> flat (N, total_dim) array.""" + return np.concatenate(list(obs.values()), axis=1) + + +def flatten_policy_obs_dict(obs: dict[str, np.ndarray]) -> np.ndarray: + """Build actor-policy inputs from the single actor observation group.""" + return obs["obs"] + + +def split_obs_dict(obs: dict[str, np.ndarray]) -> tuple[np.ndarray, np.ndarray]: + """Split observation dict into (actor_obs, critic_obs). + + When no separate critic group exists, critic_obs == actor_obs. + """ + actor = obs["obs"] + return actor, obs.get("critic", actor) + + +def get_obs_dims(obs_groups_spec: dict[str, int]) -> tuple[int, int]: + """Extract (actor_obs_dim, critic_obs_dim) from obs_groups_spec. + + When no separate critic group exists, critic_obs_dim == actor_obs_dim. + """ + obs_dim = obs_groups_spec.get("obs", 0) + return obs_dim, obs_groups_spec.get("critic", obs_dim) + + +def get_critic_base_dim(obs_groups_spec: dict[str, int]) -> int: + """Get critic observation dim, falling back to actor obs when absent.""" + critic_dim = obs_groups_spec.get("critic", 0) + return critic_dim if critic_dim > 0 else obs_groups_spec.get("obs", 0) 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 index e69de29bb..6684d663a 100644 --- a/src/unilab/config/__init__.py +++ b/src/unilab/config/__init__.py @@ -0,0 +1,3 @@ +from unilab.config.reward import RewardDict, extract_reward_config, resolve_reward_dict + +__all__ = ["RewardDict", "extract_reward_config", "resolve_reward_dict"] diff --git a/src/unilab/config/reward.py b/src/unilab/config/reward.py new file mode 100644 index 000000000..ee68d18fc --- /dev/null +++ b/src/unilab/config/reward.py @@ -0,0 +1,54 @@ +"""Utility functions for reward config handling.""" + +from typing import Any, cast + +from omegaconf import DictConfig, OmegaConf + +RewardDict = dict[str, Any] + + +def _to_reward_dict(value: object, *, error_message: str) -> RewardDict: + """Convert an OmegaConf container into a plain reward dictionary.""" + resolved = OmegaConf.to_container(value, resolve=True) + if not isinstance(resolved, dict): + raise ValueError(error_message) + # Some reward configs are mounted as a full `reward:` section. + # Env config injection expects the inner reward mapping. + if set(resolved) == {"reward"} and isinstance(resolved["reward"], dict): + return cast(RewardDict, resolved["reward"]) + return cast(RewardDict, resolved) + + +def resolve_reward_dict(cfg: DictConfig) -> RewardDict: + """Resolve the reward config from the final composed config.""" + reward_cfg = OmegaConf.select(cfg, "reward") + if not reward_cfg: + raise ValueError( + "Missing 'reward' config in Hydra. Reward config must be explicitly provided." + ) + + reward_dict = _to_reward_dict( + reward_cfg, + error_message="Reward config must resolve to a mapping.", + ) + if not reward_dict: + raise ValueError( + "Reward config resolved to empty. Please select a non-default reward override." + ) + + return reward_dict + + +def extract_reward_config(cfg: DictConfig) -> dict[str, RewardDict]: + """Extract and validate reward config from Hydra config. + + Args: + cfg: Hydra DictConfig containing reward section + + Returns: + Dictionary with reward_config key for env_cfg_override + + Raises: + ValueError: If reward config is missing + """ + return {"reward_config": resolve_reward_dict(cfg)} 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/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/envs/common/rotation.py b/src/unilab/envs/common/rotation.py new file mode 100644 index 000000000..28aa7cd55 --- /dev/null +++ b/src/unilab/envs/common/rotation.py @@ -0,0 +1,295 @@ +from __future__ import annotations + +import numpy as np + + +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 + q2_was_1d = q2.ndim == 1 + + if q1_was_1d: + q1 = q1[None, :] + if q2_was_1d: + q2 = q2[None, :] + + if q1.shape[0] == 1 and q2.shape[0] > 1: + q1 = np.broadcast_to(q1, q2.shape) + elif q2.shape[0] == 1 and q1.shape[0] > 1: + q2 = np.broadcast_to(q2, q1.shape) + + w1, x1, y1, z1 = q1[:, 0], q1[:, 1], q1[:, 2], q1[:, 3] + w2, x2, y2, z2 = q2[:, 0], q2[:, 1], q2[:, 2], q2[:, 3] + result = np.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, + ) + return result[0] if q1_was_1d and q2_was_1d else result + + +def np_quat_conjugate(q: np.ndarray) -> np.ndarray: + """Conjugate of unit quaternions (N, 4) or (4,), w-first.""" + if q.ndim == 1: + return np.array([q[0], -q[1], -q[2], -q[3]]) + conj = q.copy() + conj[:, 1:] *= -1 + return conj # type: ignore[no-any-return] + + +def np_quat_canonicalize(q: np.ndarray) -> np.ndarray: + """Flip quaternion signs so the real part is non-negative.""" + q_was_1d = q.ndim == 1 + if q_was_1d: + q = q[None, :] + + sign = np.where(q[:, 0:1] < 0.0, -1.0, 1.0) + result = q * sign + canonical: np.ndarray = result[0] if q_was_1d else result + return canonical + + +def np_quat_ensure_continuity(q: np.ndarray) -> np.ndarray: + """Flip quaternion signs in a time sequence to keep adjacent dots non-negative.""" + if q.ndim != 2 or q.shape[1] != 4: + raise ValueError(f"Expected quaternion sequence with shape (T, 4), got {q.shape}") + + result = np.array(q, copy=True) + for i in range(1, result.shape[0]): + if float(np.dot(result[i - 1], result[i])) < 0.0: + result[i] *= -1.0 + return result + + +def np_quat_to_axis_angle(q: np.ndarray) -> np.ndarray: + """Convert unit quaternion batch (N, 4), w-first, to axis-angle vectors (N, 3). + + Adapted from PyTorch3D. Uses atan2 + Taylor expansion for numerical + stability near zero rotation. + """ + q = np_quat_canonicalize(q) + xyz = q[:, 1:] # (N, 3) imaginary part + w = q[:, 0:1] # (N, 1) real part + norms = np.linalg.norm(xyz, axis=-1, keepdims=True) # (N, 1) + half_angle = np.arctan2(norms, w) # (N, 1) + angle = 2.0 * half_angle # (N, 1) + small = np.abs(angle) < 1e-6 # (N, 1) + safe_angle = np.where(small, 1.0, angle) + sin_half_over_angle = np.where( + small, + 0.5 - angle**2 / 48.0, + np.sin(half_angle) / safe_angle, + ) + axis_angle: np.ndarray = xyz / sin_half_over_angle + return axis_angle + + +def np_quat_angular_velocity(q: np.ndarray, dt: float) -> np.ndarray: + """Estimate angular velocity from a quaternion time sequence using shortest-arc diffs.""" + if q.ndim != 2 or q.shape[1] != 4: + raise ValueError(f"Expected quaternion sequence with shape (T, 4), got {q.shape}") + if dt <= 0.0: + raise ValueError(f"dt must be positive, got {dt}") + + rotations = np_quat_ensure_continuity(q) + num_frames = rotations.shape[0] + omega = np.zeros((num_frames, 3), dtype=rotations.dtype) + if num_frames <= 1: + return omega + + if num_frames == 2: + q_rel = np_quat_mul(rotations[1], np_quat_conjugate(rotations[0])) + q_rel = np_quat_canonicalize(q_rel) + angvel = np_quat_to_axis_angle(q_rel[None, :])[0] / dt + omega[:] = angvel + return omega + + q_prev = rotations[:-2] + q_next = rotations[2:] + q_rel = np_quat_mul(q_next, np_quat_conjugate(q_prev)) + q_rel = np_quat_canonicalize(q_rel) + omega[1:-1] = np_quat_to_axis_angle(q_rel) / (2.0 * dt) + omega[0] = omega[1] + omega[-1] = omega[-2] + return omega + + +def np_yaw_to_quat(yaw: np.ndarray) -> np.ndarray: + """Convert yaw batch (N,) to quaternion batch (N, 4) in NumPy.""" + half = 0.5 * yaw + return np.stack( + [ + np.cos(half), + np.zeros_like(half), + np.zeros_like(half), + np.sin(half), + ], + axis=1, + ) + + +def np_quat_inv(q: np.ndarray) -> np.ndarray: + """Inverse of unit quaternions (N, 4) or (4,), w-first.""" + return np_quat_conjugate(q) + + +def np_quat_apply(q: np.ndarray, v: np.ndarray) -> np.ndarray: + """Rotate vector(s) by quaternion(s), supports batched/scalar inputs.""" + q_was_1d = q.ndim == 1 + v_was_1d = v.ndim == 1 + + if q_was_1d: + q = q[None, :] + if v_was_1d: + v = v[None, :] + + if q.shape[0] == 1 and v.shape[0] > 1: + q = np.broadcast_to(q, (v.shape[0], 4)) + elif v.shape[0] == 1 and q.shape[0] > 1: + v = np.broadcast_to(v, (q.shape[0], 3)) + + w, x, y, z = q[:, 0], q[:, 1], q[:, 2], q[:, 3] + vx, vy, vz = v[:, 0], v[:, 1], v[:, 2] + + t = 2 * np.stack( + [ + y * vz - z * vy, + z * vx - x * vz, + x * vy - y * vx, + ], + axis=1, + ) + t += 2 * w[:, None] * v + + result = v + np.stack( + [ + y * t[:, 2] - z * t[:, 1], + z * t[:, 0] - x * t[:, 2], + x * t[:, 1] - y * t[:, 0], + ], + axis=1, + ) + + rotated: np.ndarray = result[0] if q_was_1d and v_was_1d else result + return rotated + + +def np_quat_apply_inverse(q: np.ndarray, v: np.ndarray) -> np.ndarray: + """Rotate vector(s) by inverse quaternion(s).""" + return np_quat_apply(np_quat_inv(q), v) + + +def np_quat_error_magnitude(q1: np.ndarray, q2: np.ndarray) -> np.ndarray: + """Angular error magnitude between quaternions (N,) or scalar.""" + q1_was_1d = q1.ndim == 1 + q2_was_1d = q2.ndim == 1 + + if q1_was_1d: + q1 = q1[None, :] + if q2_was_1d: + q2 = q2[None, :] + + if q1.shape[0] == 1 and q2.shape[0] > 1: + q1 = np.broadcast_to(q1, q2.shape) + elif q2.shape[0] == 1 and q1.shape[0] > 1: + q2 = np.broadcast_to(q2, q1.shape) + + # Relative rotation from q1 to q2. + q_rel = np_quat_mul(q2, np_quat_inv(q1)) + q_rel = np_quat_canonicalize(q_rel) + + # Use atan2-based angle extraction for better numerical behavior. + xyz_norm = np.linalg.norm(q_rel[:, 1:], axis=1) + w = np.clip(q_rel[:, 0], -1.0, 1.0) + error = 2.0 * np.arctan2(xyz_norm, w) + magnitude: np.ndarray = error[0] if q1_was_1d and q2_was_1d else error + return magnitude + + +def np_quat_from_euler_xyz(roll: np.ndarray, pitch: np.ndarray, yaw: np.ndarray) -> np.ndarray: + """Convert Euler angles (XYZ) to quaternions (N, 4) or (4,), w-first.""" + roll = np.atleast_1d(roll) + pitch = np.atleast_1d(pitch) + yaw = np.atleast_1d(yaw) + squeeze = roll.shape[0] == 1 + + cr = np.cos(roll * 0.5) + sr = np.sin(roll * 0.5) + cp = np.cos(pitch * 0.5) + sp = np.sin(pitch * 0.5) + cy = np.cos(yaw * 0.5) + sy = np.sin(yaw * 0.5) + + w = cr * cp * cy + sr * sp * sy + x = sr * cp * cy - cr * sp * sy + y = cr * sp * cy + sr * cp * sy + z = cr * cp * sy - sr * sp * cy + + result = np.stack([w, x, y, z], axis=1) + return result[0] if squeeze else result + + +def np_yaw_quat(q: np.ndarray) -> np.ndarray: + """Extract yaw-only quaternion from full quaternion(s), w-first.""" + q_was_1d = q.ndim == 1 + if q_was_1d: + q = q[None, :] + + w, x, y, z = q[:, 0], q[:, 1], q[:, 2], q[:, 3] + yaw = np.arctan2(2 * (w * z + x * y), 1 - 2 * (y * y + z * z)) + + half_yaw = yaw * 0.5 + result = np.stack( + [ + np.cos(half_yaw), + np.zeros_like(half_yaw), + np.zeros_like(half_yaw), + np.sin(half_yaw), + ], + axis=1, + ) + + return result[0] if q_was_1d else result + + +def np_matrix_from_quat(q: np.ndarray) -> np.ndarray: + """Convert quaternion(s) to rotation matrix (N, 3, 3) or (3, 3), w-first.""" + q_was_1d = q.ndim == 1 + if q_was_1d: + q = q[None, :] + + w, x, y, z = q[:, 0], q[:, 1], q[:, 2], q[:, 3] + + xx = x * x + yy = y * y + zz = z * z + xy = x * y + xz = x * z + yz = y * z + wx = w * x + wy = w * y + wz = w * z + + result = np.stack( + [ + np.stack([1 - 2 * (yy + zz), 2 * (xy - wz), 2 * (xz + wy)], axis=1), + np.stack([2 * (xy + wz), 1 - 2 * (xx + zz), 2 * (yz - wx)], axis=1), + np.stack([2 * (xz - wy), 2 * (yz + wx), 1 - 2 * (xx + yy)], axis=1), + ], + axis=1, + ) + + return result[0] if q_was_1d else result + + +def np_subtract_frame_transforms( + pos1: np.ndarray, quat1: np.ndarray, pos2: np.ndarray, quat2: np.ndarray +) -> tuple[np.ndarray, np.ndarray]: + """Compute relative transform from frame 1 to frame 2 in frame-1 coordinates.""" + rel_pos = np_quat_apply_inverse(quat1, pos2 - pos1) + rel_quat = np_quat_mul(np_quat_inv(quat1), quat2) + return rel_pos, rel_quat diff --git a/src/unilab/envs/locomotion/common/dr_provider.py b/src/unilab/envs/locomotion/common/dr_provider.py index dc17caa20..d77e50ba7 100644 --- a/src/unilab/envs/locomotion/common/dr_provider.py +++ b/src/unilab/envs/locomotion/common/dr_provider.py @@ -25,7 +25,7 @@ validate_interval_push_support, zero_actions, ) -from unilab.utils.math_utils import np_quat_mul, np_yaw_to_quat +from unilab.envs.common.rotation import np_quat_mul, np_yaw_to_quat class LocomotionDRProvider(DomainRandomizationProvider): diff --git a/src/unilab/envs/locomotion/go1/joystick.py b/src/unilab/envs/locomotion/go1/joystick.py index 06882f55d..cdb56fdb0 100644 --- a/src/unilab/envs/locomotion/go1/joystick.py +++ b/src/unilab/envs/locomotion/go1/joystick.py @@ -11,13 +11,13 @@ 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.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/manipulation/inhand_rot_allegro/rotation.py b/src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py index 71686bb91..ab3826c0e 100644 --- a/src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py +++ b/src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py @@ -26,8 +26,8 @@ validate_interval_push_support, zero_actions, ) +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..bc5289afb 100644 --- a/src/unilab/envs/manipulation/sharpa_inhand/base.py +++ b/src/unilab/envs/manipulation/sharpa_inhand/base.py @@ -12,7 +12,7 @@ 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.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..55f3c81d1 100644 --- a/src/unilab/envs/manipulation/sharpa_inhand/rotation.py +++ b/src/unilab/envs/manipulation/sharpa_inhand/rotation.py @@ -18,6 +18,7 @@ ResetPlan, ) from unilab.dr.dr_utils import build_common_reset_randomization, validate_common_reset_randomization +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..8caad0cfc 100644 --- a/src/unilab/envs/motion_tracking/g1/tracking.py +++ b/src/unilab/envs/motion_tracking/g1/tracking.py @@ -25,18 +25,18 @@ validate_interval_push_support, zero_actions, ) -from unilab.envs.locomotion.g1.base import G1BaseCfg, G1BaseEnv -from unilab.utils.math_utils import ( +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/training/__init__.py b/src/unilab/training/__init__.py index 1726b7fcc..d41e0c7c8 100644 --- a/src/unilab/training/__init__.py +++ b/src/unilab/training/__init__.py @@ -5,20 +5,28 @@ assert_offpolicy_task_choice_matches_algo, create_env, ensure_registries, - get_entrypoint_log_root, get_hydra_runtime_choice, + setup_logger, +) +from unilab.training.logging import ExperimentTracker, OffPolicyLogger, OnPolicyLogger +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, ) +from unilab.visualization.playback import render_play_mode __all__ = [ "BackendAdapter", + "ExperimentTracker", + "HardwareMonitor", + "OffPolicyLogger", + "OnPolicyLogger", "assert_offpolicy_task_choice_matches_algo", "create_env", "ensure_registries", diff --git a/src/unilab/training/backend_adapter.py b/src/unilab/training/backend_adapter.py index 18008750f..21baea8e3 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.config.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/training/logging/__init__.py b/src/unilab/training/logging/__init__.py new file mode 100644 index 000000000..a420aa4e5 --- /dev/null +++ b/src/unilab/training/logging/__init__.py @@ -0,0 +1,23 @@ +"""Training logging and experiment tracking helpers.""" + +from unilab.training.logging.experiment import ( + ExperimentTracker, + build_wandb_run_name, + build_wandb_settings, + get_device_info_dict, + get_git_info, + patch_rsl_rl_wandb_writer, +) +from unilab.training.logging.offpolicy import OffPolicyLogger +from unilab.training.logging.onpolicy import OnPolicyLogger + +__all__ = [ + "ExperimentTracker", + "OffPolicyLogger", + "OnPolicyLogger", + "build_wandb_run_name", + "build_wandb_settings", + "get_device_info_dict", + "get_git_info", + "patch_rsl_rl_wandb_writer", +] diff --git a/src/unilab/training/logging/common.py b/src/unilab/training/logging/common.py new file mode 100644 index 000000000..be017b7a5 --- /dev/null +++ b/src/unilab/training/logging/common.py @@ -0,0 +1,316 @@ +from __future__ import annotations + +import importlib +import os +import time +from collections import deque +from typing import Any + +from rich import box +from rich.console import Console +from rich.live import Live +from rich.panel import Panel +from rich.table import Table +from rich.text import Text + + +def _fmt_time(seconds: float) -> str: + if seconds < 60: + return f"{seconds:.0f}s" + m, s = divmod(int(seconds), 60) + if m < 60: + return f"{m}m{s:02d}s" + h, m = divmod(m, 60) + return f"{h}h{m:02d}m{s:02d}s" + + +def _fmt_number(v: float) -> str: + if abs(v) == 0: + return "0" + if abs(v) >= 1e6: + return f"{v:.2e}" + if abs(v) >= 100: + return f"{v:.1f}" + if abs(v) >= 1: + return f"{v:.3f}" + if abs(v) >= 0.001: + return f"{v:.4f}" + return f"{v:.2e}" + + +def _load_wandb() -> Any | None: + """Load wandb lazily so it remains an optional dependency.""" + try: + return importlib.import_module("wandb") + except ImportError: + return None + + +class BaseTrainingLogger: + """Shared lifecycle and backend logging setup for rich training loggers.""" + + def __init__( + self, + *, + algo_name: str, + max_iterations: int, + num_envs: int, + env_name: str, + log_dir: str, + log_backend: str, + wandb_project: str, + wandb_entity: str | None, + wandb_name: str, + wandb_group: list[str] | None | str, + wandb_job_type: str | None, + wandb_tags: list[str] | None, + wandb_notes: str | None, + refresh_per_second: int = 4, + tensorboard_subdir: str | None = "tb", + wandb_config: dict[str, Any] | None = None, + ): + self.algo_name = algo_name + self.max_iterations = max_iterations + self.num_envs = num_envs + self.env_name = env_name + + self._no_print = log_backend.lower() == "no_print" + self._log_backend = "none" if self._no_print else log_backend.lower() + + self._console = Console() + self._live: Live | None = None + self._refresh_rate = refresh_per_second + + self._start_time: float = 0.0 + self._iteration: int = 0 + self._reward_history: deque[float] = deque(maxlen=200) + self._latest_metrics: dict[str, float] = {} + self._latest_reward_components: dict[str, float] = {} + self._collect_time: float = 0.0 + self._train_time: float = 0.0 + self._mean_ep_length: float = 0.0 + self._last_save: str = "" + self._status: str = "" + + self._log_dir = log_dir + self._tb_writer: Any | None = None + self._wandb_run = None + self._owns_wandb_run = False + + if self._log_backend == "tensorboard" and log_dir: + self._init_tensorboard(log_dir, tensorboard_subdir) + elif self._log_backend == "wandb": + self._init_wandb( + project=wandb_project, + entity=wandb_entity, + name=wandb_name or f"{algo_name}_{env_name}", + log_dir=log_dir, + group=wandb_group, + job_type=wandb_job_type, + tags=wandb_tags, + notes=wandb_notes, + extra_config=wandb_config, + ) + + def _format_tensorboard_message(self, tb_dir: str) -> str: + return f"[dim]TensorBoard: {tb_dir}[/]" + + def _format_wandb_message(self, project: str, name: str) -> str: + return f"[dim]W&B: {project}/{name}[/]" + + def _init_tensorboard(self, log_dir: str, subdir: str | None): + try: + from torch.utils.tensorboard import SummaryWriter + + tb_dir = log_dir if subdir is None else os.path.join(log_dir, subdir) + os.makedirs(tb_dir, exist_ok=True) + self._tb_writer = SummaryWriter(log_dir=tb_dir) + if not self._no_print: + self._console.print(self._format_tensorboard_message(tb_dir)) + except ImportError: + if not self._no_print: + self._console.print("[yellow]tensorboard not installed[/]") + + def _init_wandb( + self, + *, + project: str, + entity: str | None, + name: str, + log_dir: str, + group: str | None | list[str], + job_type: str | None, + tags: list[str] | None, + notes: str | None, + extra_config: dict[str, Any] | None = None, + ): + wandb = _load_wandb() + if wandb is None: + if not self._no_print: + self._console.print("[yellow]wandb not installed[/]") + return + + self._wandb_run = wandb.run + if self._wandb_run is None: + config: dict[str, Any] = { + "algo": self.algo_name, + "env": self.env_name, + "num_envs": self.num_envs, + } + if extra_config: + config.update(extra_config) + + kwargs: dict[str, Any] = { + "project": project, + "name": name, + "config": config, + "dir": log_dir or None, + "reinit": True, + } + if entity: + kwargs["entity"] = entity + if group: + kwargs["group"] = group + if job_type: + kwargs["job_type"] = job_type + if tags: + kwargs["tags"] = tags + if notes: + kwargs["notes"] = notes + + self._wandb_run = wandb.init(**kwargs) + self._owns_wandb_run = True + + if not self._no_print: + self._console.print(self._format_wandb_message(project, name)) + + def start(self, *, status: str = ""): + self._start_time = time.time() + self._status = status + if not self._no_print: + self._live = Live( + self._build_display(), + console=self._console, + refresh_per_second=self._refresh_rate, + transient=False, + ) + self._live.start() + + def finish(self, *, title: str = "Training Summary", extra_summary: str = ""): + if self._live is not None: + self._live.update(self._build_display()) + self._live.stop() + self._live = None + + elapsed = time.time() - self._start_time + if not self._no_print: + summary = ( + f"[bold green]Training complete[/]\n" + f" Algo: [cyan]{self.algo_name}[/] | Env: [cyan]{self.env_name}[/]\n" + f" Iterations: [yellow]{self._iteration}[/]/{self.max_iterations}\n" + f" Total time: [yellow]{_fmt_time(elapsed)}[/]\n" + ) + if extra_summary: + summary += extra_summary + if self._last_save: + summary += f" Last checkpoint: [dim]{self._last_save}[/]" + + self._console.print() + self._console.print(Panel(summary, title=f"[bold]{title}[/]", border_style="green")) + + if self._tb_writer: + self._tb_writer.close() + if self._wandb_run and self._owns_wandb_run: + wandb = _load_wandb() + if wandb is not None: + wandb.finish() + + def update_ep_length(self, length: float): + self._mean_ep_length = length + + def log_save(self, path: str): + self._last_save = path + self._refresh() + + def _refresh(self): + if self._live is not None: + self._live.update(self._build_display()) + + def _estimate_eta(self) -> str: + if self._iteration <= 0: + return "" + elapsed = time.time() - self._start_time + remaining = self.max_iterations - self._iteration + avg_iter = elapsed / self._iteration + eta_s = remaining * avg_iter + return _fmt_time(eta_s) + + def _build_header(self, *, include_status: bool) -> Panel: + elapsed = time.time() - self._start_time if self._start_time else 0 + eta = self._estimate_eta() + + header_text = Text() + header_text.append(f" {self.algo_name}", style="bold cyan") + header_text.append(" │ ", style="dim") + header_text.append(f"{self.env_name}", style="bold white") + header_text.append(" │ ", style="dim") + header_text.append(f"iter {self._iteration}/{self.max_iterations}", style="yellow") + header_text.append(" │ ", style="dim") + header_text.append(f"⏱ {_fmt_time(elapsed)}", style="green") + if eta: + header_text.append(" │ ETA ", style="dim") + header_text.append(eta, style="bold magenta") + if include_status and self._status: + header_text.append(" │ ", style="dim") + header_text.append(self._status, style="dim italic") + + return Panel(header_text, style="dim", box=box.SIMPLE) + + def _build_reward_table_common(self, *, wait_message: str) -> Table: + table = Table( + title="[bold]Rewards[/]", + box=box.SIMPLE_HEAVY, + show_header=True, + header_style="bold green", + expand=True, + pad_edge=False, + ) + table.add_column("Component", style="white", ratio=2) + table.add_column("Value", justify="right", ratio=1) + + if self._reward_history: + recent = list(self._reward_history) + mean_rew = sum(recent[-50:]) / max(len(recent[-50:]), 1) + peak_rew = max(recent) if recent else 0 + + if len(recent) >= 10: + old = sum(recent[-20:-10]) / 10 + new = sum(recent[-10:]) / 10 + trend = ( + "[green]▲[/]" + if new > old * 1.05 + else "[red]▼[/]" + if new < old * 0.95 + else "[yellow]━[/]" + ) + else: + trend = "" + + table.add_row(f"[bold]Mean Reward[/] {trend}", f"[bold green]{mean_rew:.3f}[/]") + table.add_row(" Peak", f"[dim]{peak_rew:.3f}[/]") + if self._mean_ep_length > 0: + table.add_row(" Ep Len", f"[dim]{self._mean_ep_length:.1f}[/]") + table.add_row("", "") + else: + table.add_row(wait_message, "") + + if self._latest_reward_components: + for name, val in sorted(self._latest_reward_components.items()): + display = name.replace("reward/", "").replace("_", " ") + color = "green" if val > 0 else "red" if val < 0 else "dim" + table.add_row(f" {display}", f"[{color}]{val:+.4f}[/]") + + return table + + def _build_display(self) -> Panel: + raise NotImplementedError diff --git a/src/unilab/training/logging/experiment.py b/src/unilab/training/logging/experiment.py new file mode 100644 index 000000000..f809ca2dc --- /dev/null +++ b/src/unilab/training/logging/experiment.py @@ -0,0 +1,425 @@ +"""Shared experiment tracking utilities for local files and W&B.""" + +from __future__ import annotations + +import dataclasses +import getpass +import importlib +import importlib.util +import json +import os +import socket +import subprocess +import time +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from omegaconf import OmegaConf + + +def _cfg_get(cfg: Any, key: str, default: Any = None) -> Any: + if cfg is None: + return default + if isinstance(cfg, dict): + return cfg.get(key, default) + return getattr(cfg, key, default) + + +def _plain_dict(value: Any) -> Any: + if OmegaConf.is_config(value): + return OmegaConf.to_container(value, resolve=True) + if dataclasses.is_dataclass(value) and not isinstance(value, type): + return dataclasses.asdict(value) + return value + + +def _load_wandb() -> Any | None: + try: + return importlib.import_module("wandb") + except ImportError: + return None + + +def _json_safe(value: Any) -> Any: + if isinstance(value, dict): + return {str(k): _json_safe(v) for k, v in value.items()} + if isinstance(value, (list, tuple)): + return [_json_safe(v) for v in value] + if isinstance(value, Path): + return str(value) + if isinstance(value, (str, int, float, bool)) or value is None: + return value + try: + json.dumps(value) + return value + except TypeError: + return str(value) + + +def get_device_info_dict() -> dict[str, str]: + try: + 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: + return {"platform": os.uname().sysname if hasattr(os, "uname") else "unknown"} + + +def get_git_info(root_dir: str | Path) -> dict[str, Any]: + root = Path(root_dir) + + def _run_git(*args: str) -> str | None: + try: + result = subprocess.run( + ["git", *args], + cwd=root, + check=True, + capture_output=True, + text=True, + ) + except Exception: + return None + return result.stdout.strip() + + commit = _run_git("rev-parse", "HEAD") + branch = _run_git("rev-parse", "--abbrev-ref", "HEAD") + status = _run_git("status", "--short") + + return { + "commit": commit, + "branch": branch, + "dirty": bool(status), + } + + +def build_wandb_run_name(algo_name: str, task_name: str, log_dir: str | Path | None) -> str: + if log_dir is None: + return f"{algo_name}__{task_name}" + run_dir = Path(log_dir) + return f"{algo_name}__{task_name}__{run_dir.name}" + + +def build_wandb_settings( + training_cfg: Any, + *, + algo_name: str, + task_name: str, + sim_backend: str, + log_dir: str | Path | None, +) -> dict[str, Any]: + name = _cfg_get(training_cfg, "wandb_name") + if not name: + name = build_wandb_run_name(algo_name, task_name, log_dir) + + group = _cfg_get(training_cfg, "wandb_group") + if not group: + group = task_name + + job_type = _cfg_get(training_cfg, "wandb_job_type") + if not job_type: + job_type = algo_name + + tags = [str(tag) for tag in (_cfg_get(training_cfg, "wandb_tags", []) or [])] + auto_tags = [algo_name, task_name, sim_backend, f"user-{getpass.getuser()}"] + for tag in auto_tags: + if tag not in tags: + tags.append(tag) + + return { + "project": _cfg_get(training_cfg, "wandb_project", "unilab"), + "entity": _cfg_get(training_cfg, "wandb_entity"), + "name": name, + "group": group, + "job_type": job_type, + "tags": tags, + "notes": _cfg_get(training_cfg, "wandb_notes"), + "mode": _cfg_get(training_cfg, "wandb_mode"), + } + + +class ExperimentTracker: + """Tracks experiment metadata locally and optionally in Weights & Biases.""" + + def __init__( + self, + *, + root_dir: str | Path, + log_dir: str | Path, + algo_name: str, + task_name: str, + sim_backend: str, + training_cfg: Any, + full_cfg: Any, + device: str | None = None, + collector_device: str | None = None, + ): + self.root_dir = Path(root_dir) + self.log_dir = Path(log_dir) + self.algo_name = algo_name + self.task_name = task_name + self.sim_backend = sim_backend + self.training_cfg = training_cfg + self.full_cfg = full_cfg + self.device = device + self.collector_device = collector_device + self.enabled = str(_cfg_get(training_cfg, "logger", "tensorboard")).lower() == "wandb" + + self.log_dir.mkdir(parents=True, exist_ok=True) + + self._wandb = None + self._run = None + self._owns_run = False + self._started = False + self._start_monotonic = 0.0 + self._start_utc = "" + self._summary: dict[str, Any] = {} + + @property + def run(self) -> Any | None: + return self._run + + @property + def run_url(self) -> str | None: + return getattr(self._run, "url", None) if self._run is not None else None + + @property + def wandb_settings(self) -> dict[str, Any]: + return build_wandb_settings( + self.training_cfg, + algo_name=self.algo_name, + task_name=self.task_name, + sim_backend=self.sim_backend, + log_dir=self.log_dir, + ) + + def start(self) -> None: + if self._started: + return + + self._started = True + self._start_monotonic = time.perf_counter() + self._start_utc = datetime.now(timezone.utc).isoformat() + + metadata = { + "algo": self.algo_name, + "task": self.task_name, + "sim_backend": self.sim_backend, + "device": self.device, + "collector_device": self.collector_device, + "log_dir": str(self.log_dir), + "start_time_utc": self._start_utc, + "hostname": socket.gethostname(), + "user": getpass.getuser(), + "git": get_git_info(self.root_dir), + "hardware": get_device_info_dict(), + "wandb": self.wandb_settings, + } + + payload = { + "run": _json_safe(metadata), + "config": _json_safe(_plain_dict(self.full_cfg)), + } + self._write_json(self.log_dir / "run_config.json", payload) + + if not self.enabled: + return + + self._wandb = _load_wandb() + if self._wandb is None: + print("[experiment_tracking] wandb not installed, skipping W&B experiment tracking.") + return + + self._run = self._wandb.run + if self._run is None: + kwargs = { + "project": self.wandb_settings["project"], + "name": self.wandb_settings["name"], + "config": payload, + "dir": str(self.log_dir), + "reinit": True, + } + for key in ("entity", "group", "job_type", "tags", "notes", "mode"): + value = self.wandb_settings.get(key) + if value not in (None, "", []): + kwargs[key] = value + self._run = self._wandb.init(**kwargs) + self._owns_run = True + else: + self._run.config.update(payload, allow_val_change=True) + + if self._run is not None: + self._run.summary["algo"] = self.algo_name + self._run.summary["task"] = self.task_name + self._run.summary["sim_backend"] = self.sim_backend + if self.device: + self._run.summary["device"] = self.device + if self.collector_device: + self._run.summary["collector_device"] = self.collector_device + self._run.summary["log_dir"] = str(self.log_dir) + + def update_summary(self, summary: dict[str, Any] | None = None) -> None: + if summary: + self._summary.update(summary) + + if not self._started: + return + + wall_time_sec = time.perf_counter() - self._start_monotonic + payload = { + **self._summary, + "algo": self.algo_name, + "task": self.task_name, + "sim_backend": self.sim_backend, + "log_dir": str(self.log_dir), + "start_time_utc": self._start_utc, + "end_time_utc": datetime.now(timezone.utc).isoformat(), + "wall_time_sec": wall_time_sec, + "wandb_run_url": self.run_url, + } + self._write_json(self.log_dir / "run_summary.json", _json_safe(payload)) + + if self._run is not None: + for key, value in payload.items(): + self._run.summary[key] = _json_safe(value) + + def log_video(self, video_path: str | Path | None, key: str = "media/play_video") -> None: + if video_path is None: + return + + video = Path(video_path) + if not video.exists(): + return + + self._summary["play_video_path"] = str(video) + if self._run is not None and self._wandb is not None: + self._wandb.log({key: self._wandb.Video(str(video), format="mp4")}) + + def finish(self) -> None: + if not self._started: + return + + self.update_summary() + if self._run is not None and self._wandb is not None and self._owns_run: + self._wandb.finish() + self._run = None + self._wandb = None + + @staticmethod + def _write_json(path: Path, payload: Any) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(payload, indent=2), encoding="utf-8") + + +def patch_rsl_rl_wandb_writer() -> None: + """Patch rsl-rl W&B writer so it can reuse an already-open run.""" + try: + import rsl_rl.utils.wandb_utils as wandb_utils + except Exception: + return + + if getattr(wandb_utils, "_UNILAB_PATCHED", False): + return + + wandb = _load_wandb() + if wandb is None: + return + wandb_mod = wandb + + from torch.utils.tensorboard import SummaryWriter as TensorboardSummaryWriter + + class PatchedWandbSummaryWriter(TensorboardSummaryWriter): + def __init__(self, log_dir: str, flush_secs: int, cfg: dict) -> None: + super().__init__(log_dir, flush_secs=flush_secs) + + run_name = os.path.split(log_dir)[-1] + project = cfg.get("wandb_project", "unilab") + entity = cfg.get("wandb_entity") or os.environ.get("WANDB_USERNAME") + group = cfg.get("wandb_group") + job_type = cfg.get("wandb_job_type") + tags = cfg.get("wandb_tags") + notes = cfg.get("wandb_notes") + mode = cfg.get("wandb_mode") + + self.logged_videos: set[str] = set() + self._owns_run = wandb_mod.run is None + if self._owns_run: + kwargs = { + "project": project, + "name": run_name, + "config": {"log_dir": log_dir}, + "settings": wandb_mod.Settings(start_method="thread"), + } + if entity: + kwargs["entity"] = entity + if group: + kwargs["group"] = group + if job_type: + kwargs["job_type"] = job_type + if tags: + kwargs["tags"] = tags + if notes: + kwargs["notes"] = notes + if mode: + kwargs["mode"] = mode + wandb_mod.init(**kwargs) + else: + wandb_mod.config.update({"log_dir": log_dir}, allow_val_change=True) + + def store_config(self, env_cfg: dict | object, train_cfg: dict) -> None: + wandb_mod.config.update({"train_cfg": train_cfg}, allow_val_change=True) + env_payload: Any + if isinstance(env_cfg, dict): + env_payload = env_cfg + elif dataclasses.is_dataclass(env_cfg) and not isinstance(env_cfg, type): + env_payload = dataclasses.asdict(env_cfg) + elif hasattr(env_cfg, "to_dict"): + env_payload = env_cfg.to_dict() # type: ignore[union-attr] + else: + env_payload = str(env_cfg) + wandb_mod.config.update({"env_cfg": env_payload}, allow_val_change=True) + + def add_scalar( + self, + tag: Any, + scalar_value: Any, + global_step: Any = None, + walltime: Any = None, + new_style: Any = False, + double_precision: Any = False, + ) -> None: + super().add_scalar( + tag, + scalar_value, + global_step=global_step, + walltime=walltime, + new_style=new_style, + double_precision=double_precision, + ) + wandb_mod.log({tag: scalar_value}, step=global_step) + + def stop(self) -> None: + if self._owns_run: + wandb_mod.finish() + + def save_model(self, model_path: str, it: int) -> None: + wandb_mod.save(model_path, base_path=os.path.dirname(model_path)) + + def save_file(self, path: str) -> None: + wandb_mod.save(path, base_path=os.path.dirname(path)) + + def save_video(self, video: Path, it: int) -> None: + if video.name not in self.logged_videos: + wandb_mod.log({"video": wandb_mod.Video(str(video), format="mp4")}, step=it) + self.logged_videos.add(video.name) + + wandb_utils.WandbSummaryWriter = PatchedWandbSummaryWriter + setattr(wandb_utils, "_UNILAB_PATCHED", True) diff --git a/src/unilab/training/logging/offpolicy.py b/src/unilab/training/logging/offpolicy.py new file mode 100644 index 000000000..0e9bc6f51 --- /dev/null +++ b/src/unilab/training/logging/offpolicy.py @@ -0,0 +1,347 @@ +"""Rich-based training logger for off-policy RL algorithms (SAC, TD3, etc).""" + +from __future__ import annotations + +import time +from collections import deque +from typing import Any + +from rich import box +from rich.console import Group +from rich.panel import Panel +from rich.table import Table + +from unilab.training.logging.common import BaseTrainingLogger, _fmt_number, _fmt_time, _load_wandb + + +class OffPolicyLogger(BaseTrainingLogger): + """Rich logger for off-policy RL algorithms (SAC, TD3, etc).""" + + def __init__( + self, + algo_name: str = "RL", + max_iterations: int = 1500, + num_envs: int = 4096, + env_name: str = "", + obs_dim: int = 0, + action_dim: int = 0, + refresh_per_second: int = 4, + log_dir: str = "", + log_backend: str = "tensorboard", + wandb_project: str = "unilab", + wandb_entity: str | None = None, + wandb_name: str = "", + wandb_group: str | None = None, + wandb_job_type: str | None = None, + wandb_tags: list[str] | None = None, + wandb_notes: str | None = None, + ): + super().__init__( + algo_name=algo_name, + max_iterations=max_iterations, + num_envs=num_envs, + env_name=env_name, + log_dir=log_dir, + log_backend=log_backend, + wandb_project=wandb_project, + wandb_entity=wandb_entity, + wandb_name=wandb_name, + wandb_group=wandb_group, + wandb_job_type=wandb_job_type, + wandb_tags=wandb_tags, + wandb_notes=wandb_notes, + refresh_per_second=refresh_per_second, + tensorboard_subdir=None, + wandb_config={ + "obs_dim": obs_dim, + "action_dim": action_dim, + "max_iterations": max_iterations, + }, + ) + 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 + self._wait_time: float = 0.0 + self._iter_times: deque = deque(maxlen=50) + self._collector_timing: dict[str, float] = {} + self._timeout_rate: float = 0.0 + self._terminated_rate: float = 0.0 + self._buffer_utilization: float = 0.0 + self._sync_collection: bool = False + self._env_steps_per_sync: int = 0 + self._replay_queue_len: int = 0 + self._replay_queue_max: int = 0 + self._status: str = "Initializing..." + + def _format_tensorboard_message(self, tb_dir: str) -> str: + return f"[dim]TensorBoard logging to: {tb_dir}[/]" + + 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..."): + super().start(status=status) + + def finish(self, *, title: str = "Training Summary", extra_summary: str = ""): + super().finish( + title=title, + extra_summary=f" Total env steps: [yellow]{self._total_steps:,}[/]\n{extra_summary}", + ) + + def log_buffer_fill(self, current: int, target: int): + self._buffer_size = current + self._buffer_target = target + pct = current / max(target, 1) * 100 + self._status = f"Buffer fill: {current:,}/{target:,} ({pct:.0f}%)" + self._refresh() + + def update_collector_timing(self, timing_ms: dict[str, float]): + self._collector_timing.update(timing_ms) + + def update_done_rates(self, timeout_rate: float, terminated_rate: float): + self._timeout_rate = float(timeout_rate) + self._terminated_rate = float(terminated_rate) + + def update_buffer_utilization(self, utilization: float): + self._buffer_utilization = float(utilization) + + def update_replay_queue(self, current_len: int, max_size: int): + 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): + 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): + self._total_steps = total_steps + self._buffer_size = buffer_size + if mean_reward != 0: + self._reward_history.append(mean_reward) + self._refresh() + + def log_step( + self, + iteration: int, + metrics: dict[str, float] | None = None, + reward: float | None = None, + reward_components: dict[str, float] | None = None, + collect_time: float = 0.0, + train_time: float = 0.0, + wait_time: float = 0.0, + extra_info: dict | None = None, + ): + 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() + self._backend_log_step( + iteration, metrics, reward, reward_components, collect_time, train_time + ) + + def _backend_log_step( + self, + iteration: int, + metrics: dict[str, float] | None, + reward: float | None, + reward_components: dict[str, float] | None, + collect_time: float, + train_time: float, + ): + global_step = self._total_steps if self._total_steps > 0 else iteration + elapsed = time.time() - self._start_time if self._start_time else 0 + + if self._tb_writer: + writer = self._tb_writer + if metrics: + for key, value in metrics.items(): + writer.add_scalar(f"train/{key}", value, global_step) + if reward is not None: + writer.add_scalar("reward/mean", reward, global_step) + if reward_components: + for key, value in reward_components.items(): + writer.add_scalar(f"reward/{key}", value, global_step) + if self._mean_ep_length > 0: + 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: + 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 + ) + writer.add_scalar( + "perf/collect_train_ratio", + self._collect_time / max(self._train_time, 1e-6), + global_step, + ) + + if self._wandb_run: + wandb = _load_wandb() + if wandb is None: + return + log_dict: dict[str, Any] = {"iteration": iteration} + if metrics: + 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 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 + 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, 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): + self._status = status + self._refresh() + + def _build_display(self) -> Panel: + header_panel = self._build_header(include_status=True) + 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) + return Panel( + Group(header_panel, grid, bottom), + title="[bold] 🚀 UniLab Off-Policy Training [/]", + border_style="bright_blue", + padding=(0, 1), + ) + + def _build_metrics_table(self) -> Table: + table = Table( + title="[bold]Losses & Metrics[/]", + box=box.SIMPLE_HEAVY, + show_header=True, + header_style="bold cyan", + expand=True, + pad_edge=False, + ) + 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: + 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: + return self._build_reward_table_common(wait_message="[dim]Waiting for data...[/]") + + def _build_timing_table(self) -> Table: + table = Table( + title="[bold]Timing & System[/]", + box=box.SIMPLE_HEAVY, + show_header=True, + header_style="bold blue", + expand=True, + pad_edge=False, + ) + table.add_column("Item", style="white", ratio=2, no_wrap=True) + table.add_column("Value", style="yellow", justify="right", ratio=1, no_wrap=True) + table.add_column("Item", style="white", ratio=2, no_wrap=True) + 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_ms = self._wait_time * 1000 + wait_color = "red" if wait_ms > 1.0 else "yellow" + table.add_row( + "[dim]learner[/] Wait", + f"[{wait_color}]{wait_ms:.1f}ms[/]", + "[dim]learner[/] Train", + f"{self._train_time * 1000:.1f}ms", + ) + table.add_row( + "[dim]learner[/] Collect", + f"{self._collect_time * 1000:.1f}ms", + "", + "", + ) + timing_items = list(self._collector_timing.items()) + 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_value:.1f}ms", + f"[dim]collector[/] {right_key}", + f"{right_value:.1f}ms", + ) + else: + 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}%", + ) + 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: + utilization_str = f"[green]{utilization:.2f}[/]" + table.add_row("Write/Read", utilization_str, "", "") + table.add_row( + "Envs", + f"{self.num_envs:,}", + "Sync Collect", + f"{'✓' if self._sync_collection else '✗'} ({self._env_steps_per_sync})" + if self._sync_collection + else "✗", + ) + if self._replay_queue_max > 0: + replay_color = "green" if self._replay_queue_len < self._replay_queue_max else "yellow" + table.add_row( + "Replay Queue", + f"[{replay_color}]{self._replay_queue_len}/{self._replay_queue_max}[/]", + "", + "", + ) + if elapsed > 0 and self._total_steps > 0: + table.add_row("Steps/s", f"{self._total_steps / elapsed:,.0f}", "", "") + return table diff --git a/src/unilab/training/logging/onpolicy.py b/src/unilab/training/logging/onpolicy.py new file mode 100644 index 000000000..d166c2512 --- /dev/null +++ b/src/unilab/training/logging/onpolicy.py @@ -0,0 +1,196 @@ +from __future__ import annotations + +import time +from typing import Any + +from rich import box +from rich.console import Group +from rich.panel import Panel +from rich.table import Table + +from unilab.training.logging.common import BaseTrainingLogger, _fmt_number, _fmt_time, _load_wandb + + +class OnPolicyLogger(BaseTrainingLogger): + """Rich logger for on-policy RL (PPO, A2C, etc).""" + + def __init__( + self, + algo_name: str = "PPO", + max_iterations: int = 1500, + num_envs: int = 4096, + num_steps: int = 24, + env_name: str = "", + log_dir: str = "", + log_backend: str = "tensorboard", + wandb_project: str = "unilab", + wandb_entity: str | None = None, + wandb_name: str = "", + wandb_group: str | None = None, + wandb_job_type: str | None = None, + wandb_tags: list[str] | None = None, + wandb_notes: str | None = None, + ): + super().__init__( + algo_name=algo_name, + max_iterations=max_iterations, + num_envs=num_envs, + env_name=env_name, + log_dir=log_dir, + log_backend=log_backend, + wandb_project=wandb_project, + wandb_entity=wandb_entity, + wandb_name=wandb_name, + wandb_group=wandb_group, + wandb_job_type=wandb_job_type, + wandb_tags=wandb_tags, + wandb_notes=wandb_notes, + tensorboard_subdir="tb", + ) + self.num_steps = num_steps + + def start(self, *, status: str = ""): + super().start(status=status) + + def finish(self, *, title: str = "Training Summary", extra_summary: str = ""): + super().finish(title=title, extra_summary=extra_summary) + + def log_step( + self, + iteration: int, + metrics: dict[str, float] | None = None, + reward: float | None = None, + reward_components: dict[str, float] | None = None, + collect_time: float = 0.0, + train_time: float = 0.0, + ): + self._iteration = iteration + self._collect_time = collect_time + self._train_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._refresh() + self._backend_log_step(iteration, metrics, reward, reward_components) + + def _backend_log_step( + self, + iteration: int, + metrics: dict[str, float] | None, + reward: float | None, + reward_components: dict[str, float] | None, + ): + if self._tb_writer: + w = self._tb_writer + if metrics: + for k, v in metrics.items(): + w.add_scalar(f"train/{k}", v, iteration) + if reward is not None: + w.add_scalar("reward/mean", reward, iteration) + if reward_components: + for k, v in reward_components.items(): + w.add_scalar(f"reward/{k}", v, iteration) + if self._mean_ep_length > 0: + w.add_scalar("episode/length", self._mean_ep_length, iteration) + w.add_scalar("perf/collect_time_ms", self._collect_time * 1000, iteration) + w.add_scalar("perf/train_time_ms", self._train_time * 1000, iteration) + + 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 + 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 + if self._mean_ep_length > 0: + log_dict["episode/length"] = self._mean_ep_length + log_dict["perf/collect_time_ms"] = self._collect_time * 1000 + log_dict["perf/train_time_ms"] = self._train_time * 1000 + wandb.log(log_dict, step=iteration) + + def _build_display(self) -> Panel: + header_panel = self._build_header(include_status=False) + + 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, + title="[bold] 🚀 UniLab On-Policy Training [/]", + border_style="bright_blue", + padding=(0, 1), + ) + + def _build_metrics_table(self) -> Table: + table = Table( + title="[bold]Policy Metrics[/]", + box=box.SIMPLE_HEAVY, + show_header=True, + header_style="bold cyan", + expand=True, + pad_edge=False, + ) + 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...[/]", "") + else: + for k in sorted(self._latest_metrics.keys()): + v = self._latest_metrics[k] + name = k.replace("_", " ").title() + table.add_row(name, _fmt_number(v)) + + return table + + def _build_reward_table(self) -> Table: + return self._build_reward_table_common(wait_message="[dim]Waiting...[/]") + + def _build_timing_table(self) -> Table: + table = Table( + title="[bold]Timing[/]", + box=box.SIMPLE_HEAVY, + show_header=True, + header_style="bold blue", + expand=True, + pad_edge=False, + ) + table.add_column("Item", style="white", ratio=1) + table.add_column("Value", style="yellow", justify="right", ratio=1) + table.add_column("Item", style="white", ratio=1) + table.add_column("Value", style="yellow", justify="right", ratio=1) + + elapsed = time.time() - self._start_time if self._start_time else 0 + iter_time = self._collect_time + self._train_time + fps = int(self.num_envs * self.num_steps / max(iter_time, 1e-8)) if iter_time > 0 else 0 + + table.add_row("Elapsed", _fmt_time(elapsed), "Envs", f"{self.num_envs:,}") + table.add_row( + "Collect", + f"{self._collect_time * 1000:.1f}ms", + "Train", + f"{self._train_time * 1000:.1f}ms", + ) + table.add_row("Iter Time", f"{iter_time * 1000:.1f}ms", "Steps/s", f"{fps:,}") + + return table diff --git a/src/unilab/training/monitoring.py b/src/unilab/training/monitoring.py new file mode 100644 index 000000000..cafcd38a0 --- /dev/null +++ b/src/unilab/training/monitoring.py @@ -0,0 +1,60 @@ +"""Hardware monitoring utilities for performance profiling.""" + +from typing import Dict + +import torch + +try: + import psutil + + HAS_PSUTIL = True +except ImportError: + HAS_PSUTIL = False + + +class HardwareMonitor: + """Monitor CPU, GPU, memory usage.""" + + def __init__(self): + self.has_psutil = HAS_PSUTIL + if self.has_psutil: + self.process = psutil.Process() + + self.has_cuda = torch.cuda.is_available() + if self.has_cuda: + try: + import pynvml + + pynvml.nvmlInit() + self.nvml_handle = pynvml.nvmlDeviceGetHandleByIndex(0) + self.has_nvml = True + except Exception: + self.has_nvml = False + else: + self.has_nvml = False + + def get_metrics(self) -> Dict[str, float]: + """Get current hardware metrics.""" + metrics = {} + + # CPU & Memory (requires psutil) + if self.has_psutil: + metrics["cpu_percent"] = self.process.cpu_percent() + metrics["cpu_count"] = psutil.cpu_count() + mem = self.process.memory_info() + metrics["memory_rss_mb"] = mem.rss / 1024 / 1024 + metrics["memory_percent"] = self.process.memory_percent() + + # GPU + if self.has_cuda: + metrics["gpu_memory_allocated_mb"] = torch.cuda.memory_allocated() / 1024 / 1024 + metrics["gpu_memory_reserved_mb"] = torch.cuda.memory_reserved() / 1024 / 1024 + + if self.has_nvml: + import pynvml + + util = pynvml.nvmlDeviceGetUtilizationRates(self.nvml_handle) + metrics["gpu_utilization"] = util.gpu + metrics["gpu_memory_utilization"] = util.memory + + return metrics 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 index 767f02f3a..8dd77421e 100644 --- a/src/unilab/utils/algo_utils.py +++ b/src/unilab/utils/algo_utils.py @@ -1,135 +1,15 @@ -"""Common utilities for RL algorithms.""" - from __future__ import annotations -import importlib -import logging -from typing import Sequence - -logger = logging.getLogger(__name__) +import warnings -# Attribute name for package-level registry bootstrap contracts. -_REGISTRY_MODULES_ATTR = "__unilab_registry_modules__" +from unilab.algos.torch.common.actor_factory import build_actor +from unilab.base.registry import ensure_registries -# Default packages to import for env registration bootstrap. -_DEFAULT_REGISTRY_PACKAGES = ( - "unilab.envs.locomotion", - "unilab.envs.manipulation", - "unilab.envs.motion_tracking", +warnings.warn( + "`unilab.utils.algo_utils` is deprecated and will be removed in 0.2.0; " + "use `unilab.base.registry` and `unilab.algos.torch.common.actor_factory` instead.", + DeprecationWarning, + stacklevel=2, ) - -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}") +__all__ = ["build_actor", "ensure_registries"] 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/device_utils.py b/src/unilab/utils/device_utils.py index cee8d479a..08bca5160 100644 --- a/src/unilab/utils/device_utils.py +++ b/src/unilab/utils/device_utils.py @@ -1,29 +1,16 @@ -import torch +from __future__ import annotations -from unilab.base import registry +import warnings +from unilab.algos.torch.common.device import get_env_dims +from unilab.utils.device import get_default_device -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" +warnings.warn( + "`unilab.utils.device_utils` is deprecated and will be removed in 0.2.0; " + "use `unilab.utils.device` for `get_default_device` " + "and `unilab.algos.torch.common.device` for `get_env_dims` instead.", + DeprecationWarning, + stacklevel=2, +) - -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 - - env = registry.make( - env_name, num_envs=1, sim_backend=sim_backend, env_cfg_override=env_cfg_override - ) - obs_dim, critic_dim = get_obs_dims_from_spec(env.obs_groups_spec) - action_shape = env.action_space.shape - assert action_shape is not None - action_dim = action_shape[0] - env.close() # type: ignore[attr-defined] - return obs_dim, action_dim, critic_dim +__all__ = ["get_default_device", "get_env_dims"] diff --git a/src/unilab/utils/experiment_tracking.py b/src/unilab/utils/experiment_tracking.py index e40216dce..72a18530a 100644 --- a/src/unilab/utils/experiment_tracking.py +++ b/src/unilab/utils/experiment_tracking.py @@ -1,416 +1,28 @@ -"""Shared experiment tracking utilities for local files and W&B.""" - from __future__ import annotations -import dataclasses -import getpass -import importlib -import json -import os -import socket -import subprocess -import time -from datetime import datetime, timezone -from pathlib import Path -from typing import Any - -from omegaconf import OmegaConf - - -def _cfg_get(cfg: Any, key: str, default: Any = None) -> Any: - if cfg is None: - return default - if isinstance(cfg, dict): - return cfg.get(key, default) - return getattr(cfg, key, default) - - -def _plain_dict(value: Any) -> Any: - if OmegaConf.is_config(value): - return OmegaConf.to_container(value, resolve=True) - if dataclasses.is_dataclass(value) and not isinstance(value, type): - return dataclasses.asdict(value) - return value - - -def _load_wandb() -> Any | None: - try: - return importlib.import_module("wandb") - except ImportError: - return None - - -def _json_safe(value: Any) -> Any: - if isinstance(value, dict): - return {str(k): _json_safe(v) for k, v in value.items()} - if isinstance(value, (list, tuple)): - return [_json_safe(v) for v in value] - if isinstance(value, Path): - return str(value) - if isinstance(value, (str, int, float, bool)) or value is None: - return value - try: - json.dumps(value) - return value - except TypeError: - return str(value) - - -def get_device_info_dict() -> dict[str, str]: - try: - module = importlib.import_module("benchmark.core.device_info") - getter = getattr(module, "get_device_info_dict") - return dict(getter()) - except Exception: - return {"platform": os.uname().sysname if hasattr(os, "uname") else "unknown"} - - -def get_git_info(root_dir: str | Path) -> dict[str, Any]: - root = Path(root_dir) - - def _run_git(*args: str) -> str | None: - try: - result = subprocess.run( - ["git", *args], - cwd=root, - check=True, - capture_output=True, - text=True, - ) - except Exception: - return None - return result.stdout.strip() - - commit = _run_git("rev-parse", "HEAD") - branch = _run_git("rev-parse", "--abbrev-ref", "HEAD") - status = _run_git("status", "--short") - - return { - "commit": commit, - "branch": branch, - "dirty": bool(status), - } - - -def build_wandb_run_name(algo_name: str, task_name: str, log_dir: str | Path | None) -> str: - if log_dir is None: - return f"{algo_name}__{task_name}" - run_dir = Path(log_dir) - return f"{algo_name}__{task_name}__{run_dir.name}" - - -def build_wandb_settings( - training_cfg: Any, - *, - algo_name: str, - task_name: str, - sim_backend: str, - log_dir: str | Path | None, -) -> dict[str, Any]: - name = _cfg_get(training_cfg, "wandb_name") - if not name: - name = build_wandb_run_name(algo_name, task_name, log_dir) - - group = _cfg_get(training_cfg, "wandb_group") - if not group: - group = task_name - - job_type = _cfg_get(training_cfg, "wandb_job_type") - if not job_type: - job_type = algo_name - - tags = [str(tag) for tag in (_cfg_get(training_cfg, "wandb_tags", []) or [])] - auto_tags = [algo_name, task_name, sim_backend, f"user-{getpass.getuser()}"] - for tag in auto_tags: - if tag not in tags: - tags.append(tag) - - return { - "project": _cfg_get(training_cfg, "wandb_project", "unilab"), - "entity": _cfg_get(training_cfg, "wandb_entity"), - "name": name, - "group": group, - "job_type": job_type, - "tags": tags, - "notes": _cfg_get(training_cfg, "wandb_notes"), - "mode": _cfg_get(training_cfg, "wandb_mode"), - } - - -class ExperimentTracker: - """Tracks experiment metadata locally and optionally in Weights & Biases.""" - - def __init__( - self, - *, - root_dir: str | Path, - log_dir: str | Path, - algo_name: str, - task_name: str, - sim_backend: str, - training_cfg: Any, - full_cfg: Any, - device: str | None = None, - collector_device: str | None = None, - ): - self.root_dir = Path(root_dir) - self.log_dir = Path(log_dir) - self.algo_name = algo_name - self.task_name = task_name - self.sim_backend = sim_backend - self.training_cfg = training_cfg - self.full_cfg = full_cfg - self.device = device - self.collector_device = collector_device - self.enabled = str(_cfg_get(training_cfg, "logger", "tensorboard")).lower() == "wandb" - - self.log_dir.mkdir(parents=True, exist_ok=True) - - self._wandb = None - self._run = None - self._owns_run = False - self._started = False - self._start_monotonic = 0.0 - self._start_utc = "" - self._summary: dict[str, Any] = {} - - @property - def run(self) -> Any | None: - return self._run - - @property - def run_url(self) -> str | None: - return getattr(self._run, "url", None) if self._run is not None else None - - @property - def wandb_settings(self) -> dict[str, Any]: - return build_wandb_settings( - self.training_cfg, - algo_name=self.algo_name, - task_name=self.task_name, - sim_backend=self.sim_backend, - log_dir=self.log_dir, - ) - - def start(self) -> None: - if self._started: - return - - self._started = True - self._start_monotonic = time.perf_counter() - self._start_utc = datetime.now(timezone.utc).isoformat() - - metadata = { - "algo": self.algo_name, - "task": self.task_name, - "sim_backend": self.sim_backend, - "device": self.device, - "collector_device": self.collector_device, - "log_dir": str(self.log_dir), - "start_time_utc": self._start_utc, - "hostname": socket.gethostname(), - "user": getpass.getuser(), - "git": get_git_info(self.root_dir), - "hardware": get_device_info_dict(), - "wandb": self.wandb_settings, - } - - payload = { - "run": _json_safe(metadata), - "config": _json_safe(_plain_dict(self.full_cfg)), - } - self._write_json(self.log_dir / "run_config.json", payload) - - if not self.enabled: - return - - self._wandb = _load_wandb() - if self._wandb is None: - print("[experiment_tracking] wandb not installed, skipping W&B experiment tracking.") - return - - self._run = self._wandb.run - if self._run is None: - kwargs = { - "project": self.wandb_settings["project"], - "name": self.wandb_settings["name"], - "config": payload, - "dir": str(self.log_dir), - "reinit": True, - } - for key in ("entity", "group", "job_type", "tags", "notes", "mode"): - value = self.wandb_settings.get(key) - if value not in (None, "", []): - kwargs[key] = value - self._run = self._wandb.init(**kwargs) - self._owns_run = True - else: - self._run.config.update(payload, allow_val_change=True) - - if self._run is not None: - self._run.summary["algo"] = self.algo_name - self._run.summary["task"] = self.task_name - self._run.summary["sim_backend"] = self.sim_backend - if self.device: - self._run.summary["device"] = self.device - if self.collector_device: - self._run.summary["collector_device"] = self.collector_device - self._run.summary["log_dir"] = str(self.log_dir) - - def update_summary(self, summary: dict[str, Any] | None = None) -> None: - if summary: - self._summary.update(summary) - - if not self._started: - return - - wall_time_sec = time.perf_counter() - self._start_monotonic - payload = { - **self._summary, - "algo": self.algo_name, - "task": self.task_name, - "sim_backend": self.sim_backend, - "log_dir": str(self.log_dir), - "start_time_utc": self._start_utc, - "end_time_utc": datetime.now(timezone.utc).isoformat(), - "wall_time_sec": wall_time_sec, - "wandb_run_url": self.run_url, - } - self._write_json(self.log_dir / "run_summary.json", _json_safe(payload)) - - if self._run is not None: - for key, value in payload.items(): - self._run.summary[key] = _json_safe(value) - - def log_video(self, video_path: str | Path | None, key: str = "media/play_video") -> None: - if video_path is None: - return - - video = Path(video_path) - if not video.exists(): - return - - self._summary["play_video_path"] = str(video) - if self._run is not None and self._wandb is not None: - self._wandb.log({key: self._wandb.Video(str(video), format="mp4")}) - - def finish(self) -> None: - if not self._started: - return - - self.update_summary() - if self._run is not None and self._wandb is not None and self._owns_run: - self._wandb.finish() - self._run = None - self._wandb = None - - @staticmethod - def _write_json(path: Path, payload: Any) -> None: - path.parent.mkdir(parents=True, exist_ok=True) - path.write_text(json.dumps(payload, indent=2), encoding="utf-8") - - -def patch_rsl_rl_wandb_writer() -> None: - """Patch rsl-rl W&B writer so it can reuse an already-open run.""" - try: - import rsl_rl.utils.wandb_utils as wandb_utils - except Exception: - return - - if getattr(wandb_utils, "_UNILAB_PATCHED", False): - return - - wandb = _load_wandb() - if wandb is None: - return - wandb_mod = wandb - - from torch.utils.tensorboard import SummaryWriter as TensorboardSummaryWriter - - class PatchedWandbSummaryWriter(TensorboardSummaryWriter): - def __init__(self, log_dir: str, flush_secs: int, cfg: dict) -> None: - super().__init__(log_dir, flush_secs=flush_secs) - - run_name = os.path.split(log_dir)[-1] - project = cfg.get("wandb_project", "unilab") - entity = cfg.get("wandb_entity") or os.environ.get("WANDB_USERNAME") - group = cfg.get("wandb_group") - job_type = cfg.get("wandb_job_type") - tags = cfg.get("wandb_tags") - notes = cfg.get("wandb_notes") - mode = cfg.get("wandb_mode") - - self.logged_videos: set[str] = set() - self._owns_run = wandb_mod.run is None - if self._owns_run: - kwargs = { - "project": project, - "name": run_name, - "config": {"log_dir": log_dir}, - "settings": wandb_mod.Settings(start_method="thread"), - } - if entity: - kwargs["entity"] = entity - if group: - kwargs["group"] = group - if job_type: - kwargs["job_type"] = job_type - if tags: - kwargs["tags"] = tags - if notes: - kwargs["notes"] = notes - if mode: - kwargs["mode"] = mode - wandb_mod.init(**kwargs) - else: - wandb_mod.config.update({"log_dir": log_dir}, allow_val_change=True) - - def store_config(self, env_cfg: dict | object, train_cfg: dict) -> None: - wandb_mod.config.update({"train_cfg": train_cfg}, allow_val_change=True) - env_payload: Any - if isinstance(env_cfg, dict): - env_payload = env_cfg - elif dataclasses.is_dataclass(env_cfg) and not isinstance(env_cfg, type): - env_payload = dataclasses.asdict(env_cfg) - elif hasattr(env_cfg, "to_dict"): - env_payload = env_cfg.to_dict() # type: ignore[union-attr] - else: - env_payload = str(env_cfg) - wandb_mod.config.update({"env_cfg": env_payload}, allow_val_change=True) - - def add_scalar( - self, - tag: Any, - scalar_value: Any, - global_step: Any = None, - walltime: Any = None, - new_style: Any = False, - double_precision: Any = False, - ) -> None: - super().add_scalar( - tag, - scalar_value, - global_step=global_step, - walltime=walltime, - new_style=new_style, - double_precision=double_precision, - ) - wandb_mod.log({tag: scalar_value}, step=global_step) - - def stop(self) -> None: - if self._owns_run: - wandb_mod.finish() - - def save_model(self, model_path: str, it: int) -> None: - wandb_mod.save(model_path, base_path=os.path.dirname(model_path)) - - def save_file(self, path: str) -> None: - wandb_mod.save(path, base_path=os.path.dirname(path)) - - def save_video(self, video: Path, it: int) -> None: - if video.name not in self.logged_videos: - wandb_mod.log({"video": wandb_mod.Video(str(video), format="mp4")}, step=it) - self.logged_videos.add(video.name) - - wandb_utils.WandbSummaryWriter = PatchedWandbSummaryWriter - setattr(wandb_utils, "_UNILAB_PATCHED", True) +import warnings + +from unilab.training.logging.experiment import ( + ExperimentTracker, + build_wandb_run_name, + build_wandb_settings, + get_device_info_dict, + get_git_info, + patch_rsl_rl_wandb_writer, +) + +warnings.warn( + "`unilab.utils.experiment_tracking` is deprecated and will be removed in 0.2.0; " + "use `unilab.training.logging.experiment` instead.", + DeprecationWarning, + stacklevel=2, +) + +__all__ = [ + "ExperimentTracker", + "build_wandb_run_name", + "build_wandb_settings", + "get_device_info_dict", + "get_git_info", + "patch_rsl_rl_wandb_writer", +] diff --git a/src/unilab/utils/final_observation.py b/src/unilab/utils/final_observation.py index f18fc1ba6..e6be7472f 100644 --- a/src/unilab/utils/final_observation.py +++ b/src/unilab/utils/final_observation.py @@ -1,160 +1,26 @@ from __future__ import annotations -from dataclasses import dataclass -from typing import Any - -import numpy as np - - -@dataclass(frozen=True) -class TransitionBootstrapContract: - actor_next_obs: np.ndarray - transition_next_obs: np.ndarray - terminal_mask: np.ndarray - timeout_terminal_mask: np.ndarray - actor_next_critic: np.ndarray | None = None - transition_next_critic: np.ndarray | None = None - - -@dataclass(frozen=True) -class TerminalObservationContract: - terminal_obs: np.ndarray | None - terminal_mask: np.ndarray - timeout_terminal_mask: np.ndarray - terminal_critic: np.ndarray | None = None - - -def patch_transition_next_obs( - next_obs: np.ndarray, - final_observation: dict[str, Any] | None = None, - done: np.ndarray | None = None, - info: dict[str, Any] | None = None, - next_critic: np.ndarray | None = None, -) -> tuple[np.ndarray, np.ndarray | None, np.ndarray]: - """Patch transition next obs with final_observation without mutating actor inputs.""" - terminal_contract = resolve_terminal_observation_contract( - next_obs_batch_size=next_obs.shape[0], - final_observation=final_observation, - done=done, - info=info, - ) - if not np.any(terminal_contract.terminal_mask) or terminal_contract.terminal_obs is None: - return ( - next_obs, - next_critic, - np.zeros((next_obs.shape[0],), dtype=bool), - ) - - transition_next_obs = next_obs.copy() - transition_next_obs[terminal_contract.terminal_mask] = np.asarray( - terminal_contract.terminal_obs, dtype=next_obs.dtype - )[terminal_contract.terminal_mask] - - transition_next_critic = next_critic - if next_critic is not None and terminal_contract.terminal_critic is not None: - transition_next_critic = next_critic.copy() - transition_next_critic[terminal_contract.terminal_mask] = np.asarray( - terminal_contract.terminal_critic, dtype=next_critic.dtype - )[terminal_contract.terminal_mask] - - return ( - transition_next_obs, - transition_next_critic, - terminal_contract.terminal_mask, - ) - - -def resolve_transition_bootstrap_contract( - next_obs: np.ndarray, - info: dict[str, Any] | None = None, - final_observation: dict[str, Any] | None = None, - done: np.ndarray | None = None, - truncated: np.ndarray | None = None, - next_critic: np.ndarray | None = None, -) -> TransitionBootstrapContract: - """Resolve actor/storage observations and timeout bootstrap masks for a step.""" - ( - transition_next_obs, - transition_next_critic, - terminal_mask, - ) = patch_transition_next_obs( - next_obs, - final_observation=final_observation, - done=done, - info=info, - next_critic=next_critic, - ) - timeout_terminal_mask = terminal_mask - if truncated is not None: - timeout_terminal_mask = np.logical_and( - terminal_mask, np.asarray(truncated, dtype=bool).ravel() - ) - return TransitionBootstrapContract( - actor_next_obs=next_obs, - transition_next_obs=transition_next_obs, - terminal_mask=terminal_mask, - timeout_terminal_mask=timeout_terminal_mask, - actor_next_critic=next_critic, - transition_next_critic=transition_next_critic, - ) - - -def resolve_terminal_observation_contract( - next_obs_batch_size: int, - final_observation: dict[str, Any] | None = None, - done: np.ndarray | None = None, - info: dict[str, Any] | None = None, - truncated: np.ndarray | None = None, -) -> TerminalObservationContract: - """Resolve terminal observation facts without constructing patched next obs.""" - terminal_mask = _resolve_terminal_mask(next_obs_batch_size, done, info) - resolved_final_observation = _resolve_final_observation(final_observation, info) - - terminal_obs: np.ndarray | None = None - terminal_critic: np.ndarray | None = None - if np.any(terminal_mask) and isinstance(resolved_final_observation, dict): - terminal_obs = resolved_final_observation.get("obs") - terminal_critic = resolved_final_observation.get("critic") - - timeout_terminal_mask = terminal_mask - if truncated is not None: - timeout_terminal_mask = np.logical_and( - terminal_mask, np.asarray(truncated, dtype=bool).ravel() - ) - - return TerminalObservationContract( - terminal_obs=terminal_obs, - terminal_mask=terminal_mask, - timeout_terminal_mask=timeout_terminal_mask, - terminal_critic=terminal_critic, - ) - - -def _resolve_final_observation( - final_observation: dict[str, Any] | None, - info: dict[str, Any] | None, -) -> dict[str, Any] | None: - if isinstance(final_observation, dict): - return final_observation - if isinstance(info, dict): - final_obs = info.get("final_observation") - if isinstance(final_obs, dict): - return final_obs - return None - - -def _resolve_terminal_mask( - next_obs_batch_size: int, - done: np.ndarray | None, - info: dict[str, Any] | None, -) -> np.ndarray: - if done is not None: - done_mask = np.asarray(done, dtype=bool).ravel() - if done_mask.shape == (next_obs_batch_size,): - return done_mask - return np.zeros((next_obs_batch_size,), dtype=bool) - if isinstance(info, dict): - terminal_mask = np.asarray(info.get("_final_observation"), dtype=bool) - if terminal_mask.shape == (next_obs_batch_size,): - return terminal_mask - return np.zeros((next_obs_batch_size,), dtype=bool) +import warnings + +from unilab.base.final_observation import ( + TerminalObservationContract, + TransitionBootstrapContract, + patch_transition_next_obs, + resolve_terminal_observation_contract, + resolve_transition_bootstrap_contract, +) + +warnings.warn( + "`unilab.utils.final_observation` is deprecated and will be removed in 0.2.0; " + "use `unilab.base.final_observation` instead.", + DeprecationWarning, + stacklevel=2, +) + +__all__ = [ + "TerminalObservationContract", + "TransitionBootstrapContract", + "patch_transition_next_obs", + "resolve_terminal_observation_contract", + "resolve_transition_bootstrap_contract", +] diff --git a/src/unilab/utils/hardware_monitor.py b/src/unilab/utils/hardware_monitor.py index cafcd38a0..40b53fe89 100644 --- a/src/unilab/utils/hardware_monitor.py +++ b/src/unilab/utils/hardware_monitor.py @@ -1,60 +1,14 @@ -"""Hardware monitoring utilities for performance profiling.""" +from __future__ import annotations -from typing import Dict +import warnings -import torch +from unilab.training.monitoring import HardwareMonitor -try: - import psutil +warnings.warn( + "`unilab.utils.hardware_monitor` is deprecated and will be removed in 0.2.0; " + "use `unilab.training.monitoring` instead.", + DeprecationWarning, + stacklevel=2, +) - HAS_PSUTIL = True -except ImportError: - HAS_PSUTIL = False - - -class HardwareMonitor: - """Monitor CPU, GPU, memory usage.""" - - def __init__(self): - self.has_psutil = HAS_PSUTIL - if self.has_psutil: - self.process = psutil.Process() - - self.has_cuda = torch.cuda.is_available() - if self.has_cuda: - try: - import pynvml - - pynvml.nvmlInit() - self.nvml_handle = pynvml.nvmlDeviceGetHandleByIndex(0) - self.has_nvml = True - except Exception: - self.has_nvml = False - else: - self.has_nvml = False - - def get_metrics(self) -> Dict[str, float]: - """Get current hardware metrics.""" - metrics = {} - - # CPU & Memory (requires psutil) - if self.has_psutil: - metrics["cpu_percent"] = self.process.cpu_percent() - metrics["cpu_count"] = psutil.cpu_count() - mem = self.process.memory_info() - metrics["memory_rss_mb"] = mem.rss / 1024 / 1024 - metrics["memory_percent"] = self.process.memory_percent() - - # GPU - if self.has_cuda: - metrics["gpu_memory_allocated_mb"] = torch.cuda.memory_allocated() / 1024 / 1024 - metrics["gpu_memory_reserved_mb"] = torch.cuda.memory_reserved() / 1024 / 1024 - - if self.has_nvml: - import pynvml - - util = pynvml.nvmlDeviceGetUtilizationRates(self.nvml_handle) - metrics["gpu_utilization"] = util.gpu - metrics["gpu_memory_utilization"] = util.memory - - return metrics +__all__ = ["HardwareMonitor"] diff --git a/src/unilab/utils/logging_common.py b/src/unilab/utils/logging_common.py index be017b7a5..288594d24 100644 --- a/src/unilab/utils/logging_common.py +++ b/src/unilab/utils/logging_common.py @@ -1,316 +1,14 @@ from __future__ import annotations -import importlib -import os -import time -from collections import deque -from typing import Any +import warnings -from rich import box -from rich.console import Console -from rich.live import Live -from rich.panel import Panel -from rich.table import Table -from rich.text import Text +from unilab.training.logging.common import BaseTrainingLogger, _fmt_number, _fmt_time, _load_wandb +warnings.warn( + "`unilab.utils.logging_common` is deprecated and will be removed in 0.2.0; " + "use `unilab.training.logging.common` instead.", + DeprecationWarning, + stacklevel=2, +) -def _fmt_time(seconds: float) -> str: - if seconds < 60: - return f"{seconds:.0f}s" - m, s = divmod(int(seconds), 60) - if m < 60: - return f"{m}m{s:02d}s" - h, m = divmod(m, 60) - return f"{h}h{m:02d}m{s:02d}s" - - -def _fmt_number(v: float) -> str: - if abs(v) == 0: - return "0" - if abs(v) >= 1e6: - return f"{v:.2e}" - if abs(v) >= 100: - return f"{v:.1f}" - if abs(v) >= 1: - return f"{v:.3f}" - if abs(v) >= 0.001: - return f"{v:.4f}" - return f"{v:.2e}" - - -def _load_wandb() -> Any | None: - """Load wandb lazily so it remains an optional dependency.""" - try: - return importlib.import_module("wandb") - except ImportError: - return None - - -class BaseTrainingLogger: - """Shared lifecycle and backend logging setup for rich training loggers.""" - - def __init__( - self, - *, - algo_name: str, - max_iterations: int, - num_envs: int, - env_name: str, - log_dir: str, - log_backend: str, - wandb_project: str, - wandb_entity: str | None, - wandb_name: str, - wandb_group: list[str] | None | str, - wandb_job_type: str | None, - wandb_tags: list[str] | None, - wandb_notes: str | None, - refresh_per_second: int = 4, - tensorboard_subdir: str | None = "tb", - wandb_config: dict[str, Any] | None = None, - ): - self.algo_name = algo_name - self.max_iterations = max_iterations - self.num_envs = num_envs - self.env_name = env_name - - self._no_print = log_backend.lower() == "no_print" - self._log_backend = "none" if self._no_print else log_backend.lower() - - self._console = Console() - self._live: Live | None = None - self._refresh_rate = refresh_per_second - - self._start_time: float = 0.0 - self._iteration: int = 0 - self._reward_history: deque[float] = deque(maxlen=200) - self._latest_metrics: dict[str, float] = {} - self._latest_reward_components: dict[str, float] = {} - self._collect_time: float = 0.0 - self._train_time: float = 0.0 - self._mean_ep_length: float = 0.0 - self._last_save: str = "" - self._status: str = "" - - self._log_dir = log_dir - self._tb_writer: Any | None = None - self._wandb_run = None - self._owns_wandb_run = False - - if self._log_backend == "tensorboard" and log_dir: - self._init_tensorboard(log_dir, tensorboard_subdir) - elif self._log_backend == "wandb": - self._init_wandb( - project=wandb_project, - entity=wandb_entity, - name=wandb_name or f"{algo_name}_{env_name}", - log_dir=log_dir, - group=wandb_group, - job_type=wandb_job_type, - tags=wandb_tags, - notes=wandb_notes, - extra_config=wandb_config, - ) - - def _format_tensorboard_message(self, tb_dir: str) -> str: - return f"[dim]TensorBoard: {tb_dir}[/]" - - def _format_wandb_message(self, project: str, name: str) -> str: - return f"[dim]W&B: {project}/{name}[/]" - - def _init_tensorboard(self, log_dir: str, subdir: str | None): - try: - from torch.utils.tensorboard import SummaryWriter - - tb_dir = log_dir if subdir is None else os.path.join(log_dir, subdir) - os.makedirs(tb_dir, exist_ok=True) - self._tb_writer = SummaryWriter(log_dir=tb_dir) - if not self._no_print: - self._console.print(self._format_tensorboard_message(tb_dir)) - except ImportError: - if not self._no_print: - self._console.print("[yellow]tensorboard not installed[/]") - - def _init_wandb( - self, - *, - project: str, - entity: str | None, - name: str, - log_dir: str, - group: str | None | list[str], - job_type: str | None, - tags: list[str] | None, - notes: str | None, - extra_config: dict[str, Any] | None = None, - ): - wandb = _load_wandb() - if wandb is None: - if not self._no_print: - self._console.print("[yellow]wandb not installed[/]") - return - - self._wandb_run = wandb.run - if self._wandb_run is None: - config: dict[str, Any] = { - "algo": self.algo_name, - "env": self.env_name, - "num_envs": self.num_envs, - } - if extra_config: - config.update(extra_config) - - kwargs: dict[str, Any] = { - "project": project, - "name": name, - "config": config, - "dir": log_dir or None, - "reinit": True, - } - if entity: - kwargs["entity"] = entity - if group: - kwargs["group"] = group - if job_type: - kwargs["job_type"] = job_type - if tags: - kwargs["tags"] = tags - if notes: - kwargs["notes"] = notes - - self._wandb_run = wandb.init(**kwargs) - self._owns_wandb_run = True - - if not self._no_print: - self._console.print(self._format_wandb_message(project, name)) - - def start(self, *, status: str = ""): - self._start_time = time.time() - self._status = status - if not self._no_print: - self._live = Live( - self._build_display(), - console=self._console, - refresh_per_second=self._refresh_rate, - transient=False, - ) - self._live.start() - - def finish(self, *, title: str = "Training Summary", extra_summary: str = ""): - if self._live is not None: - self._live.update(self._build_display()) - self._live.stop() - self._live = None - - elapsed = time.time() - self._start_time - if not self._no_print: - summary = ( - f"[bold green]Training complete[/]\n" - f" Algo: [cyan]{self.algo_name}[/] | Env: [cyan]{self.env_name}[/]\n" - f" Iterations: [yellow]{self._iteration}[/]/{self.max_iterations}\n" - f" Total time: [yellow]{_fmt_time(elapsed)}[/]\n" - ) - if extra_summary: - summary += extra_summary - if self._last_save: - summary += f" Last checkpoint: [dim]{self._last_save}[/]" - - self._console.print() - self._console.print(Panel(summary, title=f"[bold]{title}[/]", border_style="green")) - - if self._tb_writer: - self._tb_writer.close() - if self._wandb_run and self._owns_wandb_run: - wandb = _load_wandb() - if wandb is not None: - wandb.finish() - - def update_ep_length(self, length: float): - self._mean_ep_length = length - - def log_save(self, path: str): - self._last_save = path - self._refresh() - - def _refresh(self): - if self._live is not None: - self._live.update(self._build_display()) - - def _estimate_eta(self) -> str: - if self._iteration <= 0: - return "" - elapsed = time.time() - self._start_time - remaining = self.max_iterations - self._iteration - avg_iter = elapsed / self._iteration - eta_s = remaining * avg_iter - return _fmt_time(eta_s) - - def _build_header(self, *, include_status: bool) -> Panel: - elapsed = time.time() - self._start_time if self._start_time else 0 - eta = self._estimate_eta() - - header_text = Text() - header_text.append(f" {self.algo_name}", style="bold cyan") - header_text.append(" │ ", style="dim") - header_text.append(f"{self.env_name}", style="bold white") - header_text.append(" │ ", style="dim") - header_text.append(f"iter {self._iteration}/{self.max_iterations}", style="yellow") - header_text.append(" │ ", style="dim") - header_text.append(f"⏱ {_fmt_time(elapsed)}", style="green") - if eta: - header_text.append(" │ ETA ", style="dim") - header_text.append(eta, style="bold magenta") - if include_status and self._status: - header_text.append(" │ ", style="dim") - header_text.append(self._status, style="dim italic") - - return Panel(header_text, style="dim", box=box.SIMPLE) - - def _build_reward_table_common(self, *, wait_message: str) -> Table: - table = Table( - title="[bold]Rewards[/]", - box=box.SIMPLE_HEAVY, - show_header=True, - header_style="bold green", - expand=True, - pad_edge=False, - ) - table.add_column("Component", style="white", ratio=2) - table.add_column("Value", justify="right", ratio=1) - - if self._reward_history: - recent = list(self._reward_history) - mean_rew = sum(recent[-50:]) / max(len(recent[-50:]), 1) - peak_rew = max(recent) if recent else 0 - - if len(recent) >= 10: - old = sum(recent[-20:-10]) / 10 - new = sum(recent[-10:]) / 10 - trend = ( - "[green]▲[/]" - if new > old * 1.05 - else "[red]▼[/]" - if new < old * 0.95 - else "[yellow]━[/]" - ) - else: - trend = "" - - table.add_row(f"[bold]Mean Reward[/] {trend}", f"[bold green]{mean_rew:.3f}[/]") - table.add_row(" Peak", f"[dim]{peak_rew:.3f}[/]") - if self._mean_ep_length > 0: - table.add_row(" Ep Len", f"[dim]{self._mean_ep_length:.1f}[/]") - table.add_row("", "") - else: - table.add_row(wait_message, "") - - if self._latest_reward_components: - for name, val in sorted(self._latest_reward_components.items()): - display = name.replace("reward/", "").replace("_", " ") - color = "green" if val > 0 else "red" if val < 0 else "dim" - table.add_row(f" {display}", f"[{color}]{val:+.4f}[/]") - - return table - - def _build_display(self) -> Panel: - raise NotImplementedError +__all__ = ["BaseTrainingLogger", "_fmt_number", "_fmt_time", "_load_wandb"] diff --git a/src/unilab/utils/math_utils.py b/src/unilab/utils/math_utils.py index 429df0927..a3d40cd73 100644 --- a/src/unilab/utils/math_utils.py +++ b/src/unilab/utils/math_utils.py @@ -1,346 +1,52 @@ 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 - q2_was_1d = q2.ndim == 1 - - if q1_was_1d: - q1 = q1[None, :] - if q2_was_1d: - q2 = q2[None, :] - - if q1.shape[0] == 1 and q2.shape[0] > 1: - q1 = np.broadcast_to(q1, q2.shape) - elif q2.shape[0] == 1 and q1.shape[0] > 1: - q2 = np.broadcast_to(q2, q1.shape) - - w1, x1, y1, z1 = q1[:, 0], q1[:, 1], q1[:, 2], q1[:, 3] - w2, x2, y2, z2 = q2[:, 0], q2[:, 1], q2[:, 2], q2[:, 3] - result = np.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, - ) - return result[0] if q1_was_1d and q2_was_1d else result - - -def np_quat_conjugate(q: np.ndarray) -> np.ndarray: - """Conjugate of unit quaternions (N, 4) or (4,), w-first.""" - if q.ndim == 1: - return np.array([q[0], -q[1], -q[2], -q[3]]) - conj = q.copy() - conj[:, 1:] *= -1 - return conj # type: ignore[no-any-return] - - -def np_quat_canonicalize(q: np.ndarray) -> np.ndarray: - """Flip quaternion signs so the real part is non-negative.""" - q_was_1d = q.ndim == 1 - if q_was_1d: - q = q[None, :] - - sign = np.where(q[:, 0:1] < 0.0, -1.0, 1.0) - result = q * sign - canonical: np.ndarray = result[0] if q_was_1d else result - return canonical - - -def np_quat_ensure_continuity(q: np.ndarray) -> np.ndarray: - """Flip quaternion signs in a time sequence to keep adjacent dots non-negative.""" - if q.ndim != 2 or q.shape[1] != 4: - raise ValueError(f"Expected quaternion sequence with shape (T, 4), got {q.shape}") - - result = np.array(q, copy=True) - for i in range(1, result.shape[0]): - if float(np.dot(result[i - 1], result[i])) < 0.0: - result[i] *= -1.0 - return result - - -def np_quat_to_axis_angle(q: np.ndarray) -> np.ndarray: - """Convert unit quaternion batch (N, 4), w-first, to axis-angle vectors (N, 3). - - Adapted from PyTorch3D. Uses atan2 + Taylor expansion for numerical - stability near zero rotation. - """ - q = np_quat_canonicalize(q) - xyz = q[:, 1:] # (N, 3) imaginary part - w = q[:, 0:1] # (N, 1) real part - norms = np.linalg.norm(xyz, axis=-1, keepdims=True) # (N, 1) - half_angle = np.arctan2(norms, w) # (N, 1) - angle = 2.0 * half_angle # (N, 1) - small = np.abs(angle) < 1e-6 # (N, 1) - safe_angle = np.where(small, 1.0, angle) - sin_half_over_angle = np.where( - small, - 0.5 - angle**2 / 48.0, - np.sin(half_angle) / safe_angle, - ) - axis_angle: np.ndarray = xyz / sin_half_over_angle - return axis_angle - - -def np_quat_angular_velocity(q: np.ndarray, dt: float) -> np.ndarray: - """Estimate angular velocity from a quaternion time sequence using shortest-arc diffs.""" - if q.ndim != 2 or q.shape[1] != 4: - raise ValueError(f"Expected quaternion sequence with shape (T, 4), got {q.shape}") - if dt <= 0.0: - raise ValueError(f"dt must be positive, got {dt}") - - rotations = np_quat_ensure_continuity(q) - num_frames = rotations.shape[0] - omega = np.zeros((num_frames, 3), dtype=rotations.dtype) - if num_frames <= 1: - return omega - - if num_frames == 2: - q_rel = np_quat_mul(rotations[1], np_quat_conjugate(rotations[0])) - q_rel = np_quat_canonicalize(q_rel) - angvel = np_quat_to_axis_angle(q_rel[None, :])[0] / dt - omega[:] = angvel - return omega - - q_prev = rotations[:-2] - q_next = rotations[2:] - q_rel = np_quat_mul(q_next, np_quat_conjugate(q_prev)) - q_rel = np_quat_canonicalize(q_rel) - omega[1:-1] = np_quat_to_axis_angle(q_rel) / (2.0 * dt) - omega[0] = omega[1] - omega[-1] = omega[-2] - return omega - - -def np_yaw_to_quat(yaw: np.ndarray) -> np.ndarray: - """Convert yaw batch (N,) to quaternion batch (N, 4) in NumPy.""" - half = 0.5 * yaw - return np.stack( - [ - np.cos(half), - np.zeros_like(half), - np.zeros_like(half), - np.sin(half), - ], - axis=1, - ) - - -def np_quat_inv(q: np.ndarray) -> np.ndarray: - """Inverse of unit quaternions (N, 4) or (4,), w-first.""" - return np_quat_conjugate(q) - - -def np_quat_apply(q: np.ndarray, v: np.ndarray) -> np.ndarray: - """Rotate vector(s) by quaternion(s), supports batched/scalar inputs.""" - q_was_1d = q.ndim == 1 - v_was_1d = v.ndim == 1 - - if q_was_1d: - q = q[None, :] - if v_was_1d: - v = v[None, :] - - if q.shape[0] == 1 and v.shape[0] > 1: - q = np.broadcast_to(q, (v.shape[0], 4)) - elif v.shape[0] == 1 and q.shape[0] > 1: - v = np.broadcast_to(v, (q.shape[0], 3)) - - w, x, y, z = q[:, 0], q[:, 1], q[:, 2], q[:, 3] - vx, vy, vz = v[:, 0], v[:, 1], v[:, 2] - - t = 2 * np.stack( - [ - y * vz - z * vy, - z * vx - x * vz, - x * vy - y * vx, - ], - axis=1, - ) - t += 2 * w[:, None] * v - - result = v + np.stack( - [ - y * t[:, 2] - z * t[:, 1], - z * t[:, 0] - x * t[:, 2], - x * t[:, 1] - y * t[:, 0], - ], - axis=1, - ) - - rotated: np.ndarray = result[0] if q_was_1d and v_was_1d else result - return rotated - - -def np_quat_apply_inverse(q: np.ndarray, v: np.ndarray) -> np.ndarray: - """Rotate vector(s) by inverse quaternion(s).""" - return np_quat_apply(np_quat_inv(q), v) - - -def np_quat_error_magnitude(q1: np.ndarray, q2: np.ndarray) -> np.ndarray: - """Angular error magnitude between quaternions (N,) or scalar.""" - q1_was_1d = q1.ndim == 1 - q2_was_1d = q2.ndim == 1 - - if q1_was_1d: - q1 = q1[None, :] - if q2_was_1d: - q2 = q2[None, :] - - if q1.shape[0] == 1 and q2.shape[0] > 1: - q1 = np.broadcast_to(q1, q2.shape) - elif q2.shape[0] == 1 and q1.shape[0] > 1: - q2 = np.broadcast_to(q2, q1.shape) - - # Relative rotation from q1 to q2. - q_rel = np_quat_mul(q2, np_quat_inv(q1)) - q_rel = np_quat_canonicalize(q_rel) - - # Use atan2-based angle extraction for better numerical behavior. - xyz_norm = np.linalg.norm(q_rel[:, 1:], axis=1) - w = np.clip(q_rel[:, 0], -1.0, 1.0) - error = 2.0 * np.arctan2(xyz_norm, w) - magnitude: np.ndarray = error[0] if q1_was_1d and q2_was_1d else error - return magnitude - - -def np_quat_from_euler_xyz(roll: np.ndarray, pitch: np.ndarray, yaw: np.ndarray) -> np.ndarray: - """Convert Euler angles (XYZ) to quaternions (N, 4) or (4,), w-first.""" - roll = np.atleast_1d(roll) - pitch = np.atleast_1d(pitch) - yaw = np.atleast_1d(yaw) - squeeze = roll.shape[0] == 1 - - cr = np.cos(roll * 0.5) - sr = np.sin(roll * 0.5) - cp = np.cos(pitch * 0.5) - sp = np.sin(pitch * 0.5) - cy = np.cos(yaw * 0.5) - sy = np.sin(yaw * 0.5) - - w = cr * cp * cy + sr * sp * sy - x = sr * cp * cy - cr * sp * sy - y = cr * sp * cy + sr * cp * sy - z = cr * cp * sy - sr * sp * cy - - result = np.stack([w, x, y, z], axis=1) - return result[0] if squeeze else result - - -def np_yaw_quat(q: np.ndarray) -> np.ndarray: - """Extract yaw-only quaternion from full quaternion(s), w-first.""" - q_was_1d = q.ndim == 1 - if q_was_1d: - q = q[None, :] - - w, x, y, z = q[:, 0], q[:, 1], q[:, 2], q[:, 3] - yaw = np.arctan2(2 * (w * z + x * y), 1 - 2 * (y * y + z * z)) - - half_yaw = yaw * 0.5 - result = np.stack( - [ - np.cos(half_yaw), - np.zeros_like(half_yaw), - np.zeros_like(half_yaw), - np.sin(half_yaw), - ], - axis=1, - ) - - return result[0] if q_was_1d else result - - -def np_matrix_from_quat(q: np.ndarray) -> np.ndarray: - """Convert quaternion(s) to rotation matrix (N, 3, 3) or (3, 3), w-first.""" - q_was_1d = q.ndim == 1 - if q_was_1d: - q = q[None, :] - - w, x, y, z = q[:, 0], q[:, 1], q[:, 2], q[:, 3] - - xx = x * x - yy = y * y - zz = z * z - xy = x * y - xz = x * z - yz = y * z - wx = w * x - wy = w * y - wz = w * z - - result = np.stack( - [ - np.stack([1 - 2 * (yy + zz), 2 * (xy - wz), 2 * (xz + wy)], axis=1), - np.stack([2 * (xy + wz), 1 - 2 * (xx + zz), 2 * (yz - wx)], axis=1), - np.stack([2 * (xz - wy), 2 * (yz + wx), 1 - 2 * (xx + yy)], axis=1), - ], - axis=1, - ) - - return result[0] if q_was_1d else result - - -def np_subtract_frame_transforms( - pos1: np.ndarray, quat1: np.ndarray, pos2: np.ndarray, quat2: np.ndarray -) -> tuple[np.ndarray, np.ndarray]: - """Compute relative transform from frame 1 to frame 2 in frame-1 coordinates.""" - 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) +import warnings + +from unilab.algos.mlx.common.rotation import axis_angle_to_quat, quat_mul +from unilab.envs.common.math import np_sample_uniform +from unilab.envs.common.rotation import ( + np_matrix_from_quat, + np_quat_angular_velocity, + np_quat_apply, + np_quat_apply_inverse, + np_quat_canonicalize, + np_quat_conjugate, + np_quat_ensure_continuity, + np_quat_error_magnitude, + np_quat_from_euler_xyz, + np_quat_inv, + np_quat_mul, + np_quat_to_axis_angle, + np_subtract_frame_transforms, + np_yaw_quat, + np_yaw_to_quat, +) + +warnings.warn( + "`unilab.utils.math_utils` is deprecated and will be removed in 0.2.0; " + "use `unilab.envs.common.rotation`, `unilab.envs.common.math`, or " + "`unilab.algos.mlx.common.rotation` instead.", + DeprecationWarning, + stacklevel=2, +) + +__all__ = [ + "axis_angle_to_quat", + "np_matrix_from_quat", + "np_quat_angular_velocity", + "np_quat_apply", + "np_quat_apply_inverse", + "np_quat_canonicalize", + "np_quat_conjugate", + "np_quat_ensure_continuity", + "np_quat_error_magnitude", + "np_quat_from_euler_xyz", + "np_quat_inv", + "np_quat_mul", + "np_quat_to_axis_angle", + "np_sample_uniform", + "np_subtract_frame_transforms", + "np_yaw_quat", + "np_yaw_to_quat", + "quat_mul", +] diff --git a/src/unilab/utils/obs_utils.py b/src/unilab/utils/obs_utils.py index 3b17f9cfa..3f1a2c384 100644 --- a/src/unilab/utils/obs_utils.py +++ b/src/unilab/utils/obs_utils.py @@ -1,37 +1,26 @@ from __future__ import annotations -import numpy as np - - -def flatten_obs_dict(obs: dict[str, np.ndarray]) -> np.ndarray: - """Concatenate obs groups in insertion order -> flat (N, total_dim) array.""" - return np.concatenate(list(obs.values()), axis=1) - - -def flatten_policy_obs_dict(obs: dict[str, np.ndarray]) -> np.ndarray: - """Build actor-policy inputs from the single actor observation group.""" - return obs["obs"] - - -def split_obs_dict(obs: dict[str, np.ndarray]) -> tuple[np.ndarray, np.ndarray]: - """Split observation dict into (actor_obs, critic_obs). - - When no separate critic group exists, critic_obs == actor_obs. - """ - actor = obs["obs"] - return actor, obs.get("critic", actor) - - -def get_obs_dims(obs_groups_spec: dict[str, int]) -> tuple[int, int]: - """Extract (actor_obs_dim, critic_obs_dim) from obs_groups_spec. - - When no separate critic group exists, critic_obs_dim == actor_obs_dim. - """ - obs_dim = obs_groups_spec.get("obs", 0) - return obs_dim, obs_groups_spec.get("critic", obs_dim) - - -def get_critic_base_dim(obs_groups_spec: dict[str, int]) -> int: - """Get critic observation dim, falling back to actor obs when absent.""" - critic_dim = obs_groups_spec.get("critic", 0) - return critic_dim if critic_dim > 0 else obs_groups_spec.get("obs", 0) +import warnings + +from unilab.base.observations import ( + flatten_obs_dict, + flatten_policy_obs_dict, + get_critic_base_dim, + get_obs_dims, + split_obs_dict, +) + +warnings.warn( + "`unilab.utils.obs_utils` is deprecated and will be removed in 0.2.0; " + "use `unilab.base.observations` instead.", + DeprecationWarning, + stacklevel=2, +) + +__all__ = [ + "flatten_obs_dict", + "flatten_policy_obs_dict", + "get_critic_base_dim", + "get_obs_dims", + "split_obs_dict", +] diff --git a/src/unilab/utils/offpolicy_logger.py b/src/unilab/utils/offpolicy_logger.py index c160dae36..467d28846 100644 --- a/src/unilab/utils/offpolicy_logger.py +++ b/src/unilab/utils/offpolicy_logger.py @@ -1,469 +1,14 @@ -"""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 -""" - from __future__ import annotations -import time -from collections import deque -from typing import Any - -from rich import box -from rich.console import Group -from rich.panel import Panel -from rich.table import Table - -from unilab.utils.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 - """ - - def __init__( - self, - algo_name: str = "RL", - max_iterations: int = 1500, - num_envs: int = 4096, - env_name: str = "", - obs_dim: int = 0, - action_dim: int = 0, - refresh_per_second: int = 4, - log_dir: str = "", - log_backend: str = "tensorboard", # "tensorboard", "wandb", "none" - wandb_project: str = "unilab", - wandb_entity: str | None = None, - wandb_name: str = "", - wandb_group: str | None = None, - wandb_job_type: str | None = None, - wandb_tags: list[str] | None = None, - wandb_notes: str | None = None, - ): - super().__init__( - algo_name=algo_name, - max_iterations=max_iterations, - num_envs=num_envs, - env_name=env_name, - log_dir=log_dir, - log_backend=log_backend, - wandb_project=wandb_project, - wandb_entity=wandb_entity, - wandb_name=wandb_name, - wandb_group=wandb_group, - wandb_job_type=wandb_job_type, - wandb_tags=wandb_tags, - wandb_notes=wandb_notes, - refresh_per_second=refresh_per_second, - tensorboard_subdir=None, - wandb_config={ - "obs_dim": obs_dim, - "action_dim": action_dim, - "max_iterations": max_iterations, - }, - ) - 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 - self._wait_time: float = 0.0 - self._iter_times: deque = deque(maxlen=50) - self._collector_timing: dict[str, float] = {} - self._timeout_rate: float = 0.0 - self._terminated_rate: float = 0.0 - self._buffer_utilization: float = 0.0 - self._sync_collection: bool = False - self._env_steps_per_sync: int = 0 - self._replay_queue_len: int = 0 - 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}[/]" - - 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 - self._status = f"Buffer fill: {current:,}/{target:,} ({pct:.0f}%)" - 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: - self._reward_history.append(mean_reward) - self._refresh() - - def log_step( - self, - iteration: int, - metrics: dict[str, float] | None = None, - reward: float | None = None, - reward_components: dict[str, float] | None = None, - collect_time: float = 0.0, - train_time: float = 0.0, - wait_time: float = 0.0, - extra_info: dict | None = None, - ): - """Log one training iteration.""" - 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 - ) - - def _backend_log_step( - self, - iteration: int, - metrics: dict[str, float] | None, - reward: float | None, - reward_components: dict[str, float] | None, - 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.) - if metrics: - for k, v in metrics.items(): - w.add_scalar(f"train/{k}", v, global_step) - - # reward/ — reward signals - if reward is not None: - w.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 - 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 - if elapsed > 0 and self._total_steps > 0: - w.add_scalar("perf/steps_per_sec", self._total_steps / elapsed, global_step) - w.add_scalar( - "perf/iter_ms", (self._collect_time + self._train_time) * 1000, global_step - ) - w.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/ - 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/ - 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/ - 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, - 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, - show_header=True, - header_style="bold cyan", - expand=True, - pad_edge=False, - ) - 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)) - - 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, - show_header=True, - header_style="bold blue", - expand=True, - pad_edge=False, - ) - table.add_column("Item", style="white", ratio=2, no_wrap=True) - table.add_column("Value", style="yellow", justify="right", ratio=1, no_wrap=True) - table.add_column("Item", style="white", ratio=2, no_wrap=True) - 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 - wait_ms = self._wait_time * 1000 - wait_color = "red" if wait_ms > 1.0 else "yellow" - table.add_row( - "[dim]learner[/] Wait", - f"[{wait_color}]{wait_ms:.1f}ms[/]", - "[dim]learner[/] Train", - f"{self._train_time * 1000:.1f}ms", - ) - table.add_row( - "[dim]learner[/] Collect", - f"{self._collect_time * 1000:.1f}ms", - "", - "", - ) - 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] - table.add_row( - f"[dim]collector[/] {left_key}", - f"{left_val:.1f}ms", - f"[dim]collector[/] {right_key}", - f"{right_val:.1f}ms", - ) - else: - table.add_row( - f"[dim]collector[/] {left_key}", - f"{left_val:.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}[/]" - else: - util_str = f"[green]{util:.2f}[/]" - table.add_row("Write/Read", util_str, "", "") - - table.add_row( - "Envs", - f"{self.num_envs:,}", - "Sync Collect", - f"{'✓' if self._sync_collection else '✗'} ({self._env_steps_per_sync})" - if self._sync_collection - else "✗", - ) +import warnings - if self._replay_queue_max > 0: - rq_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}[/]", - "", - "", - ) +from unilab.training.logging.offpolicy import OffPolicyLogger - # Steps per second - if elapsed > 0 and self._total_steps > 0: - sps = self._total_steps / elapsed - table.add_row("Steps/s", f"{sps:,.0f}", "", "") +warnings.warn( + "`unilab.utils.offpolicy_logger` is deprecated and will be removed in 0.2.0; " + "use `unilab.training.logging.offpolicy` instead.", + DeprecationWarning, + stacklevel=2, +) - return table +__all__ = ["OffPolicyLogger"] diff --git a/src/unilab/utils/onpolicy_logger.py b/src/unilab/utils/onpolicy_logger.py index f6ddd4b6a..2abe0fe23 100644 --- a/src/unilab/utils/onpolicy_logger.py +++ b/src/unilab/utils/onpolicy_logger.py @@ -1,196 +1,14 @@ from __future__ import annotations -import time -from typing import Any +import warnings -from rich import box -from rich.console import Group -from rich.panel import Panel -from rich.table import Table +from unilab.training.logging.onpolicy import OnPolicyLogger -from unilab.utils.logging_common import BaseTrainingLogger, _fmt_number, _fmt_time, _load_wandb +warnings.warn( + "`unilab.utils.onpolicy_logger` is deprecated and will be removed in 0.2.0; " + "use `unilab.training.logging.onpolicy` instead.", + DeprecationWarning, + stacklevel=2, +) - -class OnPolicyLogger(BaseTrainingLogger): - """Rich logger for on-policy RL (PPO, A2C, etc).""" - - def __init__( - self, - algo_name: str = "PPO", - max_iterations: int = 1500, - num_envs: int = 4096, - num_steps: int = 24, - env_name: str = "", - log_dir: str = "", - log_backend: str = "tensorboard", - wandb_project: str = "unilab", - wandb_entity: str | None = None, - wandb_name: str = "", - wandb_group: str | None = None, - wandb_job_type: str | None = None, - wandb_tags: list[str] | None = None, - wandb_notes: str | None = None, - ): - super().__init__( - algo_name=algo_name, - max_iterations=max_iterations, - num_envs=num_envs, - env_name=env_name, - log_dir=log_dir, - log_backend=log_backend, - wandb_project=wandb_project, - wandb_entity=wandb_entity, - wandb_name=wandb_name, - wandb_group=wandb_group, - wandb_job_type=wandb_job_type, - wandb_tags=wandb_tags, - wandb_notes=wandb_notes, - tensorboard_subdir="tb", - ) - self.num_steps = num_steps - - def start(self, *, status: str = ""): - super().start(status=status) - - def finish(self, *, title: str = "Training Summary", extra_summary: str = ""): - super().finish(title=title, extra_summary=extra_summary) - - def log_step( - self, - iteration: int, - metrics: dict[str, float] | None = None, - reward: float | None = None, - reward_components: dict[str, float] | None = None, - collect_time: float = 0.0, - train_time: float = 0.0, - ): - self._iteration = iteration - self._collect_time = collect_time - self._train_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._refresh() - self._backend_log_step(iteration, metrics, reward, reward_components) - - def _backend_log_step( - self, - iteration: int, - metrics: dict[str, float] | None, - reward: float | None, - reward_components: dict[str, float] | None, - ): - if self._tb_writer: - w = self._tb_writer - if metrics: - for k, v in metrics.items(): - w.add_scalar(f"train/{k}", v, iteration) - if reward is not None: - w.add_scalar("reward/mean", reward, iteration) - if reward_components: - for k, v in reward_components.items(): - w.add_scalar(f"reward/{k}", v, iteration) - if self._mean_ep_length > 0: - w.add_scalar("episode/length", self._mean_ep_length, iteration) - w.add_scalar("perf/collect_time_ms", self._collect_time * 1000, iteration) - w.add_scalar("perf/train_time_ms", self._train_time * 1000, iteration) - - 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 - 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 - if self._mean_ep_length > 0: - log_dict["episode/length"] = self._mean_ep_length - log_dict["perf/collect_time_ms"] = self._collect_time * 1000 - log_dict["perf/train_time_ms"] = self._train_time * 1000 - wandb.log(log_dict, step=iteration) - - def _build_display(self) -> Panel: - header_panel = self._build_header(include_status=False) - - 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, - title="[bold] 🚀 UniLab On-Policy Training [/]", - border_style="bright_blue", - padding=(0, 1), - ) - - def _build_metrics_table(self) -> Table: - table = Table( - title="[bold]Policy Metrics[/]", - box=box.SIMPLE_HEAVY, - show_header=True, - header_style="bold cyan", - expand=True, - pad_edge=False, - ) - 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...[/]", "") - else: - for k in sorted(self._latest_metrics.keys()): - v = self._latest_metrics[k] - name = k.replace("_", " ").title() - table.add_row(name, _fmt_number(v)) - - return table - - def _build_reward_table(self) -> Table: - return self._build_reward_table_common(wait_message="[dim]Waiting...[/]") - - def _build_timing_table(self) -> Table: - table = Table( - title="[bold]Timing[/]", - box=box.SIMPLE_HEAVY, - show_header=True, - header_style="bold blue", - expand=True, - pad_edge=False, - ) - table.add_column("Item", style="white", ratio=1) - table.add_column("Value", style="yellow", justify="right", ratio=1) - table.add_column("Item", style="white", ratio=1) - table.add_column("Value", style="yellow", justify="right", ratio=1) - - elapsed = time.time() - self._start_time if self._start_time else 0 - iter_time = self._collect_time + self._train_time - fps = int(self.num_envs * self.num_steps / max(iter_time, 1e-8)) if iter_time > 0 else 0 - - table.add_row("Elapsed", _fmt_time(elapsed), "Envs", f"{self.num_envs:,}") - table.add_row( - "Collect", - f"{self._collect_time * 1000:.1f}ms", - "Train", - f"{self._train_time * 1000:.1f}ms", - ) - table.add_row("Iter Time", f"{iter_time * 1000:.1f}ms", "Steps/s", f"{fps:,}") - - return table +__all__ = ["OnPolicyLogger"] diff --git a/src/unilab/utils/render_many.py b/src/unilab/utils/render_many.py index 9b22f2085..ac630630b 100644 --- a/src/unilab/utils/render_many.py +++ b/src/unilab/utils/render_many.py @@ -1,606 +1,12 @@ -"""MuJoCo-only batched rendering helpers. +from __future__ import annotations -This module renders many MuJoCo states into image frames by constructing -MuJoCo model/data/renderer objects inside worker processes. It is not available -for Motrix-only workflows. -""" +import warnings -import math -import os -import subprocess -import sys -import textwrap -from collections.abc import Sequence -from typing import Any +from unilab.visualization.render_many import * # noqa: F403 -import imageio - -_USER_MUJOCO_GL = os.environ.get("MUJOCO_GL") - -_EGL_PROBE_SCRIPT = textwrap.dedent( - ''' - import mujoco - - xml = """ - - - - - - """ - - model = mujoco.MjModel.from_xml_string(xml) - data = mujoco.MjData(model) - renderer = mujoco.Renderer(model, height=8, width=8) - mujoco.mj_forward(model, data) - renderer.update_scene(data) - renderer.render() - renderer.close() - ''' +warnings.warn( + "`unilab.utils.render_many` is deprecated and will be removed in 0.2.0; " + "use `unilab.visualization.render_many` instead.", + DeprecationWarning, + stacklevel=2, ) - - -def _egl_runtime_usable() -> bool: - env = os.environ.copy() - env["MUJOCO_GL"] = "egl" - env.setdefault("MUJOCO_EGL_DEVICE_ID", "0") - - try: - subprocess.run( - [sys.executable, "-c", _EGL_PROBE_SCRIPT], - env=env, - check=True, - stdout=subprocess.DEVNULL, - stderr=subprocess.DEVNULL, - timeout=10, - ) - except (OSError, subprocess.SubprocessError): - return False - - os.environ.setdefault("MUJOCO_EGL_DEVICE_ID", env["MUJOCO_EGL_DEVICE_ID"]) - return True - - -def _resolve_gl_backend() -> str: - """Pick a valid MUJOCO_GL backend for the current platform. - - Respects an explicit user setting unless it's provably invalid (e.g. egl - on macOS). Falls back to glfw when EGL is requested but not available. - """ - current = os.environ.get("MUJOCO_GL", "") - safe_values = {"glfw", "osmesa", "disabled"} - - if sys.platform == "darwin": - # macOS has no EGL support; glfw is the only off-screen option - return current if current in safe_values else "glfw" - - # Linux / other: honour explicit non-egl choices supplied before import. - if current in safe_values and current == _USER_MUJOCO_GL: - return current - - # Probe EGL by creating a tiny MuJoCo renderer in a clean subprocess. - if _egl_runtime_usable(): - return "egl" - - return "glfw" - - -# Must be set *before* importing mujoco (it reads the var at import time) -os.environ["MUJOCO_GL"] = _resolve_gl_backend() - -import mujoco # noqa: E402 -import numpy as np - - -def get_grid_offsets(num_envs, spacing=1.0): - rows = int(math.ceil(math.sqrt(num_envs))) - cols = int(math.ceil(num_envs / rows)) - offsets = np.zeros((num_envs, 2)) - for i in range(num_envs): - r = i // cols - c = i % cols - offsets[i, 0] = r * spacing - offsets[i, 1] = c * spacing - return offsets - - -# Worker global context -_worker_ctx: dict[str, Any] = {} - - -def _close_worker(): - """Explicitly close the renderer in the worker context.""" - if "renderer" in _worker_ctx: - _worker_ctx["renderer"].close() - - -def init_worker(model_path, shape): - """Initialize MuJoCo-only rendering context for a worker process.""" - import atexit - - def _load_model(path_like): - path = str(path_like) - loader = ( - mujoco.MjModel.from_binary_path - if path.endswith(".mjb") - else mujoco.MjModel.from_xml_path - ) - return loader(path) - - if isinstance(model_path, Sequence) and not isinstance(model_path, (str, bytes, os.PathLike)): - models = [_load_model(path) for path in model_path] - else: - models = [_load_model(model_path)] - - for model in models: - model.vis.global_.offwidth = 3840 - model.vis.global_.offheight = 2160 - - _worker_ctx["models"] = models - _worker_ctx["data_list"] = [mujoco.MjData(model) for model in models] - _worker_ctx["renderer"] = mujoco.Renderer(models[0], height=shape[1], width=shape[0]) - atexit.register(_close_worker) - - -def render_frame_job(args): - """ - Worker function to render a single frame. - args: (state_batch, offsets, transparent, cam_distance, cam_elevation, cam_azimuth, cam_lookat) - """ - state_batch, offsets, transparent, cam_distance, cam_elevation, cam_azimuth, cam_lookat = args - - models = _worker_ctx["models"] - data_list = _worker_ctx["data_list"] - renderer = _worker_ctx["renderer"] - - # Visual options - vopt = mujoco.MjvOption() - vopt.flags[mujoco.mjtVisFlag.mjVIS_TRANSPARENT] = transparent - pert = mujoco.MjvPerturb() - catmask_dynamic = mujoco.mjtCatBit.mjCAT_DYNAMIC - catmask_static = mujoco.mjtCatBit.mjCAT_STATIC - - # Helper to set state - def set_state(model, d, s, offset=None): - d.time = s[0] - d.qpos[:] = s[1 : 1 + model.nq] - d.qvel[:] = s[1 + model.nq : 1 + model.nq + model.nv] - - apply_root_offset = False - - if offset is not None: - # Check if Root (Body 1) has a free joint or slide joints allowing X/Y movement - # Body 0 is world. Body 1 is usually the robot base. - robot_moved = False - - # Heuristic: Check joint at qpos 0, 1. - # If jnt_type[0] is free (0), fine. - # If jnt_type[0] is slide (2) and axis is x/y... - - # Better check: Does the first body have a joint? - first_body_jnt = model.body_jntadr[1] if model.nbody > 1 else -1 - if first_body_jnt >= 0: - jnt_type = model.jnt_type[first_body_jnt] - # mjJNT_FREE=0 - if jnt_type == 0: - d.qpos[0] += offset[0] - d.qpos[1] += offset[1] - robot_moved = True - - # If robot wasn't moved via qpos, we need to manually offset geometries later - if not robot_moved: - apply_root_offset = True - - # 2. Box offset - box_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "box") - if box_id >= 0: - jnt_adr = model.body_jntadr[box_id] - if jnt_adr >= 0: - qpos_adr = model.jnt_qposadr[jnt_adr] - d.qpos[qpos_adr] += offset[0] - d.qpos[qpos_adr + 1] += offset[1] - - # 3. Target offset (target_x, target_y) - target_x = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_JOINT, "target_x") - if target_x >= 0: - d.qpos[model.jnt_qposadr[target_x]] += offset[0] - - target_y = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_JOINT, "target_y") - if target_y >= 0: - d.qpos[model.jnt_qposadr[target_y]] += offset[1] - - mujoco.mj_forward(model, d) - - # Post-process: Shift all geometries if robot root wasn't moved - if apply_root_offset and offset is not None: - # Shift all geoms? - # We should shift Everything that is PART OF THE ROBOT. - # Or just everything? - # Box and Target were already shifted via qpos. - # BUT qpos shift updates body_pos which updates geom_pos. - # If we shift ALL geom_pos, we double shift Box and Target! - - # So we need to shift geoms that belong to bodies which are NOT Box or Target. - # Or simpler: Shift everything, but subtract offset from Box/Target qpos first? No. - - # Let's iterate bodies. - # Simple heuristic: Shift everything except Box and Target? - box_body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "box") - target_body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "mocap_target") - - # Also target might be just a body named "mocap_target" - - for i in range(model.ngeom): - body_id = model.geom_bodyid[i] - # If it is robot body. - # We want to shift generally everything that wasn't shifted by Qpos. - # Box and Target were shifted by Qpos. - # Floor (Plane) should usually NOT be shifted (infinite). - # Everything else (Robot Base, Robot Links, Decoration) should be shifted. - - is_box_or_target = (body_id == box_body_id) or (body_id == target_body_id) - is_plane = model.geom_type[i] == mujoco.mjtGeom.mjGEOM_PLANE - - if not is_box_or_target and not is_plane: - d.geom_xpos[i, 0] += offset[0] - d.geom_xpos[i, 1] += offset[1] - - # Also update site positions if they are visualized - for i in range(model.nsite): - body_id = model.site_bodyid[i] - - is_box_or_target = (body_id == box_body_id) or (body_id == target_body_id) - - if not is_box_or_target: - d.site_xpos[i, 0] += offset[0] - d.site_xpos[i, 1] += offset[1] - - num_envs = state_batch.shape[0] - - # 1. Clear/Init Scene - primary_model = models[0] - primary_data = data_list[0] - set_state( - primary_model, primary_data, state_batch[0], offsets[0] if offsets is not None else None - ) - - # Init Camera - cam = mujoco.MjvCamera() - if offsets is not None: - center_x = np.mean(offsets[:, 0]) - center_y = np.mean(offsets[:, 1]) - if cam_lookat is None: - cam.lookat = [center_x, center_y, 0.75] - else: - cam.lookat = [float(cam_lookat[0]), float(cam_lookat[1]), float(cam_lookat[2])] - cam.distance = cam_distance - cam.elevation = cam_elevation - cam.azimuth = cam_azimuth - cam.type = mujoco.mjtCamera.mjCAMERA_FREE - else: - cam.type = mujoco.mjtCamera.mjCAMERA_FREE - - renderer.update_scene(primary_data, camera=cam, scene_option=vopt) - - # 2. Add other robots - for i in range(1, num_envs): - model = models[min(i, len(models) - 1)] - data = data_list[min(i, len(data_list) - 1)] - set_state(model, data, state_batch[i], offsets[i] if offsets is not None else None) - mujoco.mjv_addGeoms(model, data, vopt, pert, catmask_dynamic, renderer.scene) - - # Avoid duplicating world floor/static group-0 geoms across envs, while still - # adding fixed-base static visuals (e.g., Sharpa base link visual). - geomgroup0 = int(vopt.geomgroup[0]) - vopt.geomgroup[0] = 0 - mujoco.mjv_addGeoms(model, data, vopt, pert, catmask_static, renderer.scene) - vopt.geomgroup[0] = geomgroup0 - - return renderer.render() - - -def render_states_get_frames( - state_list, - model_path, - width=1280, - height=720, - num_processes=8, - camera_id=-1, - cam_distance=2.0, - cam_elevation=-20, - cam_azimuth=90, - cam_lookat=None, - render_spacing=1.0, -): - """ - Render a list of physics states and return the list of frames. - - Args: - state_list: List of numpy arrays, each shape (num_envs, state_dim). - model_path: Path to the mujoco XML model file. - width: Width of the video. - height: Height of the video. - num_processes: Number of parallel processes to use. - camera_id: Camera ID to render from. - cam_distance: Camera distance from lookat point. - cam_elevation: Camera elevation angle in degrees. - cam_azimuth: Camera azimuth angle in degrees. - cam_lookat: Optional [x, y, z] lookat override for the free camera. - render_spacing: Grid spacing used to offset each env in composed video frames. - Returns: - List of numpy arrays (H, W, 3) (RGB) - """ - if not state_list: - print("No states to render.") - return [] - - num_envs = state_list[0].shape[0] - offsets = get_grid_offsets(num_envs, spacing=render_spacing) - shape = (width, height) - - print( - f"Rendering {len(state_list)} frames for {num_envs} envs with {num_processes} processes..." - ) - - # Prepare arguments for each frame - tasks = [ - (s, offsets, False, cam_distance, cam_elevation, cam_azimuth, cam_lookat) - for s in state_list - ] - - frames = [] - - if num_processes <= 1: - # Serial execution - # Initialize context manually - init_worker(model_path, shape) - try: - for task in tasks: - res = render_frame_job(task) - frames.append(res) - finally: - _close_worker() - else: - # Use multiprocessing Pool - # On macOS, use spawn to avoid forking OpenGL/MuJoCo contexts. - import multiprocessing - - ctx = multiprocessing.get_context("spawn") - with ctx.Pool( - processes=num_processes, initializer=init_worker, initargs=(model_path, shape) - ) as pool: - results = pool.map(render_frame_job, tasks) - frames.extend(results) - - return frames - - -def _get_nearest_env_indices(offsets, primary_idx, max_extra): - """Return indices of the *max_extra* environments closest to *primary_idx*.""" - if len(offsets) <= 1 + max_extra: - return [i for i in range(len(offsets)) if i != primary_idx] - primary = offsets[primary_idx] - dists = np.linalg.norm(offsets - primary, axis=1) - dists[primary_idx] = np.inf # exclude self - return list(np.argsort(dists)[:max_extra]) - - -def render_frame_tracking_job(args): - """Render a single frame with camera tracking on the primary env's root body. - - The camera uses ``mjCAMERA_TRACKING`` so it follows the robot each frame. - Only the primary env + nearest neighbours are rendered. - """ - ( - state_batch, - offsets, - env_indices, - primary_local_idx, - cam_distance, - cam_elevation, - cam_azimuth, - ) = args - - models = _worker_ctx["models"] - data_list = _worker_ctx["data_list"] - renderer = _worker_ctx["renderer"] - - vopt = mujoco.MjvOption() - pert = mujoco.MjvPerturb() - catmask_dynamic = mujoco.mjtCatBit.mjCAT_DYNAMIC - catmask_static = mujoco.mjtCatBit.mjCAT_STATIC - - def set_state(model, d, s, offset=None): - d.time = s[0] - d.qpos[:] = s[1 : 1 + model.nq] - d.qvel[:] = s[1 + model.nq : 1 + model.nq + model.nv] - - apply_root_offset = False - - if offset is not None: - robot_moved = False - first_body_jnt = model.body_jntadr[1] if model.nbody > 1 else -1 - if first_body_jnt >= 0: - jnt_type = model.jnt_type[first_body_jnt] - if jnt_type == 0: # mjJNT_FREE - d.qpos[0] += offset[0] - d.qpos[1] += offset[1] - robot_moved = True - - if not robot_moved: - apply_root_offset = True - - box_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "box") - if box_id >= 0: - jnt_adr = model.body_jntadr[box_id] - if jnt_adr >= 0: - qpos_adr = model.jnt_qposadr[jnt_adr] - d.qpos[qpos_adr] += offset[0] - d.qpos[qpos_adr + 1] += offset[1] - - target_x = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_JOINT, "target_x") - if target_x >= 0: - d.qpos[model.jnt_qposadr[target_x]] += offset[0] - - target_y = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_JOINT, "target_y") - if target_y >= 0: - d.qpos[model.jnt_qposadr[target_y]] += offset[1] - - mujoco.mj_forward(model, d) - - if apply_root_offset and offset is not None: - box_body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "box") - target_body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "mocap_target") - - for i in range(model.ngeom): - body_id = model.geom_bodyid[i] - is_box_or_target = (body_id == box_body_id) or (body_id == target_body_id) - is_plane = model.geom_type[i] == mujoco.mjtGeom.mjGEOM_PLANE - - if not is_box_or_target and not is_plane: - d.geom_xpos[i, 0] += offset[0] - d.geom_xpos[i, 1] += offset[1] - - for i in range(model.nsite): - body_id = model.site_bodyid[i] - is_box_or_target = (body_id == box_body_id) or (body_id == target_body_id) - if not is_box_or_target: - d.site_xpos[i, 0] += offset[0] - d.site_xpos[i, 1] += offset[1] - - # Primary env first — camera tracks body 1 of this env - primary_global = env_indices[primary_local_idx] - primary_model = models[min(primary_global, len(models) - 1)] - primary_data = data_list[min(primary_global, len(data_list) - 1)] - set_state( - primary_model, - primary_data, - state_batch[primary_global], - offsets[primary_global] if offsets is not None else None, - ) - - cam = mujoco.MjvCamera() - cam.type = mujoco.mjtCamera.mjCAMERA_TRACKING - cam.trackbodyid = 1 # robot root body - cam.distance = cam_distance - cam.elevation = cam_elevation - cam.azimuth = cam_azimuth - - renderer.update_scene(primary_data, camera=cam, scene_option=vopt) - - # Add neighbour envs as background context - for local_i, global_i in enumerate(env_indices): - if local_i == primary_local_idx: - continue - model = models[min(global_i, len(models) - 1)] - data = data_list[min(global_i, len(data_list) - 1)] - set_state( - model, data, state_batch[global_i], offsets[global_i] if offsets is not None else None - ) - mujoco.mjv_addGeoms(model, data, vopt, pert, catmask_dynamic, renderer.scene) - - geomgroup0 = int(vopt.geomgroup[0]) - vopt.geomgroup[0] = 0 - mujoco.mjv_addGeoms(model, data, vopt, pert, catmask_static, renderer.scene) - vopt.geomgroup[0] = geomgroup0 - - return renderer.render() - - -def render_states_get_frames_tracking( - state_list, - model_path, - width=1280, - height=720, - tracking_env_idx=0, - max_extra_envs=2, - cam_distance=2.0, - cam_elevation=-20, - cam_azimuth=90, - render_spacing=1.0, -): - """Render with camera tracking on a single primary environment. - - Only the primary env and its nearest neighbours are shown. The camera - follows the root body of the primary env each frame (``mjCAMERA_TRACKING``). - - Args: - state_list: List of numpy arrays, each shape (num_envs, state_dim). - model_path: Path to the mujoco XML model file. - tracking_env_idx: Index of the primary environment to track. - max_extra_envs: Number of nearest-neighbour envs to render alongside. - cam_distance: Camera distance from the tracked body. - cam_elevation: Camera elevation angle in degrees. - cam_azimuth: Camera azimuth angle in degrees. - render_spacing: Grid spacing for env layout. - """ - if not state_list: - print("No states to render.") - return [] - - num_envs = state_list[0].shape[0] - offsets = get_grid_offsets(num_envs, spacing=render_spacing) - shape = (width, height) - - tracking_env_idx = min(tracking_env_idx, num_envs - 1) - neighbour_indices = _get_nearest_env_indices(offsets, tracking_env_idx, max_extra_envs) - env_indices = [tracking_env_idx] + neighbour_indices - primary_local_idx = 0 # primary is always first in env_indices - - total_shown = len(env_indices) - print( - f"Rendering {len(state_list)} frames (tracking env {tracking_env_idx} " - f"+ {total_shown - 1} neighbours) ..." - ) - - tasks = [ - (s, offsets, env_indices, primary_local_idx, cam_distance, cam_elevation, cam_azimuth) - for s in state_list - ] - - # Camera tracking changes each frame so multiprocessing gives inconsistent - # results when workers don't share state. Default to serial. - frames = [] - init_worker(model_path, shape) - try: - for task in tasks: - frames.append(render_frame_tracking_job(task)) - finally: - _close_worker() - - return frames - - -def render_states_to_video( - state_list, - model_path, - output_path, - fps=30, - width=1280, - height=720, - num_processes=8, - cam_distance=2.0, - cam_elevation=-20, - cam_azimuth=90, - cam_lookat=None, - render_spacing=1.0, -): - """ - Render a list of physics states to a video file using parallel processing. - """ - frames = render_states_get_frames( - state_list, - model_path, - width, - height, - num_processes, - cam_distance=cam_distance, - cam_elevation=cam_elevation, - cam_azimuth=cam_azimuth, - cam_lookat=cam_lookat, - render_spacing=render_spacing, - ) - - print(f"Saving video to {output_path}...") - imageio.mimsave(output_path, frames, fps=fps) - print("Done!") diff --git a/src/unilab/utils/reward_utils.py b/src/unilab/utils/reward_utils.py index ee68d18fc..091ef256d 100644 --- a/src/unilab/utils/reward_utils.py +++ b/src/unilab/utils/reward_utils.py @@ -1,54 +1,14 @@ -"""Utility functions for reward config handling.""" +from __future__ import annotations -from typing import Any, cast +import warnings -from omegaconf import DictConfig, OmegaConf +from unilab.config.reward import RewardDict, extract_reward_config, resolve_reward_dict -RewardDict = dict[str, Any] +warnings.warn( + "`unilab.utils.reward_utils` is deprecated and will be removed in 0.2.0; " + "use `unilab.config.reward` instead.", + DeprecationWarning, + stacklevel=2, +) - -def _to_reward_dict(value: object, *, error_message: str) -> RewardDict: - """Convert an OmegaConf container into a plain reward dictionary.""" - resolved = OmegaConf.to_container(value, resolve=True) - if not isinstance(resolved, dict): - raise ValueError(error_message) - # Some reward configs are mounted as a full `reward:` section. - # Env config injection expects the inner reward mapping. - if set(resolved) == {"reward"} and isinstance(resolved["reward"], dict): - return cast(RewardDict, resolved["reward"]) - return cast(RewardDict, resolved) - - -def resolve_reward_dict(cfg: DictConfig) -> RewardDict: - """Resolve the reward config from the final composed config.""" - reward_cfg = OmegaConf.select(cfg, "reward") - if not reward_cfg: - raise ValueError( - "Missing 'reward' config in Hydra. Reward config must be explicitly provided." - ) - - reward_dict = _to_reward_dict( - reward_cfg, - error_message="Reward config must resolve to a mapping.", - ) - if not reward_dict: - raise ValueError( - "Reward config resolved to empty. Please select a non-default reward override." - ) - - return reward_dict - - -def extract_reward_config(cfg: DictConfig) -> dict[str, RewardDict]: - """Extract and validate reward config from Hydra config. - - Args: - cfg: Hydra DictConfig containing reward section - - Returns: - Dictionary with reward_config key for env_cfg_override - - Raises: - ValueError: If reward config is missing - """ - return {"reward_config": resolve_reward_dict(cfg)} +__all__ = ["RewardDict", "extract_reward_config", "resolve_reward_dict"] diff --git a/src/unilab/utils/rsl_rl_compat.py b/src/unilab/utils/rsl_rl_compat.py index fb17346e8..4f5f5485b 100644 --- a/src/unilab/utils/rsl_rl_compat.py +++ b/src/unilab/utils/rsl_rl_compat.py @@ -1,230 +1,12 @@ -""" -Compatibility utilities for supporting both rsl_rl 3.x and 4.x. +from __future__ import annotations -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 +import warnings -This module provides runtime version detection and config conversion so that -the codebase can work with both versions without code duplication. -""" +from unilab.algos.torch.rsl_rl.compat import * # noqa: F403 -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 +warnings.warn( + "`unilab.utils.rsl_rl_compat` is deprecated and will be removed in 0.2.0; " + "use `unilab.algos.torch.rsl_rl.compat` instead.", + DeprecationWarning, + stacklevel=2, +) diff --git a/src/unilab/utils/rsl_rl_vec_env_wrapper.py b/src/unilab/utils/rsl_rl_vec_env_wrapper.py index 6748a6cd0..6c805cd43 100644 --- a/src/unilab/utils/rsl_rl_vec_env_wrapper.py +++ b/src/unilab/utils/rsl_rl_vec_env_wrapper.py @@ -1,166 +1,14 @@ -"""Shared RSL-RL vectorized environment wrapper. +from __future__ import annotations -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 warnings -import numpy as np -import torch -from tensordict import TensorDict +from unilab.algos.torch.rsl_rl.vec_env_wrapper import RslRlVecEnvWrapper -from unilab.utils.obs_utils import flatten_policy_obs_dict -from unilab.utils.torch_utils import to_torch +warnings.warn( + "`unilab.utils.rsl_rl_vec_env_wrapper` is deprecated and will be removed in 0.2.0; " + "use `unilab.algos.torch.rsl_rl.vec_env_wrapper` instead.", + DeprecationWarning, + stacklevel=2, +) - -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) +__all__ = ["RslRlVecEnvWrapper"] diff --git a/src/unilab/utils/run_utils.py b/src/unilab/utils/run_utils.py index 34817f2bc..0f92bad83 100644 --- a/src/unilab/utils/run_utils.py +++ b/src/unilab/utils/run_utils.py @@ -1,35 +1,30 @@ -import os +from __future__ import annotations +import warnings -def get_latest_run(log_dir: str) -> str | None: - """Find the latest run in the log directory that contains a model. +from unilab.training.run import ( + get_entrypoint_log_root, + get_latest_checkpoint, + get_latest_run, + get_log_root, + parse_checkpoint_path, + resolve_checkpoint_path, + resolve_task_checkpoint_path, +) - Args: - log_dir: Path to the base log directory (e.g., logs/fast_sac_Go2LocoFlatTerrain) +warnings.warn( + "`unilab.utils.run_utils` is deprecated and will be removed in 0.2.0; " + "use `unilab.training.run` instead.", + DeprecationWarning, + stacklevel=2, +) - 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 +__all__ = [ + "get_entrypoint_log_root", + "get_latest_checkpoint", + "get_latest_run", + "get_log_root", + "parse_checkpoint_path", + "resolve_checkpoint_path", + "resolve_task_checkpoint_path", +] diff --git a/src/unilab/utils/tensor.py b/src/unilab/utils/tensor.py new file mode 100644 index 000000000..8e45d5025 --- /dev/null +++ b/src/unilab/utils/tensor.py @@ -0,0 +1,33 @@ +"""Generic array <-> torch conversion utilities.""" + +from __future__ import annotations + +import numpy as np +import torch + + +def to_torch(x, device: str | torch.device) -> torch.Tensor: + """Convert numpy-like input to torch on the target device. + + Supports torch tensors, numpy arrays, and any array exposing ``__dlpack__``. + """ + if isinstance(x, torch.Tensor): + return x.to(device) + if isinstance(x, np.ndarray): + return torch.from_numpy(x).to(device) + try: + if hasattr(x, "__dlpack__"): + return torch.from_dlpack(x).to(device) # pyright: ignore[reportPrivateImportUsage] + except Exception: + pass + arr = np.asarray(x, dtype=np.float32) + return torch.from_numpy(arr).to(device) + + +def to_numpy(x) -> np.ndarray: + """Convert torch tensor or numpy-like input to numpy.""" + if isinstance(x, np.ndarray): + return x + if isinstance(x, torch.Tensor): + return x.detach().cpu().numpy() + return np.asarray(x) diff --git a/src/unilab/utils/torch_utils.py b/src/unilab/utils/torch_utils.py index 8e45d5025..41697153a 100644 --- a/src/unilab/utils/torch_utils.py +++ b/src/unilab/utils/torch_utils.py @@ -1,33 +1,14 @@ -"""Generic array <-> torch conversion utilities.""" - from __future__ import annotations -import numpy as np -import torch - - -def to_torch(x, device: str | torch.device) -> torch.Tensor: - """Convert numpy-like input to torch on the target device. +import warnings - Supports torch tensors, numpy arrays, and any array exposing ``__dlpack__``. - """ - if isinstance(x, torch.Tensor): - return x.to(device) - if isinstance(x, np.ndarray): - return torch.from_numpy(x).to(device) - try: - if hasattr(x, "__dlpack__"): - return torch.from_dlpack(x).to(device) # pyright: ignore[reportPrivateImportUsage] - except Exception: - pass - arr = np.asarray(x, dtype=np.float32) - return torch.from_numpy(arr).to(device) +from unilab.utils.tensor import to_numpy, to_torch +warnings.warn( + "`unilab.utils.torch_utils` is deprecated and will be removed in 0.2.0; " + "use `unilab.utils.tensor` instead.", + DeprecationWarning, + stacklevel=2, +) -def to_numpy(x) -> np.ndarray: - """Convert torch tensor or numpy-like input to numpy.""" - if isinstance(x, np.ndarray): - return x - if isinstance(x, torch.Tensor): - return x.detach().cpu().numpy() - return np.asarray(x) +__all__ = ["to_numpy", "to_torch"] diff --git a/src/unilab/utils/viser_scene.py b/src/unilab/utils/viser_scene.py index f6fbaf858..bda072393 100644 --- a/src/unilab/utils/viser_scene.py +++ b/src/unilab/utils/viser_scene.py @@ -1,281 +1,18 @@ -# pyright: reportMissingImports=false -"""MuJoCo-to-viser scene adapter for interactive web-based 3D visualization. - -This module renders MuJoCo scenes via a viser web server, providing browser-based -interactive 3D viewing without requiring a local display or GLFW. It is gated -behind the ``viser`` optional-dependency group and is **not** imported by default. - -Usage (from ``scripts/play_viser.py``):: - - from unilab.utils.viser_scene import MujocoViserScene, VISER_AVAILABLE -""" - from __future__ import annotations -import math -from typing import Any - -import mujoco -import numpy as np - -try: - import trimesh - import viser - - VISER_AVAILABLE = True -except ImportError: - VISER_AVAILABLE = False - - -# --------------------------------------------------------------------------- # -# Rotation helpers (pure numpy, no scipy dependency) # -# --------------------------------------------------------------------------- # - - -def _rotmat_to_wxyz(mat: np.ndarray) -> tuple[float, float, float, float]: - """Convert a 3x3 rotation matrix to a (w, x, y, z) quaternion.""" - m = np.asarray(mat, dtype=np.float64).reshape(3, 3) - trace = m[0, 0] + m[1, 1] + m[2, 2] - - if trace > 0: - s = 0.5 / math.sqrt(trace + 1.0) - w = 0.25 / s - x = (m[2, 1] - m[1, 2]) * s - y = (m[0, 2] - m[2, 0]) * s - z = (m[1, 0] - m[0, 1]) * s - elif m[0, 0] > m[1, 1] and m[0, 0] > m[2, 2]: - s = 2.0 * math.sqrt(1.0 + m[0, 0] - m[1, 1] - m[2, 2]) - w = (m[2, 1] - m[1, 2]) / s - x = 0.25 * s - y = (m[0, 1] + m[1, 0]) / s - z = (m[0, 2] + m[2, 0]) / s - elif m[1, 1] > m[2, 2]: - s = 2.0 * math.sqrt(1.0 + m[1, 1] - m[0, 0] - m[2, 2]) - w = (m[0, 2] - m[2, 0]) / s - x = (m[0, 1] + m[1, 0]) / s - y = 0.25 * s - z = (m[1, 2] + m[2, 1]) / s - else: - s = 2.0 * math.sqrt(1.0 + m[2, 2] - m[0, 0] - m[1, 1]) - w = (m[1, 0] - m[0, 1]) / s - x = (m[0, 2] + m[2, 0]) / s - y = (m[1, 2] + m[2, 1]) / s - z = 0.25 * s - - return (float(w), float(x), float(y), float(z)) - - -# --------------------------------------------------------------------------- # -# Geometry extraction helpers # -# --------------------------------------------------------------------------- # - - -def _rgba_to_color(rgba: np.ndarray) -> tuple[int, int, int]: - """Convert MuJoCo float RGBA [0,1] to viser int RGB [0,255].""" - return ( - int(np.clip(rgba[0] * 255, 0, 255)), - int(np.clip(rgba[1] * 255, 0, 255)), - int(np.clip(rgba[2] * 255, 0, 255)), - ) - - -def _rgba_to_opacity(rgba: np.ndarray) -> float: - return float(np.clip(rgba[3], 0.0, 1.0)) - - -def _extract_mesh(model: mujoco.MjModel, geom_dataid: int) -> tuple[np.ndarray, np.ndarray]: - """Extract vertices and faces for a MuJoCo mesh geom.""" - vert_adr = model.mesh_vertadr[geom_dataid] - vert_num = model.mesh_vertnum[geom_dataid] - face_adr = model.mesh_faceadr[geom_dataid] - face_num = model.mesh_facenum[geom_dataid] - - vertices = model.mesh_vert[vert_adr : vert_adr + vert_num].copy() - faces = model.mesh_face[face_adr : face_adr + face_num].copy() - return vertices, faces - - -def build_visible_env_indices(num_envs: int, visible_envs: int) -> np.ndarray: - """Select a stable subset of env indices spread across the full batch. - - Args: - num_envs: Total number of runtime environments. - visible_envs: Number of env slots exposed in the viewer. - - Returns: - A monotonically increasing array of runtime env indices. - """ - if visible_envs <= 0: - raise ValueError(f"visible_envs must be positive, got {visible_envs}") - if visible_envs >= num_envs: - return np.arange(num_envs, dtype=np.int32) - return np.floor(np.linspace(0, num_envs, visible_envs, endpoint=False)).astype(np.int32) - - -# --------------------------------------------------------------------------- # -# MujocoViserScene # -# --------------------------------------------------------------------------- # - - -class MujocoViserScene: - """Bridges a ``mujoco.MjModel`` to a ``viser.ViserServer`` scene graph. - - Call :meth:`build` once to populate the scene with geometry handles, then - call :meth:`update` each frame to sync body transforms from ``MjData``. - """ - - def __init__( - self, - server: Any, - model: mujoco.MjModel, - *, - name_prefix: str = "/mujoco", - position_offset: tuple[float, float, float] = (0.0, 0.0, 0.0), - render_plane: bool = True, - ) -> None: - if not VISER_AVAILABLE: - raise ImportError("viser is not installed. Install with: uv sync --extra viser") - self._server: viser.ViserServer = server - self._model = model - self._name_prefix = name_prefix.rstrip("/") or "/mujoco" - self._position_offset = np.asarray(position_offset, dtype=np.float64) - self._render_plane = bool(render_plane) - self._handles: dict[int, Any] = {} - self._build() - - def reset( - self, - model: mujoco.MjModel, - *, - position_offset: tuple[float, float, float] | None = None, - render_plane: bool | None = None, - ) -> None: - """Rebuild the viser scene for a new MuJoCo model. - - Args: - model: MuJoCo model whose geoms should populate the scene. - position_offset: Optional XYZ offset applied to all geoms. - render_plane: Optional override for whether plane geoms should be built. - - Returns: - None. - """ - self.close() - self._model = model - if position_offset is not None: - self._position_offset = np.asarray(position_offset, dtype=np.float64) - if render_plane is not None: - self._render_plane = bool(render_plane) - self._build() - - def close(self) -> None: - """Remove all scene handles owned by this adapter.""" - for handle in self._handles.values(): - handle.remove() - self._handles.clear() - - # ------------------------------------------------------------------ # - # Scene construction # - # ------------------------------------------------------------------ # - - def _build(self) -> None: - """Create viser scene nodes for every MuJoCo geom.""" - model = self._model - server = self._server - - server.scene.set_up_direction("+z") - - for i in range(model.ngeom): - geom_type = model.geom_type[i] - size = model.geom_size[i] - rgba = model.geom_rgba[i] - color = _rgba_to_color(rgba) - opacity = _rgba_to_opacity(rgba) - name = f"{self._name_prefix}/geom/{i}" - - handle: Any | None = None - - if geom_type == mujoco.mjtGeom.mjGEOM_PLANE: - if not self._render_plane: - continue - # Render ground plane as a grid - plane_size = float(size[0]) if size[0] > 0 else 10.0 - handle = server.scene.add_grid( - name, - width=plane_size * 2, - height=plane_size * 2, - cell_size=0.5, - ) - - elif geom_type == mujoco.mjtGeom.mjGEOM_SPHERE: - handle = server.scene.add_icosphere( - name, - radius=float(size[0]), - color=color, - opacity=opacity, - ) - - elif geom_type == mujoco.mjtGeom.mjGEOM_CAPSULE: - half_len = float(size[1]) - radius = float(size[0]) - mesh = trimesh.creation.capsule(height=half_len * 2, radius=radius) - handle = server.scene.add_mesh_trimesh(name, mesh=mesh) - # Manually set color since trimesh mesh may not carry it - if hasattr(handle, "color"): - handle.color = color - - elif geom_type == mujoco.mjtGeom.mjGEOM_ELLIPSOID: - # Use a unit sphere mesh scaled non-uniformly - mesh = trimesh.creation.icosphere(subdivisions=3, radius=1.0) - mesh.vertices *= np.array([float(size[0]), float(size[1]), float(size[2])]) - handle = server.scene.add_mesh_trimesh(name, mesh=mesh) - - elif geom_type == mujoco.mjtGeom.mjGEOM_CYLINDER: - handle = server.scene.add_cylinder( - name, - radius=float(size[0]), - height=float(size[1]) * 2, - color=color, - opacity=opacity, - ) - - elif geom_type == mujoco.mjtGeom.mjGEOM_BOX: - handle = server.scene.add_box( - name, - dimensions=( - float(size[0]) * 2, - float(size[1]) * 2, - float(size[2]) * 2, - ), - color=color, - opacity=opacity, - ) - - elif geom_type == mujoco.mjtGeom.mjGEOM_MESH: - dataid = model.geom_dataid[i] - if dataid >= 0: - vertices, faces = _extract_mesh(model, dataid) - handle = server.scene.add_mesh_simple( - name, - vertices=vertices.astype(np.float32), - faces=faces.astype(np.int32), - color=color, - opacity=opacity, - ) - - if handle is not None: - self._handles[i] = handle +import warnings - # ------------------------------------------------------------------ # - # Per-frame update # - # ------------------------------------------------------------------ # +from unilab.visualization.viser_scene import ( + VISER_AVAILABLE, + MujocoViserScene, + build_visible_env_indices, +) - def update(self, data: mujoco.MjData) -> None: - """Sync all geom transforms from *data* into the viser scene.""" - with self._server.atomic(): - for i, handle in self._handles.items(): - xpos = data.geom_xpos[i] + self._position_offset - xmat = data.geom_xmat[i] +warnings.warn( + "`unilab.utils.viser_scene` is deprecated and will be removed in 0.2.0; " + "use `unilab.visualization.viser_scene` instead.", + DeprecationWarning, + stacklevel=2, +) - handle.position = (float(xpos[0]), float(xpos[1]), float(xpos[2])) - handle.wxyz = _rotmat_to_wxyz(xmat) +__all__ = ["VISER_AVAILABLE", "MujocoViserScene", "build_visible_env_indices"] diff --git a/src/unilab/utils/xml_utils.py b/src/unilab/utils/xml_utils.py index d85697c70..00ac762ab 100644 --- a/src/unilab/utils/xml_utils.py +++ b/src/unilab/utils/xml_utils.py @@ -1,307 +1,12 @@ from __future__ import annotations -import os -import tempfile -import xml.etree.ElementTree as ET -from collections.abc import Iterator, Sequence -from pathlib import Path +import warnings +from unilab.base.backend.xml import * # noqa: F403 -def _enable_discardvisual(root: ET.Element) -> None: - compiler_tag = root.find("compiler") - if compiler_tag is None: - compiler_tag = ET.Element("compiler") - root.insert(0, compiler_tag) - compiler_tag.set("discardvisual", "true") - - -def create_discardvisual_xml(model_file: str) -> str: - tree = ET.parse(model_file) - _enable_discardvisual(tree.getroot()) - return _write_temp_xml(tree, model_file) - - -def _iter_expanded_children( - parent: ET.Element, base_dir: Path -) -> Iterator[tuple[ET.Element, Path]]: - for child in parent: - if child.tag != "include": - yield child, base_dir - continue - - include_file = child.get("file") - if not include_file: - raise ValueError(f"Invalid without file attribute in {base_dir}") - include_path = (base_dir / include_file).resolve() - include_root = ET.parse(include_path).getroot() - yield from _iter_expanded_children(include_root, include_path.parent) - - -def _iter_named_bodies(root: ET.Element, base_dir: Path) -> Iterator[str]: - for child, child_base_dir in _iter_expanded_children(root, base_dir): - if child.tag == "body": - body_name = child.get("name") - if body_name: - yield body_name - yield from _iter_named_bodies(child, child_base_dir) - - -def _get_named_bodies(model_file: str) -> tuple[list[int], list[str]]: - model_path = Path(model_file).resolve() - names = list(_iter_named_bodies(ET.parse(model_path).getroot(), model_path.parent)) - ids = list(range(1, len(names) + 1)) - return ids, names - - -def get_named_body_ids(model_file: str, names: Sequence[str]) -> list[int]: - """Resolve MuJoCo-style body ids from XML without importing mujoco.""" - body_ids, body_names = _get_named_bodies(model_file) - body_id_by_name = dict(zip(body_names, body_ids, strict=True)) - missing = [name for name in names if name not in body_id_by_name] - if missing: - missing_str = ", ".join(missing) - raise ValueError(f"Bodies not found in XML '{model_file}': {missing_str}") - return [body_id_by_name[name] for name in names] - - -def _add_w_sensors(sensor_tag: ET.Element, valid_bnames: list[str]) -> None: - for bname in valid_bnames: - ET.SubElement( - sensor_tag, "framepos", name=f"track_pos_w_{bname}", objtype="xbody", objname=bname - ) - for bname in valid_bnames: - ET.SubElement( - sensor_tag, "framequat", name=f"track_quat_w_{bname}", objtype="xbody", objname=bname - ) - for bname in valid_bnames: - ET.SubElement( - sensor_tag, - "framelinvel", - name=f"track_linvel_w_{bname}", - objtype="xbody", - objname=bname, - ) - for bname in valid_bnames: - ET.SubElement( - sensor_tag, - "frameangvel", - name=f"track_angvel_w_{bname}", - objtype="xbody", - objname=bname, - ) - - -def _add_b_sensors(sensor_tag: ET.Element, valid_bnames: list[str], baselink_name: str) -> None: - for bname in valid_bnames: - ET.SubElement( - sensor_tag, - "framepos", - name=f"track_pos_b_{bname}", - objtype="xbody", - objname=bname, - reftype="xbody", - refname=baselink_name, - ) - for bname in valid_bnames: - ET.SubElement( - sensor_tag, - "framequat", - name=f"track_quat_b_{bname}", - objtype="xbody", - objname=bname, - reftype="xbody", - refname=baselink_name, - ) - for bname in valid_bnames: - ET.SubElement( - sensor_tag, - "framelinvel", - name=f"track_linvel_b_{bname}", - objtype="xbody", - objname=bname, - reftype="xbody", - refname=baselink_name, - ) - for bname in valid_bnames: - ET.SubElement( - sensor_tag, - "frameangvel", - name=f"track_angvel_b_{bname}", - objtype="xbody", - objname=bname, - reftype="xbody", - refname=baselink_name, - ) - - -def _write_temp_xml(tree: ET.ElementTree[ET.Element], model_file: str) -> str: # type: ignore[type-arg] - fd, output_path = tempfile.mkstemp( - suffix=".xml", dir=os.path.dirname(os.path.abspath(model_file)) - ) - os.close(fd) - tree.write(output_path) - return output_path - - -def _format_values(values: list[float] | tuple[float, ...]) -> str: - return " ".join(str(float(value)) for value in values) - - -def materialize_scene_visual_override( - source_model_file: str, - *, - ground_texture_file: str | None = None, - ground_texrepeat: list[float] | tuple[float, float] | None = None, - skybox_rgb1: list[float] | tuple[float, float, float] | None = None, - skybox_rgb2: list[float] | tuple[float, float, float] | None = None, -) -> str: - """Create a temporary scene XML with visual-only overrides applied.""" - tree = ET.parse(source_model_file) - root = tree.getroot() - asset_tag = root.find("asset") - if asset_tag is None: - raise ValueError(f"Scene '{source_model_file}' is missing an tag.") - - if skybox_rgb1 is not None or skybox_rgb2 is not None: - skybox = asset_tag.find("./texture[@type='skybox']") - if skybox is None: - raise ValueError(f"Scene '{source_model_file}' is missing a skybox texture.") - if skybox_rgb1 is not None: - skybox.set("rgb1", _format_values(tuple(skybox_rgb1))) - if skybox_rgb2 is not None: - skybox.set("rgb2", _format_values(tuple(skybox_rgb2))) - - if ground_texture_file is not None: - ground_texture = asset_tag.find("./texture[@name='groundplane']") - if ground_texture is None: - raise ValueError(f"Scene '{source_model_file}' is missing the groundplane texture.") - for attr in ("builtin", "mark", "rgb1", "rgb2", "markrgb", "width", "height"): - ground_texture.attrib.pop(attr, None) - ground_texture.set("file", str(Path(ground_texture_file))) - - if ground_texrepeat is not None: - ground_material = asset_tag.find("./material[@name='groundplane']") - if ground_material is None: - raise ValueError(f"Scene '{source_model_file}' is missing the groundplane material.") - ground_material.set("texrepeat", _format_values(tuple(ground_texrepeat))) - - return _write_temp_xml(tree, source_model_file) - - -def inject_mujoco_tracking_sensors( - model_file: str, - baselink_name: str | None = None, -) -> tuple[str, list, list]: - """为 MuJoCo 后端注入 tracking sensors。 - - 注入所有 body 的世界系 (_w) sensors;若指定 baselink_name, - 同时注入相对 baselink 坐标系的 (_b) sensors。 - - Returns: - (tmp_xml_path, tracked_body_ids, valid_bnames) - """ - tracked_body_ids, valid_bnames = _get_named_bodies(model_file) - - tree = ET.parse(model_file) - root = tree.getroot() - sensor_tag = root.find("sensor") - if sensor_tag is None: - sensor_tag = ET.SubElement(root, "sensor") - - _add_w_sensors(sensor_tag, valid_bnames) - if baselink_name and baselink_name in valid_bnames: - _add_b_sensors(sensor_tag, valid_bnames, baselink_name) - - return _write_temp_xml(tree, model_file), tracked_body_ids, valid_bnames - - -def inject_motrix_tracking_sensors(model_file: str, baselink_name: str) -> tuple[str, list, list]: - """为 MotrixSim 后端注入 tracking sensors。 - - 只注入相对 baselink 坐标系的 (_b) sensors。 - 世界系 (_w) 数据由 motrixsim body API 直接提供,无需 sensor 注入。 - - Returns: - (tmp_xml_path, tracked_body_ids, valid_bnames) - """ - tracked_body_ids, valid_bnames = _get_named_bodies(model_file) - - tree = ET.parse(model_file) - root = tree.getroot() - sensor_tag = root.find("sensor") - if sensor_tag is None: - sensor_tag = ET.SubElement(root, "sensor") - - _add_b_sensors(sensor_tag, valid_bnames, baselink_name) - - return _write_temp_xml(tree, model_file), tracked_body_ids, valid_bnames - - -def processed_xml(xml_path): - xml_dir = os.path.dirname(os.path.abspath(xml_path)) - - tree = ET.parse(xml_path) - root = tree.getroot() - - compiler = root.find("compiler") - if compiler is not None: - meshdir = compiler.get("meshdir") - if meshdir: - abs_meshdir = os.path.normpath(os.path.join(xml_dir, meshdir)) - compiler.set("meshdir", abs_meshdir) - - bodys = root.findall(".//body") - - geom_names = [] - for body in bodys: - body_name = body.get("name", "unnamed_body") - geoms = body.findall("geom") - - if geoms: - filtered_geoms = [] - for geom in geoms: - geom_class = geom.get("class") - if geom_class != "visual": - filtered_geoms.append(geom) - - if filtered_geoms: - i = 0 - for geom in filtered_geoms: - geom_name = geom.get("name", "unnamed_geom") - if geom_name == "unnamed_geom": - new_name = f"{body_name}_geom{i}" - i += 1 - geom.set("name", new_name) - geom_name = new_name - geom_names.append(geom_name) - - new_xml_string = ET.tostring(root, encoding="unicode") - return new_xml_string, geom_names - - -def add_sensor(root, sensor_type, name, **kwargs): - """ - 在 MuJoCo XML 的 sensor 节点下添加传感器的通用函数。 - - 参数: - - root: XML 的根节点 - - sensor_type: 传感器标签名 (如 'gyro', 'contact', 'framepos') - - name: 传感器的 name 属性 - - **kwargs: 其他任意属性 (如 site='imu', geom1='floor' 等) - """ - # 1. 查找或创建 标签 - sensor_element = root.find("sensor") - if sensor_element is None: - sensor_element = ET.SubElement(root, "sensor") - - # 2. 创建具体的传感器子节点 - sensor = ET.SubElement(sensor_element, sensor_type) - - # 3. 设置必选的 name 属性 - sensor.set("name", name) - - # 4. 循环设置其他传入的属性 - for key, value in kwargs.items(): - sensor.set(key, str(value)) - - return sensor +warnings.warn( + "`unilab.utils.xml_utils` is deprecated and will be removed in 0.2.0; " + "use `unilab.base.backend.xml` instead.", + DeprecationWarning, + stacklevel=2, +) 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/visualization/render_many.py b/src/unilab/visualization/render_many.py new file mode 100644 index 000000000..9b22f2085 --- /dev/null +++ b/src/unilab/visualization/render_many.py @@ -0,0 +1,606 @@ +"""MuJoCo-only batched rendering helpers. + +This module renders many MuJoCo states into image frames by constructing +MuJoCo model/data/renderer objects inside worker processes. It is not available +for Motrix-only workflows. +""" + +import math +import os +import subprocess +import sys +import textwrap +from collections.abc import Sequence +from typing import Any + +import imageio + +_USER_MUJOCO_GL = os.environ.get("MUJOCO_GL") + +_EGL_PROBE_SCRIPT = textwrap.dedent( + ''' + import mujoco + + xml = """ + + + + + + """ + + model = mujoco.MjModel.from_xml_string(xml) + data = mujoco.MjData(model) + renderer = mujoco.Renderer(model, height=8, width=8) + mujoco.mj_forward(model, data) + renderer.update_scene(data) + renderer.render() + renderer.close() + ''' +) + + +def _egl_runtime_usable() -> bool: + env = os.environ.copy() + env["MUJOCO_GL"] = "egl" + env.setdefault("MUJOCO_EGL_DEVICE_ID", "0") + + try: + subprocess.run( + [sys.executable, "-c", _EGL_PROBE_SCRIPT], + env=env, + check=True, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + timeout=10, + ) + except (OSError, subprocess.SubprocessError): + return False + + os.environ.setdefault("MUJOCO_EGL_DEVICE_ID", env["MUJOCO_EGL_DEVICE_ID"]) + return True + + +def _resolve_gl_backend() -> str: + """Pick a valid MUJOCO_GL backend for the current platform. + + Respects an explicit user setting unless it's provably invalid (e.g. egl + on macOS). Falls back to glfw when EGL is requested but not available. + """ + current = os.environ.get("MUJOCO_GL", "") + safe_values = {"glfw", "osmesa", "disabled"} + + if sys.platform == "darwin": + # macOS has no EGL support; glfw is the only off-screen option + return current if current in safe_values else "glfw" + + # Linux / other: honour explicit non-egl choices supplied before import. + if current in safe_values and current == _USER_MUJOCO_GL: + return current + + # Probe EGL by creating a tiny MuJoCo renderer in a clean subprocess. + if _egl_runtime_usable(): + return "egl" + + return "glfw" + + +# Must be set *before* importing mujoco (it reads the var at import time) +os.environ["MUJOCO_GL"] = _resolve_gl_backend() + +import mujoco # noqa: E402 +import numpy as np + + +def get_grid_offsets(num_envs, spacing=1.0): + rows = int(math.ceil(math.sqrt(num_envs))) + cols = int(math.ceil(num_envs / rows)) + offsets = np.zeros((num_envs, 2)) + for i in range(num_envs): + r = i // cols + c = i % cols + offsets[i, 0] = r * spacing + offsets[i, 1] = c * spacing + return offsets + + +# Worker global context +_worker_ctx: dict[str, Any] = {} + + +def _close_worker(): + """Explicitly close the renderer in the worker context.""" + if "renderer" in _worker_ctx: + _worker_ctx["renderer"].close() + + +def init_worker(model_path, shape): + """Initialize MuJoCo-only rendering context for a worker process.""" + import atexit + + def _load_model(path_like): + path = str(path_like) + loader = ( + mujoco.MjModel.from_binary_path + if path.endswith(".mjb") + else mujoco.MjModel.from_xml_path + ) + return loader(path) + + if isinstance(model_path, Sequence) and not isinstance(model_path, (str, bytes, os.PathLike)): + models = [_load_model(path) for path in model_path] + else: + models = [_load_model(model_path)] + + for model in models: + model.vis.global_.offwidth = 3840 + model.vis.global_.offheight = 2160 + + _worker_ctx["models"] = models + _worker_ctx["data_list"] = [mujoco.MjData(model) for model in models] + _worker_ctx["renderer"] = mujoco.Renderer(models[0], height=shape[1], width=shape[0]) + atexit.register(_close_worker) + + +def render_frame_job(args): + """ + Worker function to render a single frame. + args: (state_batch, offsets, transparent, cam_distance, cam_elevation, cam_azimuth, cam_lookat) + """ + state_batch, offsets, transparent, cam_distance, cam_elevation, cam_azimuth, cam_lookat = args + + models = _worker_ctx["models"] + data_list = _worker_ctx["data_list"] + renderer = _worker_ctx["renderer"] + + # Visual options + vopt = mujoco.MjvOption() + vopt.flags[mujoco.mjtVisFlag.mjVIS_TRANSPARENT] = transparent + pert = mujoco.MjvPerturb() + catmask_dynamic = mujoco.mjtCatBit.mjCAT_DYNAMIC + catmask_static = mujoco.mjtCatBit.mjCAT_STATIC + + # Helper to set state + def set_state(model, d, s, offset=None): + d.time = s[0] + d.qpos[:] = s[1 : 1 + model.nq] + d.qvel[:] = s[1 + model.nq : 1 + model.nq + model.nv] + + apply_root_offset = False + + if offset is not None: + # Check if Root (Body 1) has a free joint or slide joints allowing X/Y movement + # Body 0 is world. Body 1 is usually the robot base. + robot_moved = False + + # Heuristic: Check joint at qpos 0, 1. + # If jnt_type[0] is free (0), fine. + # If jnt_type[0] is slide (2) and axis is x/y... + + # Better check: Does the first body have a joint? + first_body_jnt = model.body_jntadr[1] if model.nbody > 1 else -1 + if first_body_jnt >= 0: + jnt_type = model.jnt_type[first_body_jnt] + # mjJNT_FREE=0 + if jnt_type == 0: + d.qpos[0] += offset[0] + d.qpos[1] += offset[1] + robot_moved = True + + # If robot wasn't moved via qpos, we need to manually offset geometries later + if not robot_moved: + apply_root_offset = True + + # 2. Box offset + box_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "box") + if box_id >= 0: + jnt_adr = model.body_jntadr[box_id] + if jnt_adr >= 0: + qpos_adr = model.jnt_qposadr[jnt_adr] + d.qpos[qpos_adr] += offset[0] + d.qpos[qpos_adr + 1] += offset[1] + + # 3. Target offset (target_x, target_y) + target_x = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_JOINT, "target_x") + if target_x >= 0: + d.qpos[model.jnt_qposadr[target_x]] += offset[0] + + target_y = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_JOINT, "target_y") + if target_y >= 0: + d.qpos[model.jnt_qposadr[target_y]] += offset[1] + + mujoco.mj_forward(model, d) + + # Post-process: Shift all geometries if robot root wasn't moved + if apply_root_offset and offset is not None: + # Shift all geoms? + # We should shift Everything that is PART OF THE ROBOT. + # Or just everything? + # Box and Target were already shifted via qpos. + # BUT qpos shift updates body_pos which updates geom_pos. + # If we shift ALL geom_pos, we double shift Box and Target! + + # So we need to shift geoms that belong to bodies which are NOT Box or Target. + # Or simpler: Shift everything, but subtract offset from Box/Target qpos first? No. + + # Let's iterate bodies. + # Simple heuristic: Shift everything except Box and Target? + box_body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "box") + target_body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "mocap_target") + + # Also target might be just a body named "mocap_target" + + for i in range(model.ngeom): + body_id = model.geom_bodyid[i] + # If it is robot body. + # We want to shift generally everything that wasn't shifted by Qpos. + # Box and Target were shifted by Qpos. + # Floor (Plane) should usually NOT be shifted (infinite). + # Everything else (Robot Base, Robot Links, Decoration) should be shifted. + + is_box_or_target = (body_id == box_body_id) or (body_id == target_body_id) + is_plane = model.geom_type[i] == mujoco.mjtGeom.mjGEOM_PLANE + + if not is_box_or_target and not is_plane: + d.geom_xpos[i, 0] += offset[0] + d.geom_xpos[i, 1] += offset[1] + + # Also update site positions if they are visualized + for i in range(model.nsite): + body_id = model.site_bodyid[i] + + is_box_or_target = (body_id == box_body_id) or (body_id == target_body_id) + + if not is_box_or_target: + d.site_xpos[i, 0] += offset[0] + d.site_xpos[i, 1] += offset[1] + + num_envs = state_batch.shape[0] + + # 1. Clear/Init Scene + primary_model = models[0] + primary_data = data_list[0] + set_state( + primary_model, primary_data, state_batch[0], offsets[0] if offsets is not None else None + ) + + # Init Camera + cam = mujoco.MjvCamera() + if offsets is not None: + center_x = np.mean(offsets[:, 0]) + center_y = np.mean(offsets[:, 1]) + if cam_lookat is None: + cam.lookat = [center_x, center_y, 0.75] + else: + cam.lookat = [float(cam_lookat[0]), float(cam_lookat[1]), float(cam_lookat[2])] + cam.distance = cam_distance + cam.elevation = cam_elevation + cam.azimuth = cam_azimuth + cam.type = mujoco.mjtCamera.mjCAMERA_FREE + else: + cam.type = mujoco.mjtCamera.mjCAMERA_FREE + + renderer.update_scene(primary_data, camera=cam, scene_option=vopt) + + # 2. Add other robots + for i in range(1, num_envs): + model = models[min(i, len(models) - 1)] + data = data_list[min(i, len(data_list) - 1)] + set_state(model, data, state_batch[i], offsets[i] if offsets is not None else None) + mujoco.mjv_addGeoms(model, data, vopt, pert, catmask_dynamic, renderer.scene) + + # Avoid duplicating world floor/static group-0 geoms across envs, while still + # adding fixed-base static visuals (e.g., Sharpa base link visual). + geomgroup0 = int(vopt.geomgroup[0]) + vopt.geomgroup[0] = 0 + mujoco.mjv_addGeoms(model, data, vopt, pert, catmask_static, renderer.scene) + vopt.geomgroup[0] = geomgroup0 + + return renderer.render() + + +def render_states_get_frames( + state_list, + model_path, + width=1280, + height=720, + num_processes=8, + camera_id=-1, + cam_distance=2.0, + cam_elevation=-20, + cam_azimuth=90, + cam_lookat=None, + render_spacing=1.0, +): + """ + Render a list of physics states and return the list of frames. + + Args: + state_list: List of numpy arrays, each shape (num_envs, state_dim). + model_path: Path to the mujoco XML model file. + width: Width of the video. + height: Height of the video. + num_processes: Number of parallel processes to use. + camera_id: Camera ID to render from. + cam_distance: Camera distance from lookat point. + cam_elevation: Camera elevation angle in degrees. + cam_azimuth: Camera azimuth angle in degrees. + cam_lookat: Optional [x, y, z] lookat override for the free camera. + render_spacing: Grid spacing used to offset each env in composed video frames. + Returns: + List of numpy arrays (H, W, 3) (RGB) + """ + if not state_list: + print("No states to render.") + return [] + + num_envs = state_list[0].shape[0] + offsets = get_grid_offsets(num_envs, spacing=render_spacing) + shape = (width, height) + + print( + f"Rendering {len(state_list)} frames for {num_envs} envs with {num_processes} processes..." + ) + + # Prepare arguments for each frame + tasks = [ + (s, offsets, False, cam_distance, cam_elevation, cam_azimuth, cam_lookat) + for s in state_list + ] + + frames = [] + + if num_processes <= 1: + # Serial execution + # Initialize context manually + init_worker(model_path, shape) + try: + for task in tasks: + res = render_frame_job(task) + frames.append(res) + finally: + _close_worker() + else: + # Use multiprocessing Pool + # On macOS, use spawn to avoid forking OpenGL/MuJoCo contexts. + import multiprocessing + + ctx = multiprocessing.get_context("spawn") + with ctx.Pool( + processes=num_processes, initializer=init_worker, initargs=(model_path, shape) + ) as pool: + results = pool.map(render_frame_job, tasks) + frames.extend(results) + + return frames + + +def _get_nearest_env_indices(offsets, primary_idx, max_extra): + """Return indices of the *max_extra* environments closest to *primary_idx*.""" + if len(offsets) <= 1 + max_extra: + return [i for i in range(len(offsets)) if i != primary_idx] + primary = offsets[primary_idx] + dists = np.linalg.norm(offsets - primary, axis=1) + dists[primary_idx] = np.inf # exclude self + return list(np.argsort(dists)[:max_extra]) + + +def render_frame_tracking_job(args): + """Render a single frame with camera tracking on the primary env's root body. + + The camera uses ``mjCAMERA_TRACKING`` so it follows the robot each frame. + Only the primary env + nearest neighbours are rendered. + """ + ( + state_batch, + offsets, + env_indices, + primary_local_idx, + cam_distance, + cam_elevation, + cam_azimuth, + ) = args + + models = _worker_ctx["models"] + data_list = _worker_ctx["data_list"] + renderer = _worker_ctx["renderer"] + + vopt = mujoco.MjvOption() + pert = mujoco.MjvPerturb() + catmask_dynamic = mujoco.mjtCatBit.mjCAT_DYNAMIC + catmask_static = mujoco.mjtCatBit.mjCAT_STATIC + + def set_state(model, d, s, offset=None): + d.time = s[0] + d.qpos[:] = s[1 : 1 + model.nq] + d.qvel[:] = s[1 + model.nq : 1 + model.nq + model.nv] + + apply_root_offset = False + + if offset is not None: + robot_moved = False + first_body_jnt = model.body_jntadr[1] if model.nbody > 1 else -1 + if first_body_jnt >= 0: + jnt_type = model.jnt_type[first_body_jnt] + if jnt_type == 0: # mjJNT_FREE + d.qpos[0] += offset[0] + d.qpos[1] += offset[1] + robot_moved = True + + if not robot_moved: + apply_root_offset = True + + box_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "box") + if box_id >= 0: + jnt_adr = model.body_jntadr[box_id] + if jnt_adr >= 0: + qpos_adr = model.jnt_qposadr[jnt_adr] + d.qpos[qpos_adr] += offset[0] + d.qpos[qpos_adr + 1] += offset[1] + + target_x = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_JOINT, "target_x") + if target_x >= 0: + d.qpos[model.jnt_qposadr[target_x]] += offset[0] + + target_y = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_JOINT, "target_y") + if target_y >= 0: + d.qpos[model.jnt_qposadr[target_y]] += offset[1] + + mujoco.mj_forward(model, d) + + if apply_root_offset and offset is not None: + box_body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "box") + target_body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "mocap_target") + + for i in range(model.ngeom): + body_id = model.geom_bodyid[i] + is_box_or_target = (body_id == box_body_id) or (body_id == target_body_id) + is_plane = model.geom_type[i] == mujoco.mjtGeom.mjGEOM_PLANE + + if not is_box_or_target and not is_plane: + d.geom_xpos[i, 0] += offset[0] + d.geom_xpos[i, 1] += offset[1] + + for i in range(model.nsite): + body_id = model.site_bodyid[i] + is_box_or_target = (body_id == box_body_id) or (body_id == target_body_id) + if not is_box_or_target: + d.site_xpos[i, 0] += offset[0] + d.site_xpos[i, 1] += offset[1] + + # Primary env first — camera tracks body 1 of this env + primary_global = env_indices[primary_local_idx] + primary_model = models[min(primary_global, len(models) - 1)] + primary_data = data_list[min(primary_global, len(data_list) - 1)] + set_state( + primary_model, + primary_data, + state_batch[primary_global], + offsets[primary_global] if offsets is not None else None, + ) + + cam = mujoco.MjvCamera() + cam.type = mujoco.mjtCamera.mjCAMERA_TRACKING + cam.trackbodyid = 1 # robot root body + cam.distance = cam_distance + cam.elevation = cam_elevation + cam.azimuth = cam_azimuth + + renderer.update_scene(primary_data, camera=cam, scene_option=vopt) + + # Add neighbour envs as background context + for local_i, global_i in enumerate(env_indices): + if local_i == primary_local_idx: + continue + model = models[min(global_i, len(models) - 1)] + data = data_list[min(global_i, len(data_list) - 1)] + set_state( + model, data, state_batch[global_i], offsets[global_i] if offsets is not None else None + ) + mujoco.mjv_addGeoms(model, data, vopt, pert, catmask_dynamic, renderer.scene) + + geomgroup0 = int(vopt.geomgroup[0]) + vopt.geomgroup[0] = 0 + mujoco.mjv_addGeoms(model, data, vopt, pert, catmask_static, renderer.scene) + vopt.geomgroup[0] = geomgroup0 + + return renderer.render() + + +def render_states_get_frames_tracking( + state_list, + model_path, + width=1280, + height=720, + tracking_env_idx=0, + max_extra_envs=2, + cam_distance=2.0, + cam_elevation=-20, + cam_azimuth=90, + render_spacing=1.0, +): + """Render with camera tracking on a single primary environment. + + Only the primary env and its nearest neighbours are shown. The camera + follows the root body of the primary env each frame (``mjCAMERA_TRACKING``). + + Args: + state_list: List of numpy arrays, each shape (num_envs, state_dim). + model_path: Path to the mujoco XML model file. + tracking_env_idx: Index of the primary environment to track. + max_extra_envs: Number of nearest-neighbour envs to render alongside. + cam_distance: Camera distance from the tracked body. + cam_elevation: Camera elevation angle in degrees. + cam_azimuth: Camera azimuth angle in degrees. + render_spacing: Grid spacing for env layout. + """ + if not state_list: + print("No states to render.") + return [] + + num_envs = state_list[0].shape[0] + offsets = get_grid_offsets(num_envs, spacing=render_spacing) + shape = (width, height) + + tracking_env_idx = min(tracking_env_idx, num_envs - 1) + neighbour_indices = _get_nearest_env_indices(offsets, tracking_env_idx, max_extra_envs) + env_indices = [tracking_env_idx] + neighbour_indices + primary_local_idx = 0 # primary is always first in env_indices + + total_shown = len(env_indices) + print( + f"Rendering {len(state_list)} frames (tracking env {tracking_env_idx} " + f"+ {total_shown - 1} neighbours) ..." + ) + + tasks = [ + (s, offsets, env_indices, primary_local_idx, cam_distance, cam_elevation, cam_azimuth) + for s in state_list + ] + + # Camera tracking changes each frame so multiprocessing gives inconsistent + # results when workers don't share state. Default to serial. + frames = [] + init_worker(model_path, shape) + try: + for task in tasks: + frames.append(render_frame_tracking_job(task)) + finally: + _close_worker() + + return frames + + +def render_states_to_video( + state_list, + model_path, + output_path, + fps=30, + width=1280, + height=720, + num_processes=8, + cam_distance=2.0, + cam_elevation=-20, + cam_azimuth=90, + cam_lookat=None, + render_spacing=1.0, +): + """ + Render a list of physics states to a video file using parallel processing. + """ + frames = render_states_get_frames( + state_list, + model_path, + width, + height, + num_processes, + cam_distance=cam_distance, + cam_elevation=cam_elevation, + cam_azimuth=cam_azimuth, + cam_lookat=cam_lookat, + render_spacing=render_spacing, + ) + + print(f"Saving video to {output_path}...") + imageio.mimsave(output_path, frames, fps=fps) + print("Done!") diff --git a/src/unilab/visualization/viser_scene.py b/src/unilab/visualization/viser_scene.py new file mode 100644 index 000000000..fdb64da61 --- /dev/null +++ b/src/unilab/visualization/viser_scene.py @@ -0,0 +1,281 @@ +# pyright: reportMissingImports=false +"""MuJoCo-to-viser scene adapter for interactive web-based 3D visualization. + +This module renders MuJoCo scenes via a viser web server, providing browser-based +interactive 3D viewing without requiring a local display or GLFW. It is gated +behind the ``viser`` optional-dependency group and is **not** imported by default. + +Usage (from ``scripts/play_viser.py``):: + + from unilab.visualization.viser_scene import MujocoViserScene, VISER_AVAILABLE +""" + +from __future__ import annotations + +import math +from typing import Any + +import mujoco +import numpy as np + +try: + import trimesh + import viser + + VISER_AVAILABLE = True +except ImportError: + VISER_AVAILABLE = False + + +# --------------------------------------------------------------------------- # +# Rotation helpers (pure numpy, no scipy dependency) # +# --------------------------------------------------------------------------- # + + +def _rotmat_to_wxyz(mat: np.ndarray) -> tuple[float, float, float, float]: + """Convert a 3x3 rotation matrix to a (w, x, y, z) quaternion.""" + m = np.asarray(mat, dtype=np.float64).reshape(3, 3) + trace = m[0, 0] + m[1, 1] + m[2, 2] + + if trace > 0: + s = 0.5 / math.sqrt(trace + 1.0) + w = 0.25 / s + x = (m[2, 1] - m[1, 2]) * s + y = (m[0, 2] - m[2, 0]) * s + z = (m[1, 0] - m[0, 1]) * s + elif m[0, 0] > m[1, 1] and m[0, 0] > m[2, 2]: + s = 2.0 * math.sqrt(1.0 + m[0, 0] - m[1, 1] - m[2, 2]) + w = (m[2, 1] - m[1, 2]) / s + x = 0.25 * s + y = (m[0, 1] + m[1, 0]) / s + z = (m[0, 2] + m[2, 0]) / s + elif m[1, 1] > m[2, 2]: + s = 2.0 * math.sqrt(1.0 + m[1, 1] - m[0, 0] - m[2, 2]) + w = (m[0, 2] - m[2, 0]) / s + x = (m[0, 1] + m[1, 0]) / s + y = 0.25 * s + z = (m[1, 2] + m[2, 1]) / s + else: + s = 2.0 * math.sqrt(1.0 + m[2, 2] - m[0, 0] - m[1, 1]) + w = (m[1, 0] - m[0, 1]) / s + x = (m[0, 2] + m[2, 0]) / s + y = (m[1, 2] + m[2, 1]) / s + z = 0.25 * s + + return (float(w), float(x), float(y), float(z)) + + +# --------------------------------------------------------------------------- # +# Geometry extraction helpers # +# --------------------------------------------------------------------------- # + + +def _rgba_to_color(rgba: np.ndarray) -> tuple[int, int, int]: + """Convert MuJoCo float RGBA [0,1] to viser int RGB [0,255].""" + return ( + int(np.clip(rgba[0] * 255, 0, 255)), + int(np.clip(rgba[1] * 255, 0, 255)), + int(np.clip(rgba[2] * 255, 0, 255)), + ) + + +def _rgba_to_opacity(rgba: np.ndarray) -> float: + return float(np.clip(rgba[3], 0.0, 1.0)) + + +def _extract_mesh(model: mujoco.MjModel, geom_dataid: int) -> tuple[np.ndarray, np.ndarray]: + """Extract vertices and faces for a MuJoCo mesh geom.""" + vert_adr = model.mesh_vertadr[geom_dataid] + vert_num = model.mesh_vertnum[geom_dataid] + face_adr = model.mesh_faceadr[geom_dataid] + face_num = model.mesh_facenum[geom_dataid] + + vertices = model.mesh_vert[vert_adr : vert_adr + vert_num].copy() + faces = model.mesh_face[face_adr : face_adr + face_num].copy() + return vertices, faces + + +def build_visible_env_indices(num_envs: int, visible_envs: int) -> np.ndarray: + """Select a stable subset of env indices spread across the full batch. + + Args: + num_envs: Total number of runtime environments. + visible_envs: Number of env slots exposed in the viewer. + + Returns: + A monotonically increasing array of runtime env indices. + """ + if visible_envs <= 0: + raise ValueError(f"visible_envs must be positive, got {visible_envs}") + if visible_envs >= num_envs: + return np.arange(num_envs, dtype=np.int32) + return np.floor(np.linspace(0, num_envs, visible_envs, endpoint=False)).astype(np.int32) + + +# --------------------------------------------------------------------------- # +# MujocoViserScene # +# --------------------------------------------------------------------------- # + + +class MujocoViserScene: + """Bridges a ``mujoco.MjModel`` to a ``viser.ViserServer`` scene graph. + + Call :meth:`build` once to populate the scene with geometry handles, then + call :meth:`update` each frame to sync body transforms from ``MjData``. + """ + + def __init__( + self, + server: Any, + model: mujoco.MjModel, + *, + name_prefix: str = "/mujoco", + position_offset: tuple[float, float, float] = (0.0, 0.0, 0.0), + render_plane: bool = True, + ) -> None: + if not VISER_AVAILABLE: + raise ImportError("viser is not installed. Install with: uv sync --extra viser") + self._server: viser.ViserServer = server + self._model = model + self._name_prefix = name_prefix.rstrip("/") or "/mujoco" + self._position_offset = np.asarray(position_offset, dtype=np.float64) + self._render_plane = bool(render_plane) + self._handles: dict[int, Any] = {} + self._build() + + def reset( + self, + model: mujoco.MjModel, + *, + position_offset: tuple[float, float, float] | None = None, + render_plane: bool | None = None, + ) -> None: + """Rebuild the viser scene for a new MuJoCo model. + + Args: + model: MuJoCo model whose geoms should populate the scene. + position_offset: Optional XYZ offset applied to all geoms. + render_plane: Optional override for whether plane geoms should be built. + + Returns: + None. + """ + self.close() + self._model = model + if position_offset is not None: + self._position_offset = np.asarray(position_offset, dtype=np.float64) + if render_plane is not None: + self._render_plane = bool(render_plane) + self._build() + + def close(self) -> None: + """Remove all scene handles owned by this adapter.""" + for handle in self._handles.values(): + handle.remove() + self._handles.clear() + + # ------------------------------------------------------------------ # + # Scene construction # + # ------------------------------------------------------------------ # + + def _build(self) -> None: + """Create viser scene nodes for every MuJoCo geom.""" + model = self._model + server = self._server + + server.scene.set_up_direction("+z") + + for i in range(model.ngeom): + geom_type = model.geom_type[i] + size = model.geom_size[i] + rgba = model.geom_rgba[i] + color = _rgba_to_color(rgba) + opacity = _rgba_to_opacity(rgba) + name = f"{self._name_prefix}/geom/{i}" + + handle: Any | None = None + + if geom_type == mujoco.mjtGeom.mjGEOM_PLANE: + if not self._render_plane: + continue + # Render ground plane as a grid + plane_size = float(size[0]) if size[0] > 0 else 10.0 + handle = server.scene.add_grid( + name, + width=plane_size * 2, + height=plane_size * 2, + cell_size=0.5, + ) + + elif geom_type == mujoco.mjtGeom.mjGEOM_SPHERE: + handle = server.scene.add_icosphere( + name, + radius=float(size[0]), + color=color, + opacity=opacity, + ) + + elif geom_type == mujoco.mjtGeom.mjGEOM_CAPSULE: + half_len = float(size[1]) + radius = float(size[0]) + mesh = trimesh.creation.capsule(height=half_len * 2, radius=radius) + handle = server.scene.add_mesh_trimesh(name, mesh=mesh) + # Manually set color since trimesh mesh may not carry it + if hasattr(handle, "color"): + handle.color = color + + elif geom_type == mujoco.mjtGeom.mjGEOM_ELLIPSOID: + # Use a unit sphere mesh scaled non-uniformly + mesh = trimesh.creation.icosphere(subdivisions=3, radius=1.0) + mesh.vertices *= np.array([float(size[0]), float(size[1]), float(size[2])]) + handle = server.scene.add_mesh_trimesh(name, mesh=mesh) + + elif geom_type == mujoco.mjtGeom.mjGEOM_CYLINDER: + handle = server.scene.add_cylinder( + name, + radius=float(size[0]), + height=float(size[1]) * 2, + color=color, + opacity=opacity, + ) + + elif geom_type == mujoco.mjtGeom.mjGEOM_BOX: + handle = server.scene.add_box( + name, + dimensions=( + float(size[0]) * 2, + float(size[1]) * 2, + float(size[2]) * 2, + ), + color=color, + opacity=opacity, + ) + + elif geom_type == mujoco.mjtGeom.mjGEOM_MESH: + dataid = model.geom_dataid[i] + if dataid >= 0: + vertices, faces = _extract_mesh(model, dataid) + handle = server.scene.add_mesh_simple( + name, + vertices=vertices.astype(np.float32), + faces=faces.astype(np.int32), + color=color, + opacity=opacity, + ) + + if handle is not None: + self._handles[i] = handle + + # ------------------------------------------------------------------ # + # Per-frame update # + # ------------------------------------------------------------------ # + + def update(self, data: mujoco.MjData) -> None: + """Sync all geom transforms from *data* into the viser scene.""" + with self._server.atomic(): + for i, handle in self._handles.items(): + xpos = data.geom_xpos[i] + self._position_offset + xmat = data.geom_xmat[i] + + handle.position = (float(xpos[0]), float(xpos[1]), float(xpos[2])) + handle.wxyz = _rotmat_to_wxyz(xmat) 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..b935c23e2 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.base.observations import flatten_obs_dict + from unilab.base.registry import ensure_registries 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 ensure_registries() diff --git a/tests/algos/test_rsl_rl_runner.py b/tests/algos/test_rsl_rl_runner.py index 30d966d19..1d4a34210 100644 --- a/tests/algos/test_rsl_rl_runner.py +++ b/tests/algos/test_rsl_rl_runner.py @@ -19,11 +19,11 @@ import torch from tensordict import TensorDict +from unilab.algos.torch.rsl_rl.compat import convert_config_v5, is_rsl_rl_v5 from unilab.base import registry +from unilab.base.registry import ensure_registries 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.utils.tensor import to_torch ensure_registries() 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_reward_injection.py b/tests/config/test_reward_injection.py index 0135528c6..a319c60a2 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.config.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..961222759 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.config.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.algos.torch.rsl_rl.vec_env_wrapper 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.algos.torch.rsl_rl.vec_env_wrapper 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.algos.torch.rsl_rl.vec_env_wrapper 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.algos.torch.rsl_rl.vec_env_wrapper 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.algos.torch.rsl_rl.vec_env_wrapper import RslRlVecEnvWrapper class FakeEnv: def __init__(self): diff --git a/tests/training/test_training_helpers.py b/tests/training/test_training_helpers.py index d6bdc8e0f..6400947d1 100644 --- a/tests/training/test_training_helpers.py +++ b/tests/training/test_training_helpers.py @@ -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..4e4a0ee38 100644 --- a/tests/utils/test_experiment_tracking.py +++ b/tests/utils/test_experiment_tracking.py @@ -4,9 +4,9 @@ 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.training.logging.experiment import ExperimentTracker, build_wandb_settings +from unilab.training.logging.offpolicy import OffPolicyLogger +from unilab.training.logging.onpolicy import OnPolicyLogger 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..568fe1e73 --- /dev/null +++ b/tests/utils/test_utils_package_policy.py @@ -0,0 +1,38 @@ +import importlib +import sys +import warnings +from pathlib import Path + +import unilab.utils + +ALLOWED_UTILS_API = {"get_default_device", "to_numpy", "to_torch"} + + +def test_utils_api_is_whitelisted() -> None: + assert set(unilab.utils.__all__) == ALLOWED_UTILS_API + + +def test_repo_has_no_package_level_utils_imports() -> None: + current_file = Path(__file__).resolve() + for root in (Path("src"), Path("tests"), Path("scripts")): + 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_each_utils_shim_is_importable_and_warns_with_removal_target() -> None: + shim_modules = sorted( + f"unilab.utils.{path.stem}" + for path in Path("src/unilab/utils").glob("*.py") + if path.stem not in {"__init__", "device", "tensor"} + ) + + for module_name in shim_modules: + sys.modules.pop(module_name, None) + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always", DeprecationWarning) + module = importlib.import_module(module_name) + assert module is not None + assert any(item.category is DeprecationWarning for item in caught), module_name + assert any("0.2.0" in str(item.message) for item in caught), module_name 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: From 8df738b770cf1e3f7d11cc2b6b258b7e1ac796dd Mon Sep 17 00:00:00 2001 From: tatp-yf Date: Fri, 24 Apr 2026 01:16:51 +0800 Subject: [PATCH 2/9] refactor: remove utils shims and owner aliases --- .../benchmark_mujoco_backend_step_detail.py | 2 +- src/unilab/algos/torch/common/__init__.py | 5 -- src/unilab/algos/torch/common/device.py | 3 +- src/unilab/algos/torch/common/tensor.py | 3 - src/unilab/algos/torch/offpolicy/logging.py | 14 ---- src/unilab/utils/algo_utils.py | 15 ----- src/unilab/utils/device_utils.py | 16 ----- src/unilab/utils/experiment_tracking.py | 28 -------- src/unilab/utils/final_observation.py | 26 -------- src/unilab/utils/hardware_monitor.py | 14 ---- src/unilab/utils/logging_common.py | 14 ---- src/unilab/utils/math_utils.py | 52 --------------- src/unilab/utils/obs_utils.py | 26 -------- src/unilab/utils/offpolicy_logger.py | 14 ---- src/unilab/utils/onpolicy_logger.py | 14 ---- src/unilab/utils/render_many.py | 12 ---- src/unilab/utils/reward_utils.py | 14 ---- src/unilab/utils/rsl_rl_compat.py | 12 ---- src/unilab/utils/rsl_rl_vec_env_wrapper.py | 14 ---- src/unilab/utils/run_utils.py | 30 --------- src/unilab/utils/torch_utils.py | 14 ---- src/unilab/utils/viser_scene.py | 18 ------ src/unilab/utils/xml_utils.py | 12 ---- tests/utils/test_utils_package_policy.py | 64 +++++++++++++------ 24 files changed, 48 insertions(+), 388 deletions(-) delete mode 100644 src/unilab/algos/torch/common/tensor.py delete mode 100644 src/unilab/algos/torch/offpolicy/logging.py delete mode 100644 src/unilab/utils/algo_utils.py delete mode 100644 src/unilab/utils/device_utils.py delete mode 100644 src/unilab/utils/experiment_tracking.py delete mode 100644 src/unilab/utils/final_observation.py delete mode 100644 src/unilab/utils/hardware_monitor.py delete mode 100644 src/unilab/utils/logging_common.py delete mode 100644 src/unilab/utils/math_utils.py delete mode 100644 src/unilab/utils/obs_utils.py delete mode 100644 src/unilab/utils/offpolicy_logger.py delete mode 100644 src/unilab/utils/onpolicy_logger.py delete mode 100644 src/unilab/utils/render_many.py delete mode 100644 src/unilab/utils/reward_utils.py delete mode 100644 src/unilab/utils/rsl_rl_compat.py delete mode 100644 src/unilab/utils/rsl_rl_vec_env_wrapper.py delete mode 100644 src/unilab/utils/run_utils.py delete mode 100644 src/unilab/utils/torch_utils.py delete mode 100644 src/unilab/utils/viser_scene.py delete mode 100644 src/unilab/utils/xml_utils.py diff --git a/benchmark/benchmark_mujoco_backend_step_detail.py b/benchmark/benchmark_mujoco_backend_step_detail.py index 676274487..da5485553 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.backend.xml import create_discardvisual_xml from unilab.base.dtype_config import get_global_dtype -from unilab.utils.xml_utils import create_discardvisual_xml matplotlib.use("Agg") import matplotlib.pyplot as plt diff --git a/src/unilab/algos/torch/common/__init__.py b/src/unilab/algos/torch/common/__init__.py index c4ca017fa..75b190f0b 100644 --- a/src/unilab/algos/torch/common/__init__.py +++ b/src/unilab/algos/torch/common/__init__.py @@ -4,17 +4,12 @@ from unilab.algos.torch.common.normalization import EmpiricalNormalization from unilab.algos.torch.common.stability import check_nan_loss, clip_gradients, safe_tensor from unilab.base.registry import ensure_registries -from unilab.utils.device import get_default_device -from unilab.utils.tensor import to_numpy, to_torch __all__ = [ "EmpiricalNormalization", "DistributionalQNetwork", "Critic", - "get_default_device", "get_env_dims", - "to_numpy", - "to_torch", "check_nan_loss", "clip_gradients", "safe_tensor", diff --git a/src/unilab/algos/torch/common/device.py b/src/unilab/algos/torch/common/device.py index 99f61379f..1b1703ba5 100644 --- a/src/unilab/algos/torch/common/device.py +++ b/src/unilab/algos/torch/common/device.py @@ -1,5 +1,4 @@ from unilab.base import registry -from unilab.utils.device import get_default_device def get_env_dims( @@ -19,4 +18,4 @@ def get_env_dims( return obs_dim, action_dim, critic_dim -__all__ = ["get_default_device", "get_env_dims"] +__all__ = ["get_env_dims"] diff --git a/src/unilab/algos/torch/common/tensor.py b/src/unilab/algos/torch/common/tensor.py deleted file mode 100644 index ce2f5e258..000000000 --- a/src/unilab/algos/torch/common/tensor.py +++ /dev/null @@ -1,3 +0,0 @@ -from unilab.utils.tensor import to_numpy, to_torch - -__all__ = ["to_numpy", "to_torch"] diff --git a/src/unilab/algos/torch/offpolicy/logging.py b/src/unilab/algos/torch/offpolicy/logging.py deleted file mode 100644 index 8ccd0e92f..000000000 --- a/src/unilab/algos/torch/offpolicy/logging.py +++ /dev/null @@ -1,14 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.training.logging.offpolicy import OffPolicyLogger - -warnings.warn( - "`unilab.algos.torch.offpolicy.logging` is deprecated and will be removed in 0.2.0; " - "use `unilab.training.logging.offpolicy` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = ["OffPolicyLogger"] diff --git a/src/unilab/utils/algo_utils.py b/src/unilab/utils/algo_utils.py deleted file mode 100644 index 8dd77421e..000000000 --- a/src/unilab/utils/algo_utils.py +++ /dev/null @@ -1,15 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.algos.torch.common.actor_factory import build_actor -from unilab.base.registry import ensure_registries - -warnings.warn( - "`unilab.utils.algo_utils` is deprecated and will be removed in 0.2.0; " - "use `unilab.base.registry` and `unilab.algos.torch.common.actor_factory` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = ["build_actor", "ensure_registries"] diff --git a/src/unilab/utils/device_utils.py b/src/unilab/utils/device_utils.py deleted file mode 100644 index 08bca5160..000000000 --- a/src/unilab/utils/device_utils.py +++ /dev/null @@ -1,16 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.algos.torch.common.device import get_env_dims -from unilab.utils.device import get_default_device - -warnings.warn( - "`unilab.utils.device_utils` is deprecated and will be removed in 0.2.0; " - "use `unilab.utils.device` for `get_default_device` " - "and `unilab.algos.torch.common.device` for `get_env_dims` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = ["get_default_device", "get_env_dims"] diff --git a/src/unilab/utils/experiment_tracking.py b/src/unilab/utils/experiment_tracking.py deleted file mode 100644 index 72a18530a..000000000 --- a/src/unilab/utils/experiment_tracking.py +++ /dev/null @@ -1,28 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.training.logging.experiment import ( - ExperimentTracker, - build_wandb_run_name, - build_wandb_settings, - get_device_info_dict, - get_git_info, - patch_rsl_rl_wandb_writer, -) - -warnings.warn( - "`unilab.utils.experiment_tracking` is deprecated and will be removed in 0.2.0; " - "use `unilab.training.logging.experiment` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = [ - "ExperimentTracker", - "build_wandb_run_name", - "build_wandb_settings", - "get_device_info_dict", - "get_git_info", - "patch_rsl_rl_wandb_writer", -] diff --git a/src/unilab/utils/final_observation.py b/src/unilab/utils/final_observation.py deleted file mode 100644 index e6be7472f..000000000 --- a/src/unilab/utils/final_observation.py +++ /dev/null @@ -1,26 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.base.final_observation import ( - TerminalObservationContract, - TransitionBootstrapContract, - patch_transition_next_obs, - resolve_terminal_observation_contract, - resolve_transition_bootstrap_contract, -) - -warnings.warn( - "`unilab.utils.final_observation` is deprecated and will be removed in 0.2.0; " - "use `unilab.base.final_observation` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = [ - "TerminalObservationContract", - "TransitionBootstrapContract", - "patch_transition_next_obs", - "resolve_terminal_observation_contract", - "resolve_transition_bootstrap_contract", -] diff --git a/src/unilab/utils/hardware_monitor.py b/src/unilab/utils/hardware_monitor.py deleted file mode 100644 index 40b53fe89..000000000 --- a/src/unilab/utils/hardware_monitor.py +++ /dev/null @@ -1,14 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.training.monitoring import HardwareMonitor - -warnings.warn( - "`unilab.utils.hardware_monitor` is deprecated and will be removed in 0.2.0; " - "use `unilab.training.monitoring` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = ["HardwareMonitor"] diff --git a/src/unilab/utils/logging_common.py b/src/unilab/utils/logging_common.py deleted file mode 100644 index 288594d24..000000000 --- a/src/unilab/utils/logging_common.py +++ /dev/null @@ -1,14 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.training.logging.common import BaseTrainingLogger, _fmt_number, _fmt_time, _load_wandb - -warnings.warn( - "`unilab.utils.logging_common` is deprecated and will be removed in 0.2.0; " - "use `unilab.training.logging.common` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = ["BaseTrainingLogger", "_fmt_number", "_fmt_time", "_load_wandb"] diff --git a/src/unilab/utils/math_utils.py b/src/unilab/utils/math_utils.py deleted file mode 100644 index a3d40cd73..000000000 --- a/src/unilab/utils/math_utils.py +++ /dev/null @@ -1,52 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.algos.mlx.common.rotation import axis_angle_to_quat, quat_mul -from unilab.envs.common.math import np_sample_uniform -from unilab.envs.common.rotation import ( - np_matrix_from_quat, - np_quat_angular_velocity, - np_quat_apply, - np_quat_apply_inverse, - np_quat_canonicalize, - np_quat_conjugate, - np_quat_ensure_continuity, - np_quat_error_magnitude, - np_quat_from_euler_xyz, - np_quat_inv, - np_quat_mul, - np_quat_to_axis_angle, - np_subtract_frame_transforms, - np_yaw_quat, - np_yaw_to_quat, -) - -warnings.warn( - "`unilab.utils.math_utils` is deprecated and will be removed in 0.2.0; " - "use `unilab.envs.common.rotation`, `unilab.envs.common.math`, or " - "`unilab.algos.mlx.common.rotation` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = [ - "axis_angle_to_quat", - "np_matrix_from_quat", - "np_quat_angular_velocity", - "np_quat_apply", - "np_quat_apply_inverse", - "np_quat_canonicalize", - "np_quat_conjugate", - "np_quat_ensure_continuity", - "np_quat_error_magnitude", - "np_quat_from_euler_xyz", - "np_quat_inv", - "np_quat_mul", - "np_quat_to_axis_angle", - "np_sample_uniform", - "np_subtract_frame_transforms", - "np_yaw_quat", - "np_yaw_to_quat", - "quat_mul", -] diff --git a/src/unilab/utils/obs_utils.py b/src/unilab/utils/obs_utils.py deleted file mode 100644 index 3f1a2c384..000000000 --- a/src/unilab/utils/obs_utils.py +++ /dev/null @@ -1,26 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.base.observations import ( - flatten_obs_dict, - flatten_policy_obs_dict, - get_critic_base_dim, - get_obs_dims, - split_obs_dict, -) - -warnings.warn( - "`unilab.utils.obs_utils` is deprecated and will be removed in 0.2.0; " - "use `unilab.base.observations` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = [ - "flatten_obs_dict", - "flatten_policy_obs_dict", - "get_critic_base_dim", - "get_obs_dims", - "split_obs_dict", -] diff --git a/src/unilab/utils/offpolicy_logger.py b/src/unilab/utils/offpolicy_logger.py deleted file mode 100644 index 467d28846..000000000 --- a/src/unilab/utils/offpolicy_logger.py +++ /dev/null @@ -1,14 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.training.logging.offpolicy import OffPolicyLogger - -warnings.warn( - "`unilab.utils.offpolicy_logger` is deprecated and will be removed in 0.2.0; " - "use `unilab.training.logging.offpolicy` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = ["OffPolicyLogger"] diff --git a/src/unilab/utils/onpolicy_logger.py b/src/unilab/utils/onpolicy_logger.py deleted file mode 100644 index 2abe0fe23..000000000 --- a/src/unilab/utils/onpolicy_logger.py +++ /dev/null @@ -1,14 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.training.logging.onpolicy import OnPolicyLogger - -warnings.warn( - "`unilab.utils.onpolicy_logger` is deprecated and will be removed in 0.2.0; " - "use `unilab.training.logging.onpolicy` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = ["OnPolicyLogger"] diff --git a/src/unilab/utils/render_many.py b/src/unilab/utils/render_many.py deleted file mode 100644 index ac630630b..000000000 --- a/src/unilab/utils/render_many.py +++ /dev/null @@ -1,12 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.visualization.render_many import * # noqa: F403 - -warnings.warn( - "`unilab.utils.render_many` is deprecated and will be removed in 0.2.0; " - "use `unilab.visualization.render_many` instead.", - DeprecationWarning, - stacklevel=2, -) diff --git a/src/unilab/utils/reward_utils.py b/src/unilab/utils/reward_utils.py deleted file mode 100644 index 091ef256d..000000000 --- a/src/unilab/utils/reward_utils.py +++ /dev/null @@ -1,14 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.config.reward import RewardDict, extract_reward_config, resolve_reward_dict - -warnings.warn( - "`unilab.utils.reward_utils` is deprecated and will be removed in 0.2.0; " - "use `unilab.config.reward` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = ["RewardDict", "extract_reward_config", "resolve_reward_dict"] diff --git a/src/unilab/utils/rsl_rl_compat.py b/src/unilab/utils/rsl_rl_compat.py deleted file mode 100644 index 4f5f5485b..000000000 --- a/src/unilab/utils/rsl_rl_compat.py +++ /dev/null @@ -1,12 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.algos.torch.rsl_rl.compat import * # noqa: F403 - -warnings.warn( - "`unilab.utils.rsl_rl_compat` is deprecated and will be removed in 0.2.0; " - "use `unilab.algos.torch.rsl_rl.compat` instead.", - DeprecationWarning, - stacklevel=2, -) 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 6c805cd43..000000000 --- a/src/unilab/utils/rsl_rl_vec_env_wrapper.py +++ /dev/null @@ -1,14 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.algos.torch.rsl_rl.vec_env_wrapper import RslRlVecEnvWrapper - -warnings.warn( - "`unilab.utils.rsl_rl_vec_env_wrapper` is deprecated and will be removed in 0.2.0; " - "use `unilab.algos.torch.rsl_rl.vec_env_wrapper` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = ["RslRlVecEnvWrapper"] diff --git a/src/unilab/utils/run_utils.py b/src/unilab/utils/run_utils.py deleted file mode 100644 index 0f92bad83..000000000 --- a/src/unilab/utils/run_utils.py +++ /dev/null @@ -1,30 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.training.run import ( - get_entrypoint_log_root, - get_latest_checkpoint, - get_latest_run, - get_log_root, - parse_checkpoint_path, - resolve_checkpoint_path, - resolve_task_checkpoint_path, -) - -warnings.warn( - "`unilab.utils.run_utils` is deprecated and will be removed in 0.2.0; " - "use `unilab.training.run` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = [ - "get_entrypoint_log_root", - "get_latest_checkpoint", - "get_latest_run", - "get_log_root", - "parse_checkpoint_path", - "resolve_checkpoint_path", - "resolve_task_checkpoint_path", -] diff --git a/src/unilab/utils/torch_utils.py b/src/unilab/utils/torch_utils.py deleted file mode 100644 index 41697153a..000000000 --- a/src/unilab/utils/torch_utils.py +++ /dev/null @@ -1,14 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.utils.tensor import to_numpy, to_torch - -warnings.warn( - "`unilab.utils.torch_utils` is deprecated and will be removed in 0.2.0; " - "use `unilab.utils.tensor` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = ["to_numpy", "to_torch"] diff --git a/src/unilab/utils/viser_scene.py b/src/unilab/utils/viser_scene.py deleted file mode 100644 index bda072393..000000000 --- a/src/unilab/utils/viser_scene.py +++ /dev/null @@ -1,18 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.visualization.viser_scene import ( - VISER_AVAILABLE, - MujocoViserScene, - build_visible_env_indices, -) - -warnings.warn( - "`unilab.utils.viser_scene` is deprecated and will be removed in 0.2.0; " - "use `unilab.visualization.viser_scene` instead.", - DeprecationWarning, - stacklevel=2, -) - -__all__ = ["VISER_AVAILABLE", "MujocoViserScene", "build_visible_env_indices"] diff --git a/src/unilab/utils/xml_utils.py b/src/unilab/utils/xml_utils.py deleted file mode 100644 index 00ac762ab..000000000 --- a/src/unilab/utils/xml_utils.py +++ /dev/null @@ -1,12 +0,0 @@ -from __future__ import annotations - -import warnings - -from unilab.base.backend.xml import * # noqa: F403 - -warnings.warn( - "`unilab.utils.xml_utils` is deprecated and will be removed in 0.2.0; " - "use `unilab.base.backend.xml` instead.", - DeprecationWarning, - stacklevel=2, -) diff --git a/tests/utils/test_utils_package_policy.py b/tests/utils/test_utils_package_policy.py index 568fe1e73..31f6ce018 100644 --- a/tests/utils/test_utils_package_policy.py +++ b/tests/utils/test_utils_package_policy.py @@ -1,38 +1,66 @@ import importlib -import sys -import warnings 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")): + 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_each_utils_shim_is_importable_and_warns_with_removal_target() -> None: - shim_modules = sorted( - f"unilab.utils.{path.stem}" - for path in Path("src/unilab/utils").glob("*.py") - if path.stem not in {"__init__", "device", "tensor"} - ) - - for module_name in shim_modules: - sys.modules.pop(module_name, None) - with warnings.catch_warnings(record=True) as caught: - warnings.simplefilter("always", DeprecationWarning) - module = importlib.import_module(module_name) - assert module is not None - assert any(item.category is DeprecationWarning for item in caught), module_name - assert any("0.2.0" in str(item.message) for item in caught), module_name +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__ From e4a07f14eebe3198f45f96895b072a59cd77681b Mon Sep 17 00:00:00 2001 From: Yves Date: Fri, 24 Apr 2026 20:56:49 +0800 Subject: [PATCH 3/9] refactor: centralize rsl-rl training helpers --- scripts/play_interactive.py | 21 +-- scripts/play_viser.py | 21 +-- scripts/train_appo.py | 7 - scripts/train_rsl_rl.py | 13 +- src/unilab/algos/torch/appo/runner.py | 5 - src/unilab/algos/torch/appo/worker.py | 5 - src/unilab/training/rsl_rl.py | 221 ++++++++++++++++++++++++++ tests/algos/test_rsl_rl_runner.py | 5 +- tests/scripts/test_train_scripts.py | 12 +- 9 files changed, 245 insertions(+), 65 deletions(-) create mode 100644 src/unilab/training/rsl_rl.py diff --git a/scripts/play_interactive.py b/scripts/play_interactive.py index 94ecfe0a6..f9b5f79c9 100644 --- a/scripts/play_interactive.py +++ b/scripts/play_interactive.py @@ -50,16 +50,14 @@ 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.algos.torch.rsl_rl.compat import ( - convert_config_v3_to_v4, - convert_config_v5, - is_rsl_rl_v4, - is_rsl_rl_v5, -) -from unilab.algos.torch.rsl_rl.vec_env_wrapper import RslRlVecEnvWrapper from unilab.base import registry from unilab.config.structured_configs import PPOConfig as _StructuredPPOConfig @@ -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 8c6e27ac7..ccd44a82c 100644 --- a/scripts/play_viser.py +++ b/scripts/play_viser.py @@ -47,17 +47,15 @@ if str(ROOT_DIR) not in sys.path: sys.path.insert(0, str(ROOT_DIR)) -from unilab.algos.torch.rsl_rl.compat import ( - convert_config_v3_to_v4, - convert_config_v5, - is_rsl_rl_v4, - is_rsl_rl_v5, -) -from unilab.algos.torch.rsl_rl.vec_env_wrapper import RslRlVecEnvWrapper from unilab.training import ( ensure_registries, get_entrypoint_log_root, ) +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, @@ -244,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") @@ -274,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 1825d0fdb..b50401c00 100644 --- a/scripts/train_appo.py +++ b/scripts/train_appo.py @@ -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.algos.torch.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() @@ -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_rsl_rl.py b/scripts/train_rsl_rl.py index 1dd1f7847..20a69f8e3 100644 --- a/scripts/train_rsl_rl.py +++ b/scripts/train_rsl_rl.py @@ -16,7 +16,6 @@ if str(ROOT_DIR) not in sys.path: sys.path.insert(0, str(ROOT_DIR)) -from unilab.algos.torch.rsl_rl.vec_env_wrapper import RslRlVecEnvWrapper from unilab.base.backend.xml import materialize_scene_visual_override from unilab.training import ( BackendAdapter, @@ -29,6 +28,7 @@ render_play_mode, ) from unilab.training.logging.experiment import ExperimentTracker, patch_rsl_rl_wandb_writer +from unilab.training.rsl_rl import RslRlVecEnvWrapper, normalize_ppo_train_cfg 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.algos.torch.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/torch/appo/runner.py b/src/unilab/algos/torch/appo/runner.py index 84fc62013..e0a373a75 100644 --- a/src/unilab/algos/torch/appo/runner.py +++ b/src/unilab/algos/torch/appo/runner.py @@ -19,7 +19,6 @@ from unilab.algos.torch.appo.learner import APPOLearner from unilab.algos.torch.appo.worker import appo_collector_fn -from unilab.algos.torch.rsl_rl.compat import convert_config_v3_to_v4, is_rsl_rl_v4, is_rsl_rl_v5 from unilab.ipc import AsyncRunner, SharedOnPolicyStorage, SharedWeightSync from unilab.training.logging.offpolicy import OffPolicyLogger @@ -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 36f8c402f..ab5c78adb 100644 --- a/src/unilab/algos/torch/appo/worker.py +++ b/src/unilab/algos/torch/appo/worker.py @@ -76,7 +76,6 @@ def appo_collector_fn( from tensordict import TensorDict - from unilab.algos.torch.rsl_rl.compat import convert_config_v3_to_v4, is_rsl_rl_v4, is_rsl_rl_v5 from unilab.base import registry from unilab.ipc import SharedOnPolicyStorage, SharedWeightSync @@ -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/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/tests/algos/test_rsl_rl_runner.py b/tests/algos/test_rsl_rl_runner.py index 1d4a34210..7fa78af19 100644 --- a/tests/algos/test_rsl_rl_runner.py +++ b/tests/algos/test_rsl_rl_runner.py @@ -19,10 +19,10 @@ import torch from tensordict import TensorDict -from unilab.algos.torch.rsl_rl.compat import convert_config_v5, is_rsl_rl_v5 from unilab.base import registry from unilab.base.registry import ensure_registries from unilab.config.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/scripts/test_train_scripts.py b/tests/scripts/test_train_scripts.py index 961222759..f969be912 100644 --- a/tests/scripts/test_train_scripts.py +++ b/tests/scripts/test_train_scripts.py @@ -1043,7 +1043,7 @@ def _play_interactive(): def test_play_wrapper_imports_shared_implementation(): """Verify play_interactive.py uses shared RslRlVecEnvWrapper.""" - from unilab.algos.torch.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.algos.torch.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.algos.torch.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.algos.torch.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.algos.torch.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) From 57bfc2592385b1867d5efd4d00e5de0f3023eaf8 Mon Sep 17 00:00:00 2001 From: tatp-yf Date: Fri, 24 Apr 2026 22:32:43 +0800 Subject: [PATCH 4/9] refactor: move shared buffer base into ipc --- src/unilab/ipc/replay_buffer.py | 2 +- src/unilab/{algos/torch/common => ipc}/shared_buffer.py | 0 2 files changed, 1 insertion(+), 1 deletion(-) rename src/unilab/{algos/torch/common => ipc}/shared_buffer.py (100%) 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 From aadc3f46d18ace209fb745367a5acc0a0f8b77d7 Mon Sep 17 00:00:00 2001 From: tatp-yf Date: Fri, 24 Apr 2026 23:13:07 +0800 Subject: [PATCH 5/9] refactor: move dtype config out of base --- src/unilab/base/backend/mujoco_backend.py | 2 +- src/unilab/base/np_env.py | 2 +- src/unilab/dr/__init__.py | 5 +++++ src/unilab/dr/dr_utils.py | 2 +- src/unilab/{base => }/dtype_config.py | 0 src/unilab/envs/locomotion/common/base.py | 2 +- src/unilab/envs/locomotion/common/commands.py | 2 +- src/unilab/envs/locomotion/common/dr_provider.py | 2 +- src/unilab/envs/locomotion/common/rewards.py | 2 +- src/unilab/envs/locomotion/g1/joystick.py | 2 +- src/unilab/envs/locomotion/go1/joystick.py | 2 +- src/unilab/envs/locomotion/go2/handstand.py | 3 +-- src/unilab/envs/locomotion/go2/joystick.py | 2 +- src/unilab/envs/manipulation/inhand_rot_allegro/base.py | 2 +- src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py | 2 +- src/unilab/envs/manipulation/sharpa_inhand/base.py | 2 +- src/unilab/envs/manipulation/sharpa_inhand/rotation.py | 2 +- src/unilab/envs/motion_tracking/g1/tracking.py | 2 +- src/unilab/envs/motion_tracking/g1/tracking_sac.py | 2 +- 19 files changed, 22 insertions(+), 18 deletions(-) rename src/unilab/{base => }/dtype_config.py (100%) diff --git a/src/unilab/base/backend/mujoco_backend.py b/src/unilab/base/backend/mujoco_backend.py index 1b537878d..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 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/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/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 d77e50ba7..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,6 +24,7 @@ validate_interval_push_support, zero_actions, ) +from unilab.dtype_config import get_global_dtype from unilab.envs.common.rotation import np_quat_mul, np_yaw_to_quat 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 cdb56fdb0..2af270868 100644 --- a/src/unilab/envs/locomotion/go1/joystick.py +++ b/src/unilab/envs/locomotion/go1/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.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 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 ab3826c0e..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,6 +25,7 @@ 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 diff --git a/src/unilab/envs/manipulation/sharpa_inhand/base.py b/src/unilab/envs/manipulation/sharpa_inhand/base.py index bc5289afb..b0bb49ccd 100644 --- a/src/unilab/envs/manipulation/sharpa_inhand/base.py +++ b/src/unilab/envs/manipulation/sharpa_inhand/base.py @@ -10,8 +10,8 @@ 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.dtype_config import get_global_dtype from unilab.envs.common.rotation import np_quat_apply, np_quat_mul DEFAULT_ACTUATED_JOINT_NAMES: list[str] = [ diff --git a/src/unilab/envs/manipulation/sharpa_inhand/rotation.py b/src/unilab/envs/manipulation/sharpa_inhand/rotation.py index 55f3c81d1..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,7 @@ 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, diff --git a/src/unilab/envs/motion_tracking/g1/tracking.py b/src/unilab/envs/motion_tracking/g1/tracking.py index 8caad0cfc..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,6 +24,7 @@ validate_interval_push_support, zero_actions, ) +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, 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 From 79a28c8ee332695e19e2b6afddb9edddcffc24cc Mon Sep 17 00:00:00 2001 From: tatp-yf Date: Fri, 24 Apr 2026 23:24:19 +0800 Subject: [PATCH 6/9] refactor: move rich training loggers to unilab.logging --- scripts/train_appo.py | 2 +- scripts/train_mlx_ppo.py | 4 ++-- scripts/train_offpolicy.py | 2 +- scripts/train_rsl_rl.py | 2 +- src/unilab/algos/torch/appo/runner.py | 2 +- src/unilab/algos/torch/offpolicy/__init__.py | 2 +- .../algos/torch/offpolicy/multi_gpu_runner.py | 2 +- src/unilab/algos/torch/offpolicy/runner.py | 2 +- src/unilab/logging/__init__.py | 11 +++++++++ src/unilab/{training => }/logging/common.py | 0 .../{training => }/logging/offpolicy.py | 2 +- src/unilab/{training => }/logging/onpolicy.py | 2 +- src/unilab/training/__init__.py | 4 +--- .../training/{logging => }/experiment.py | 0 src/unilab/training/logging/__init__.py | 23 ------------------- tests/utils/test_experiment_tracking.py | 5 ++-- 16 files changed, 25 insertions(+), 40 deletions(-) create mode 100644 src/unilab/logging/__init__.py rename src/unilab/{training => }/logging/common.py (100%) rename src/unilab/{training => }/logging/offpolicy.py (99%) rename src/unilab/{training => }/logging/onpolicy.py (98%) rename src/unilab/training/{logging => }/experiment.py (100%) delete mode 100644 src/unilab/training/logging/__init__.py diff --git a/scripts/train_appo.py b/scripts/train_appo.py index b50401c00..295b87ef2 100644 --- a/scripts/train_appo.py +++ b/scripts/train_appo.py @@ -22,7 +22,7 @@ get_log_root, render_play_mode, ) -from unilab.training.logging.experiment import ExperimentTracker +from unilab.training.experiment import ExperimentTracker def build_appo_runner_kwargs( diff --git a/scripts/train_mlx_ppo.py b/scripts/train_mlx_ppo.py index b0fc3e67d..b5037050a 100644 --- a/scripts/train_mlx_ppo.py +++ b/scripts/train_mlx_ppo.py @@ -28,6 +28,7 @@ sys.path.append(str(ROOT_DIR)) from unilab.base.observations import flatten_obs_dict +from unilab.logging import OnPolicyLogger from unilab.training import ( BackendAdapter, create_env, @@ -43,8 +44,7 @@ from unilab.training import ( get_latest_run as get_latest_run_common, ) -from unilab.training.logging.experiment import ExperimentTracker -from unilab.training.logging.onpolicy import OnPolicyLogger +from unilab.training.experiment import ExperimentTracker ensure_registries() diff --git a/scripts/train_offpolicy.py b/scripts/train_offpolicy.py index 200919cb4..ab5fe6aaf 100644 --- a/scripts/train_offpolicy.py +++ b/scripts/train_offpolicy.py @@ -25,7 +25,7 @@ from unilab.training import ( resolve_checkpoint_path as resolve_checkpoint_path_common, ) -from unilab.training.logging.experiment import ExperimentTracker +from unilab.training.experiment import ExperimentTracker def default_device(torch_module, preferred: str | None = None) -> str: diff --git a/scripts/train_rsl_rl.py b/scripts/train_rsl_rl.py index 20a69f8e3..071a75716 100644 --- a/scripts/train_rsl_rl.py +++ b/scripts/train_rsl_rl.py @@ -27,7 +27,7 @@ parse_checkpoint_path, render_play_mode, ) -from unilab.training.logging.experiment import ExperimentTracker, patch_rsl_rl_wandb_writer +from unilab.training.experiment import ExperimentTracker, patch_rsl_rl_wandb_writer from unilab.training.rsl_rl import RslRlVecEnvWrapper, normalize_ppo_train_cfg try: diff --git a/src/unilab/algos/torch/appo/runner.py b/src/unilab/algos/torch/appo/runner.py index e0a373a75..af7b95770 100644 --- a/src/unilab/algos/torch/appo/runner.py +++ b/src/unilab/algos/torch/appo/runner.py @@ -20,7 +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.training.logging.offpolicy import OffPolicyLogger +from unilab.logging import OffPolicyLogger class APPORunner(AsyncRunner): diff --git a/src/unilab/algos/torch/offpolicy/__init__.py b/src/unilab/algos/torch/offpolicy/__init__.py index 69068d011..96d6afd3e 100644 --- a/src/unilab/algos/torch/offpolicy/__init__.py +++ b/src/unilab/algos/torch/offpolicy/__init__.py @@ -3,7 +3,7 @@ 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.training.logging.offpolicy import OffPolicyLogger +from unilab.logging import OffPolicyLogger __all__ = [ "OffPolicyLogger", diff --git a/src/unilab/algos/torch/offpolicy/multi_gpu_runner.py b/src/unilab/algos/torch/offpolicy/multi_gpu_runner.py index cc2f25c07..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.training.logging.offpolicy 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 cf548c668..0f8869b3b 100644 --- a/src/unilab/algos/torch/offpolicy/runner.py +++ b/src/unilab/algos/torch/offpolicy/runner.py @@ -14,7 +14,7 @@ from unilab.ipc import SharedObsNormStats, SharedWeightSync from unilab.ipc.async_runner import _SPAWN_CTX, AsyncRunner from unilab.ipc.replay_buffer import ReplayBuffer -from unilab.training.logging.offpolicy import OffPolicyLogger +from unilab.logging import OffPolicyLogger from unilab.utils.device import get_default_device 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/training/logging/common.py b/src/unilab/logging/common.py similarity index 100% rename from src/unilab/training/logging/common.py rename to src/unilab/logging/common.py diff --git a/src/unilab/training/logging/offpolicy.py b/src/unilab/logging/offpolicy.py similarity index 99% rename from src/unilab/training/logging/offpolicy.py rename to src/unilab/logging/offpolicy.py index 0e9bc6f51..c13c9360e 100644 --- a/src/unilab/training/logging/offpolicy.py +++ b/src/unilab/logging/offpolicy.py @@ -11,7 +11,7 @@ from rich.panel import Panel from rich.table import Table -from unilab.training.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): diff --git a/src/unilab/training/logging/onpolicy.py b/src/unilab/logging/onpolicy.py similarity index 98% rename from src/unilab/training/logging/onpolicy.py rename to src/unilab/logging/onpolicy.py index d166c2512..e72c54530 100644 --- a/src/unilab/training/logging/onpolicy.py +++ b/src/unilab/logging/onpolicy.py @@ -8,7 +8,7 @@ from rich.panel import Panel from rich.table import Table -from unilab.training.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/training/__init__.py b/src/unilab/training/__init__.py index d41e0c7c8..d36f7b5ea 100644 --- a/src/unilab/training/__init__.py +++ b/src/unilab/training/__init__.py @@ -8,7 +8,7 @@ get_hydra_runtime_choice, setup_logger, ) -from unilab.training.logging import ExperimentTracker, OffPolicyLogger, OnPolicyLogger +from unilab.training.experiment import ExperimentTracker from unilab.training.monitoring import HardwareMonitor from unilab.training.run import ( get_entrypoint_log_root, @@ -25,8 +25,6 @@ "BackendAdapter", "ExperimentTracker", "HardwareMonitor", - "OffPolicyLogger", - "OnPolicyLogger", "assert_offpolicy_task_choice_matches_algo", "create_env", "ensure_registries", diff --git a/src/unilab/training/logging/experiment.py b/src/unilab/training/experiment.py similarity index 100% rename from src/unilab/training/logging/experiment.py rename to src/unilab/training/experiment.py diff --git a/src/unilab/training/logging/__init__.py b/src/unilab/training/logging/__init__.py deleted file mode 100644 index a420aa4e5..000000000 --- a/src/unilab/training/logging/__init__.py +++ /dev/null @@ -1,23 +0,0 @@ -"""Training logging and experiment tracking helpers.""" - -from unilab.training.logging.experiment import ( - ExperimentTracker, - build_wandb_run_name, - build_wandb_settings, - get_device_info_dict, - get_git_info, - patch_rsl_rl_wandb_writer, -) -from unilab.training.logging.offpolicy import OffPolicyLogger -from unilab.training.logging.onpolicy import OnPolicyLogger - -__all__ = [ - "ExperimentTracker", - "OffPolicyLogger", - "OnPolicyLogger", - "build_wandb_run_name", - "build_wandb_settings", - "get_device_info_dict", - "get_git_info", - "patch_rsl_rl_wandb_writer", -] diff --git a/tests/utils/test_experiment_tracking.py b/tests/utils/test_experiment_tracking.py index 4e4a0ee38..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.training.logging.experiment import ExperimentTracker, build_wandb_settings -from unilab.training.logging.offpolicy import OffPolicyLogger -from unilab.training.logging.onpolicy import OnPolicyLogger +from unilab.logging import OffPolicyLogger, OnPolicyLogger +from unilab.training.experiment import ExperimentTracker, build_wandb_settings class _FakeConfig(dict): From 303b34d3a0b2b4812eb350e49c93d91d6930b781 Mon Sep 17 00:00:00 2001 From: tatp-yf Date: Fri, 24 Apr 2026 23:43:50 +0800 Subject: [PATCH 7/9] refactor: stop re-exporting playback from training --- scripts/train_appo.py | 2 +- scripts/train_mlx_ppo.py | 2 +- scripts/train_offpolicy.py | 2 +- scripts/train_rsl_rl.py | 2 +- src/unilab/training/__init__.py | 2 -- tests/training/test_training_helpers.py | 2 +- 6 files changed, 5 insertions(+), 7 deletions(-) diff --git a/scripts/train_appo.py b/scripts/train_appo.py index 295b87ef2..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.training.experiment import ExperimentTracker +from unilab.visualization import render_play_mode def build_appo_runner_kwargs( diff --git a/scripts/train_mlx_ppo.py b/scripts/train_mlx_ppo.py index b5037050a..29000dd66 100644 --- a/scripts/train_mlx_ppo.py +++ b/scripts/train_mlx_ppo.py @@ -35,7 +35,6 @@ ensure_registries, get_log_root, parse_checkpoint_path, - render_play_mode, setup_logger, ) from unilab.training import ( @@ -45,6 +44,7 @@ get_latest_run as get_latest_run_common, ) from unilab.training.experiment import ExperimentTracker +from unilab.visualization import render_play_mode ensure_registries() diff --git a/scripts/train_offpolicy.py b/scripts/train_offpolicy.py index ab5fe6aaf..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.training.experiment import ExperimentTracker +from unilab.visualization import render_play_mode def default_device(torch_module, preferred: str | None = None) -> str: diff --git a/scripts/train_rsl_rl.py b/scripts/train_rsl_rl.py index 071a75716..13d913eb1 100644 --- a/scripts/train_rsl_rl.py +++ b/scripts/train_rsl_rl.py @@ -25,10 +25,10 @@ get_latest_run, get_log_root, parse_checkpoint_path, - render_play_mode, ) 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 diff --git a/src/unilab/training/__init__.py b/src/unilab/training/__init__.py index d36f7b5ea..9da3921db 100644 --- a/src/unilab/training/__init__.py +++ b/src/unilab/training/__init__.py @@ -19,7 +19,6 @@ resolve_checkpoint_path, resolve_task_checkpoint_path, ) -from unilab.visualization.playback import render_play_mode __all__ = [ "BackendAdapter", @@ -34,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/tests/training/test_training_helpers.py b/tests/training/test_training_helpers.py index 6400947d1..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" From 02eddfd8d442dc77258b0493f842e70a5fdfef85 Mon Sep 17 00:00:00 2001 From: tatp-yf Date: Sat, 25 Apr 2026 00:31:04 +0800 Subject: [PATCH 8/9] refactor: retire unilab.config package structured_configs promoted to top-level unilab.structured_configs; reward helpers folded into unilab.training.reward. The two *_params modules held only unused KNOWN_TASKS stubs and are deleted. Eliminates the config/ vs conf/ naming collision at the source. --- AGENTS.md | 2 +- docs/developers/zh_CN/development-standard.md | 6 +++--- scripts/play_interactive.py | 2 +- src/unilab/config/__init__.py | 3 --- src/unilab/config/locomotion_params.py | 16 ---------------- src/unilab/config/manipulation_params.py | 11 ----------- src/unilab/{config => }/structured_configs.py | 0 src/unilab/training/backend_adapter.py | 2 +- src/unilab/{config => training}/reward.py | 0 tests/algos/test_appo_runner.py | 2 +- tests/algos/test_mlx_ppo.py | 2 +- tests/algos/test_offpolicy_runner.py | 2 +- tests/algos/test_rsl_rl_runner.py | 2 +- tests/config/test_locomotion_params.py | 12 ++++++------ tests/config/test_reward_injection.py | 2 +- tests/scripts/test_train_scripts.py | 2 +- 16 files changed, 18 insertions(+), 48 deletions(-) delete mode 100644 src/unilab/config/__init__.py delete mode 100644 src/unilab/config/locomotion_params.py delete mode 100644 src/unilab/config/manipulation_params.py rename src/unilab/{config => }/structured_configs.py (100%) rename src/unilab/{config => training}/reward.py (100%) diff --git a/AGENTS.md b/AGENTS.md index 450a8030a..35c76fa4e 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -36,7 +36,7 @@ UniLab 是一个 **高性能、模块化、contract 驱动** 的 RL infrastructu - 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/config/structured_configs.py` +- config schema: `src/unilab/structured_configs.py` - async runner: `src/unilab/ipc/async_runner.py` ## GitHub CLI (gh) 速查 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/scripts/play_interactive.py b/scripts/play_interactive.py index f9b5f79c9..5e5f1b72b 100644 --- a/scripts/play_interactive.py +++ b/scripts/play_interactive.py @@ -59,7 +59,7 @@ ensure_registries() from unilab.base import registry -from unilab.config.structured_configs import PPOConfig as _StructuredPPOConfig +from unilab.structured_configs import PPOConfig as _StructuredPPOConfig PPOConfig = _StructuredPPOConfig diff --git a/src/unilab/config/__init__.py b/src/unilab/config/__init__.py deleted file mode 100644 index 6684d663a..000000000 --- a/src/unilab/config/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from unilab.config.reward import RewardDict, extract_reward_config, resolve_reward_dict - -__all__ = ["RewardDict", "extract_reward_config", "resolve_reward_dict"] 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/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/backend_adapter.py b/src/unilab/training/backend_adapter.py index 21baea8e3..c915cf92a 100644 --- a/src/unilab/training/backend_adapter.py +++ b/src/unilab/training/backend_adapter.py @@ -8,7 +8,7 @@ from omegaconf import DictConfig, OmegaConf from unilab.base.backend.xml import materialize_scene_visual_override -from unilab.config.reward import extract_reward_config +from unilab.training.reward import extract_reward_config class BackendAdapter: diff --git a/src/unilab/config/reward.py b/src/unilab/training/reward.py similarity index 100% rename from src/unilab/config/reward.py rename to src/unilab/training/reward.py 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_mlx_ppo.py b/tests/algos/test_mlx_ppo.py index b935c23e2..1bbfe19a8 100644 --- a/tests/algos/test_mlx_ppo.py +++ b/tests/algos/test_mlx_ppo.py @@ -330,7 +330,7 @@ def test_mlx_ppo_one_iteration_real_env(default_go2_reward_config): from unilab.base import registry from unilab.base.observations import flatten_obs_dict from unilab.base.registry import ensure_registries - from unilab.config.structured_configs import PPOConfig as PPOStructuredConfig + 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 7fa78af19..28f9873d0 100644 --- a/tests/algos/test_rsl_rl_runner.py +++ b/tests/algos/test_rsl_rl_runner.py @@ -21,7 +21,7 @@ from unilab.base import registry from unilab.base.registry import ensure_registries -from unilab.config.structured_configs import PPOConfig +from unilab.structured_configs import PPOConfig from unilab.training.rsl_rl import normalize_ppo_train_cfg from unilab.utils.tensor import to_torch 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 a319c60a2..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.config.reward import resolve_reward_dict + from unilab.training.reward import resolve_reward_dict with initialize(config_path="../../conf/ppo", version_base="1.3"): cfg = compose( diff --git a/tests/scripts/test_train_scripts.py b/tests/scripts/test_train_scripts.py index f969be912..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.config.reward import extract_reward_config + from unilab.training.reward import extract_reward_config cfg = _appo_cfg(["task=g1_motion_tracking/motrix"]) From b0a1473350861dfc71502fc2ea4cd6dad58bf526 Mon Sep 17 00:00:00 2001 From: tatp-yf Date: Sat, 25 Apr 2026 01:31:09 +0800 Subject: [PATCH 9/9] fix: update stale imports in benchmark scripts benchmark/ wasn't covered by the earlier refactor commits, so two scripts still pointed at unilab.config.structured_configs and unilab.base.dtype_config. Ruff flagged the resulting import blocks as unsorted because the phantom paths confused isort grouping. Also silence a pre-existing mypy no-redef on the try/except import fallback so the pre-commit hook passes. --- benchmark/benchmark_fast_sac_backends.py | 7 +++++-- benchmark/benchmark_mujoco_backend_step_detail.py | 2 +- 2 files changed, 6 insertions(+), 3 deletions(-) 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 da5485553..4771ee6ee 100644 --- a/benchmark/benchmark_mujoco_backend_step_detail.py +++ b/benchmark/benchmark_mujoco_backend_step_detail.py @@ -46,7 +46,7 @@ from mujoco.batch_env import BatchEnvPool from unilab.base.backend.xml import create_discardvisual_xml -from unilab.base.dtype_config import get_global_dtype +from unilab.dtype_config import get_global_dtype matplotlib.use("Agg") import matplotlib.pyplot as plt