diff --git a/scripts/play_interactive.py b/scripts/play_interactive.py index 3c4e22ca0..e552ec2cc 100644 --- a/scripts/play_interactive.py +++ b/scripts/play_interactive.py @@ -50,6 +50,7 @@ sys.path.insert(0, str(ROOT_DIR)) from unilab.training import ( + algo_config_dict, ensure_registries, get_entrypoint_log_root, resolve_task_checkpoint_path, @@ -187,10 +188,7 @@ def _algo_config_dict(cfg: DictConfig | None) -> dict[str, Any]: """ if cfg is None: return cast(dict[str, Any], PPOConfig().to_dict()) - train_cfg_raw = OmegaConf.to_container(cfg.algo, resolve=True) - if not isinstance(train_cfg_raw, dict): - raise TypeError("cfg.algo must resolve to a dict") - return cast(dict[str, Any], train_cfg_raw) + return algo_config_dict(cfg) SUPPORTED_INTERACTIVE_ALGOS = ("ppo", "appo", "sac", "flashsac", "hora_distill") diff --git a/scripts/play_viser.py b/scripts/play_viser.py index 222d9dc6c..911abe4af 100644 --- a/scripts/play_viser.py +++ b/scripts/play_viser.py @@ -32,7 +32,7 @@ import sys import time from pathlib import Path -from typing import Any, cast +from typing import Any import hydra import numpy as np @@ -93,26 +93,13 @@ resolve_checkpoint, ) +from unilab.training import algo_config_dict # noqa: E402 + # --------------------------------------------------------------------------- # # Core viewer # # --------------------------------------------------------------------------- # -def _algo_config_dict(cfg: DictConfig) -> dict[str, Any]: - """Return the composed PPO algo config as a plain dict. - - Args: - cfg: Hydra config for the current playback run. - - Returns: - The resolved ``cfg.algo`` subtree as a mutable dict for rsl_rl. - """ - train_cfg_raw = OmegaConf.to_container(cfg.algo, resolve=True) - if not isinstance(train_cfg_raw, dict): - raise TypeError("cfg.algo must resolve to a dict") - return cast(dict[str, Any], train_cfg_raw) - - def _load_env_playback_model(env: Any, env_index: int) -> mujoco.MjModel: """Resolve the exact MuJoCo model for one playback env. @@ -247,7 +234,7 @@ def _create_env(env_count: int): playback_session, _policy_obs_mode, _checkpoint_path = create_rsl_rl_playback_session( playback_cfg=_build_playback_config(args, num_envs=num_envs), env_factory=_create_env, - algo_config=_algo_config_dict(cfg), + algo_config=algo_config_dict(cfg), root_dir=ROOT_DIR, device=device, checkpoint_resolver=resolve_checkpoint, diff --git a/scripts/train_appo.py b/scripts/train_appo.py index f3d3daf9d..9f94f55b5 100644 --- a/scripts/train_appo.py +++ b/scripts/train_appo.py @@ -19,12 +19,15 @@ from unilab.algos.appo.runtime import resolve_appo_runtime from unilab.training import ( BackendAdapter, + algo_config_dict, apply_configured_training_seed, + build_run_dir_name, create_env, ensure_registries, get_log_root, log_playback_plan, resolve_appo_checkpoint_path, + resolve_nan_guard_cfg, should_run_playback, ) from unilab.training.experiment import ExperimentTracker @@ -45,10 +48,7 @@ def build_appo_runner_kwargs( rl_cfg: dict[str, Any] | None = None, ) -> dict: if rl_cfg is None: - rl_cfg_raw = OmegaConf.to_container(cfg.algo, resolve=True) - if not isinstance(rl_cfg_raw, dict): - raise TypeError("cfg.algo must resolve to a dict") - rl_cfg = cast(dict[str, Any], rl_cfg_raw) + rl_cfg = algo_config_dict(cfg) runner_kwargs = { "env_name": cfg.training.task_name, @@ -73,16 +73,9 @@ def build_appo_runner_kwargs( raise FileNotFoundError(f"Could not resolve APPO resume checkpoint: {load_run}") runner_kwargs["resume_path"] = resume_path - nan_guard_cfg = getattr(cfg.training, "nan_guard", None) - if nan_guard_cfg is not None and getattr(nan_guard_cfg, "enabled", False): - from unilab.utils.nan_guard import NanGuardCfg - - runner_kwargs["nan_guard_cfg"] = NanGuardCfg( - enabled=True, - buffer_size=int(getattr(nan_guard_cfg, "buffer_size", 100)), - max_envs_to_dump=int(getattr(nan_guard_cfg, "max_envs_to_dump", 5)), - output_dir=getattr(nan_guard_cfg, "output_dir", None), - ) + nan_guard_cfg = resolve_nan_guard_cfg(cfg.training) + if nan_guard_cfg is not None: + runner_kwargs["nan_guard_cfg"] = nan_guard_cfg return runner_kwargs @@ -330,10 +323,7 @@ def main(cfg: DictConfig) -> None: ).build_task_env_cfg_override() # Convert algo config to plain dict for APPORunner / RSL-RL internals - rl_cfg_raw = OmegaConf.to_container(cfg.algo, resolve=True) - if not isinstance(rl_cfg_raw, dict): - raise TypeError("cfg.algo must resolve to a dict") - rl_cfg = cast(dict[str, Any], rl_cfg_raw) + rl_cfg = algo_config_dict(cfg) apply_appo_runtime_flags(rl_cfg, cfg, training_enabled=not cfg.training.play_only) appo_runtime = resolve_appo_runtime(rl_cfg, default_play_fn=play_appo) @@ -343,7 +333,7 @@ def main(cfg: DictConfig) -> None: log_dir = os.path.join( log_root, cfg.training.task_name, - f"{timestamp}_{cfg.training.sim_backend}", + build_run_dir_name(timestamp, str(cfg.training.sim_backend)), ) else: log_dir = cfg.training.log_dir diff --git a/scripts/train_him_ppo.py b/scripts/train_him_ppo.py index 559f7ad47..d2b3680e4 100644 --- a/scripts/train_him_ppo.py +++ b/scripts/train_him_ppo.py @@ -7,7 +7,7 @@ import hydra import torch -from omegaconf import DictConfig, OmegaConf +from omegaconf import DictConfig EXPORT_POLICY = False # set to True in __main__ block @@ -22,10 +22,12 @@ from unilab.base.backend import materialize_scene_visual_override from unilab.training import ( BackendAdapter, + algo_config_dict, + apply_env_nan_guard, + build_run_dir_name, create_env, ensure_registries, - get_latest_checkpoint, - get_latest_run, + format_play_checkpoint_error, get_log_root, parse_checkpoint_path, ) @@ -47,57 +49,15 @@ def _get_log_root(cfg: DictConfig) -> str: return str(get_log_root(ROOT_DIR, cfg)) -def _algo_config_dict(cfg: DictConfig) -> dict[str, Any]: - raw = OmegaConf.to_container(cfg.algo, resolve=True) - if not isinstance(raw, dict): - raise TypeError("cfg.algo must resolve to a dict") - return cast(dict[str, Any], raw) - - -def _format_play_checkpoint_error( - cfg: DictConfig, - *, - task_log_root: Path, - load_path: Path | None, - load_path_dir: Path | None, -) -> str: - selected_checkpoint = OmegaConf.select(cfg, "algo.checkpoint", default=-1) - checkpoint_hint = ( - f" algo.checkpoint={selected_checkpoint!r}" - if selected_checkpoint not in (None, "", -1, "-1") - else "" - ) - if load_path_dir is not None and load_path is None and checkpoint_hint: - reason = f"Requested checkpoint was not found under resolved_run={load_path_dir}." - elif not task_log_root.exists(): - reason = "Task log root does not exist." - else: - latest_run = get_latest_run(task_log_root) - if latest_run is None: - reason = "No run directories were found under the task log root." - elif get_latest_checkpoint(latest_run) is None: - reason = f"Resolved latest run has no model_*.pt checkpoint files: {latest_run}." - else: - reason = "Requested run or checkpoint could not be resolved." - - return ( - "Could not resolve a checkpoint for play mode. " - f"{reason} task={cfg.training.task_name} task_log_root={task_log_root} " - f"algo.load_run={cfg.algo.load_run!r}{checkpoint_hint}." - " Use algo.load_run= " - "and optionally algo.checkpoint=." - ) - - def play_him_ppo(cfg: DictConfig, device: str) -> str | None: """Play mode for HIM-PPO.""" - rl_cfg = _algo_config_dict(cfg) + rl_cfg = algo_config_dict(cfg) task_log_root = get_log_root(ROOT_DIR, cfg) / str(cfg.training.task_name) load_path, load_path_dir = parse_checkpoint_path(cfg, root_dir=ROOT_DIR) if load_path is None or load_path_dir is None or not load_path.exists(): print( - _format_play_checkpoint_error( + format_play_checkpoint_error( cfg, task_log_root=task_log_root, load_path=load_path, @@ -202,7 +162,9 @@ def main(cfg: DictConfig) -> None: timestamp = datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S") log_root = _get_log_root(cfg) log_dir = str( - Path(log_root) / cfg.training.task_name / f"{timestamp}_{cfg.training.sim_backend}" + Path(log_root) + / cfg.training.task_name + / build_run_dir_name(timestamp, str(cfg.training.sim_backend)) ) else: log_dir = None @@ -226,24 +188,10 @@ def main(cfg: DictConfig) -> None: env = create_env(cfg, num_envs=cfg.algo.num_envs, env_cfg_override=env_cfg_override) from unilab.training.rsl_rl import RslRlVecEnvWrapper - nan_guard_cfg = getattr(cfg.training, "nan_guard", None) - if nan_guard_cfg is not None and getattr(nan_guard_cfg, "enabled", False): - from unilab.utils.nan_guard import NanGuard, NanGuardCfg - - guard = NanGuard( - NanGuardCfg( - enabled=True, - buffer_size=int(getattr(nan_guard_cfg, "buffer_size", 100)), - max_envs_to_dump=int(getattr(nan_guard_cfg, "max_envs_to_dump", 5)), - output_dir=getattr(nan_guard_cfg, "output_dir", None), - ), - num_envs=env.num_envs, - supports_state_playback=env.play_capabilities.supports_physics_state_playback, - ) - env.set_nan_guard(guard) + apply_env_nan_guard(env, cfg.training) wrapped_env = RslRlVecEnvWrapper(env, device=device) - rl_cfg = _algo_config_dict(cfg) + rl_cfg = algo_config_dict(cfg) runner = HIMOnPolicyRunner(wrapped_env, rl_cfg, log_dir=log_dir, device=device) if cfg.algo.load_run != "-1": diff --git a/scripts/train_hora_distill.py b/scripts/train_hora_distill.py index b16383398..fab734b0c 100644 --- a/scripts/train_hora_distill.py +++ b/scripts/train_hora_distill.py @@ -1,5 +1,4 @@ import datetime -import json import sys from pathlib import Path from typing import Any, cast @@ -50,7 +49,7 @@ setup_logger, should_run_playback, ) -from unilab.training.experiment import get_device_info_dict +from unilab.training.experiment import get_device_info_dict, write_run_config_snapshot def _write_distill_run_config( @@ -69,8 +68,9 @@ def _write_distill_run_config( Returns: None. Writes `distill_run_config.json` into `log_dir`. """ - payload = { - "run": { + write_run_config_snapshot( + log_dir, + run_metadata={ "algo": "hora_distill", "task": str(OmegaConf.select(cfg, "training.task_name")), "sim_backend": str(OmegaConf.select(cfg, "training.sim_backend")), @@ -78,11 +78,10 @@ def _write_distill_run_config( "hardware": get_device_info_dict(), "teacher": teacher_metadata, }, - "config": OmegaConf.to_container(cfg, resolve=True), - } - with (log_dir / "distill_run_config.json").open("w", encoding="utf-8") as f: - json.dump(payload, f, indent=2, ensure_ascii=True) - f.write("\n") + full_cfg=cfg, + filename="distill_run_config.json", + trailing_newline=True, + ) def _build_env_cfg_override(cfg: DictConfig) -> dict[str, Any]: diff --git a/scripts/train_offpolicy.py b/scripts/train_offpolicy.py index cbb9e31bf..5b32993f0 100644 --- a/scripts/train_offpolicy.py +++ b/scripts/train_offpolicy.py @@ -33,10 +33,12 @@ from unilab.training import ( apply_configured_training_seed, assert_offpolicy_task_choice_matches_algo, + build_run_dir_name, create_env, ensure_registries, get_log_root, log_playback_plan, + resolve_nan_guard_cfg, should_run_playback, ) from unilab.training.experiment import ExperimentTracker @@ -57,7 +59,6 @@ resolve_offpolicy_checkpoint_path as resolve_checkpoint_path, ) from unilab.training.sim2sim import policy_load_dim_guard, resolve_sim2sim_config -from unilab.utils.nan_guard import NanGuardCfg def enable_faulthandler() -> None: @@ -84,11 +85,6 @@ def build_failure_summary(exc: BaseException, run_summary: Any | None = None) -> return summary -def build_run_dir_name(timestamp: str, sim_backend: str, *, world_size: int = 1) -> str: - gpu_suffix = f"_gpux{world_size}" if world_size > 1 else "" - return f"{timestamp}_{sim_backend}{gpu_suffix}" - - def build_offpolicy_env_cfg_override(algo_name: str, cfg: DictConfig) -> dict[str, Any] | None: return _build_offpolicy_env_cfg_override(algo_name, cfg, root_dir=ROOT_DIR) @@ -149,15 +145,7 @@ def build_runner(algo_name: str, cfg: DictConfig, log_dir: str | None = None): ) apply_torch_thread_runtime(torch_thread_runtime, role="learner") - nan_guard_cfg = getattr(cfg.training, "nan_guard", None) - _nan_guard_cfg: NanGuardCfg | None = None - if nan_guard_cfg is not None and getattr(nan_guard_cfg, "enabled", False): - _nan_guard_cfg = NanGuardCfg( - enabled=True, - buffer_size=int(getattr(nan_guard_cfg, "buffer_size", 100)), - max_envs_to_dump=int(getattr(nan_guard_cfg, "max_envs_to_dump", 5)), - output_dir=getattr(nan_guard_cfg, "output_dir", None), - ) + _nan_guard_cfg = resolve_nan_guard_cfg(cfg.training) replay_prefetch_mode = getattr(cfg.training, "replay_prefetch_mode", "one_tick") if replay_prefetch_mode != "one_tick": diff --git a/scripts/train_rsl_rl.py b/scripts/train_rsl_rl.py index bdadf732b..ccab57ac0 100644 --- a/scripts/train_rsl_rl.py +++ b/scripts/train_rsl_rl.py @@ -31,11 +31,13 @@ ) from unilab.training import ( BackendAdapter, + algo_config_dict, apply_configured_training_seed, + apply_env_nan_guard, + build_run_dir_name, create_env, ensure_registries, - get_latest_checkpoint, - get_latest_run, + format_play_checkpoint_error, get_log_root, log_playback_plan, parse_checkpoint_path, @@ -107,11 +109,6 @@ def _get_log_root(cfg: DictConfig) -> str: return str(get_log_root(ROOT_DIR, cfg)) -def build_ppo_run_dir_name(timestamp: str, sim_backend: str, *, world_size: int = 1) -> str: - gpu_suffix = f"_gpux{world_size}" if world_size > 1 else "" - return f"{timestamp}_{sim_backend}{gpu_suffix}" - - def resolve_ppo_log_dir( cfg: DictConfig, *, @@ -129,7 +126,7 @@ def resolve_ppo_log_dir( return str( Path(_get_log_root(cfg)) / str(cfg.training.task_name) - / build_ppo_run_dir_name( + / build_run_dir_name( timestamp, str(cfg.training.sim_backend), world_size=world_size, @@ -137,13 +134,6 @@ def resolve_ppo_log_dir( ) -def _algo_config_dict(cfg: DictConfig) -> dict[str, Any]: - train_cfg_raw = OmegaConf.to_container(cfg.algo, resolve=True) - if not isinstance(train_cfg_raw, dict): - raise TypeError("cfg.algo must resolve to a dict") - return cast(dict[str, Any], train_cfg_raw) - - def _resolve_ppo_wrapper_cls(rl_cfg: dict[str, Any]) -> type[RslRlVecEnvWrapper]: """Resolve the VecEnv wrapper class from the owner-selected PPO runtime. @@ -192,42 +182,6 @@ def validate_ppo_run_completion_topology( ) -def _format_play_checkpoint_error( - cfg: DictConfig, - *, - task_log_root: Path, - load_path: Path | None, - load_path_dir: Path | None, -) -> str: - selected_checkpoint = OmegaConf.select(cfg, "algo.checkpoint", default=-1) - checkpoint_hint = ( - f" algo.checkpoint={selected_checkpoint!r}" - if selected_checkpoint not in (None, "", -1, "-1") - else "" - ) - - if load_path_dir is not None and load_path is None and checkpoint_hint: - reason = f"Requested checkpoint was not found under resolved_run={load_path_dir}." - elif not task_log_root.exists(): - reason = "Task log root does not exist." - else: - latest_run = get_latest_run(task_log_root) - if latest_run is None: - reason = "No run directories were found under the task log root." - elif get_latest_checkpoint(latest_run) is None: - reason = f"Resolved latest run has no model_*.pt checkpoint files: {latest_run}." - else: - reason = "Requested run or checkpoint could not be resolved." - - return ( - "Could not resolve a checkpoint for play mode. " - f"{reason} task={cfg.training.task_name} task_log_root={task_log_root} " - f"algo.load_run={cfg.algo.load_run!r}{checkpoint_hint}." - " Use algo.load_run= " - "and optionally algo.checkpoint=." - ) - - def _resolve_play_num_steps(cfg: DictConfig) -> int | None: play_steps = OmegaConf.select(cfg, "training.play_steps", default=None) if play_steps is None: @@ -237,14 +191,14 @@ def _resolve_play_num_steps(cfg: DictConfig) -> int | None: def play_rsl_rl(cfg: DictConfig, device: str) -> str | None: """Play mode for RSL-RL.""" - rl_cfg = _algo_config_dict(cfg) + rl_cfg = algo_config_dict(cfg) wrapper_cls = _resolve_ppo_wrapper_cls(rl_cfg) task_log_root = get_log_root(ROOT_DIR, cfg) / str(cfg.training.task_name) load_path, load_path_dir = parse_checkpoint_path(cfg, root_dir=ROOT_DIR) if load_path is None or load_path_dir is None or not load_path.exists(): print( - _format_play_checkpoint_error( + format_play_checkpoint_error( cfg, task_log_root=task_log_root, load_path=load_path, @@ -450,26 +404,10 @@ def main(cfg: DictConfig) -> None: env_cfg_override=env_cfg_override, ) try: - rl_cfg = _algo_config_dict(cfg) + rl_cfg = algo_config_dict(cfg) wrapper_cls = _resolve_ppo_wrapper_cls(rl_cfg) - nan_guard_cfg = getattr(cfg.training, "nan_guard", None) - if nan_guard_cfg is not None and getattr(nan_guard_cfg, "enabled", False): - from unilab.utils.nan_guard import NanGuard, NanGuardCfg - - guard = NanGuard( - NanGuardCfg( - enabled=True, - buffer_size=int(getattr(nan_guard_cfg, "buffer_size", 100)), - max_envs_to_dump=int(getattr(nan_guard_cfg, "max_envs_to_dump", 5)), - output_dir=getattr(nan_guard_cfg, "output_dir", None), - ), - num_envs=env.num_envs, - supports_state_playback=( - env.play_capabilities.supports_physics_state_playback - ), - ) - env.set_nan_guard(guard) + apply_env_nan_guard(env, cfg.training) wrapped_env = wrapper_cls(env, device=device) diff --git a/src/unilab/training/__init__.py b/src/unilab/training/__init__.py index b02b343b7..117304dc2 100644 --- a/src/unilab/training/__init__.py +++ b/src/unilab/training/__init__.py @@ -11,7 +11,11 @@ from unilab.training.experiment import ExperimentTracker from unilab.training.monitoring import HardwareMonitor from unilab.training.run import ( + algo_config_dict, + apply_env_nan_guard, + build_run_dir_name, format_hora_stage2_checkpoint_error, + format_play_checkpoint_error, get_entrypoint_log_root, get_latest_checkpoint, get_latest_run, @@ -21,6 +25,7 @@ resolve_appo_checkpoint_path, resolve_checkpoint_path, resolve_hora_stage2_checkpoint_path, + resolve_nan_guard_cfg, resolve_offpolicy_checkpoint_path, resolve_task_checkpoint_path, should_run_playback, @@ -37,9 +42,13 @@ "BackendAdapter", "ExperimentTracker", "HardwareMonitor", + "algo_config_dict", + "apply_env_nan_guard", "assert_offpolicy_task_choice_matches_algo", + "build_run_dir_name", "create_env", "ensure_registries", + "format_play_checkpoint_error", "get_entrypoint_log_root", "get_hydra_runtime_choice", "get_latest_checkpoint", @@ -48,6 +57,7 @@ "log_playback_plan", "parse_checkpoint_path", "resolve_checkpoint_path", + "resolve_nan_guard_cfg", "resolve_task_checkpoint_path", "should_run_playback", "TrainingSeedInfo", diff --git a/src/unilab/training/experiment.py b/src/unilab/training/experiment.py index ffc9cda44..1dba2efdd 100644 --- a/src/unilab/training/experiment.py +++ b/src/unilab/training/experiment.py @@ -86,6 +86,49 @@ def _run_git(*args: str) -> str | None: } +def _write_json(path: Path, payload: Any, *, trailing_newline: bool = False) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + text = json.dumps(payload, indent=2) + if trailing_newline: + text += "\n" + path.write_text(text, encoding="utf-8") + + +def build_run_config_payload( + *, + run_metadata: dict[str, Any], + full_cfg: Any, + contract_snapshot: Any | None = None, +) -> dict[str, Any]: + """Assemble the canonical ``{"run", "config", ...}`` run-config payload.""" + payload: dict[str, Any] = { + "run": _json_safe(run_metadata), + "config": _json_safe(_plain_dict(full_cfg)), + } + if contract_snapshot is not None: + payload["contract_snapshot"] = _json_safe(contract_snapshot) + return payload + + +def write_run_config_snapshot( + log_dir: str | Path, + *, + run_metadata: dict[str, Any], + full_cfg: Any, + filename: str = "run_config.json", + contract_snapshot: Any | None = None, + trailing_newline: bool = False, +) -> dict[str, Any]: + """Write the run-config snapshot into ``log_dir`` and return the payload.""" + payload = build_run_config_payload( + run_metadata=run_metadata, + full_cfg=full_cfg, + contract_snapshot=contract_snapshot, + ) + _write_json(Path(log_dir) / filename, payload, trailing_newline=trailing_newline) + return payload + + 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}" @@ -219,12 +262,12 @@ def start(self) -> None: seed_payload = {"effective_seed": self.seed_info} metadata.update(seed_payload) - payload = { - "run": _json_safe(metadata), - "config": _json_safe(_plain_dict(self.full_cfg)), - "contract_snapshot": _json_safe(extract_contract_snapshot(self.full_cfg)), - } - self._write_json(self.log_dir / "run_config.json", payload) + payload = write_run_config_snapshot( + self.log_dir, + run_metadata=metadata, + full_cfg=self.full_cfg, + contract_snapshot=extract_contract_snapshot(self.full_cfg), + ) if not self.enabled: return @@ -318,8 +361,7 @@ def finish(self) -> 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") + _write_json(path, payload) def patch_rsl_rl_action_std_logging(runner: Any) -> None: diff --git a/src/unilab/training/run.py b/src/unilab/training/run.py index 86720b304..3a7ab5691 100644 --- a/src/unilab/training/run.py +++ b/src/unilab/training/run.py @@ -5,14 +5,100 @@ import os from os import PathLike from pathlib import Path +from typing import TYPE_CHECKING, Any, cast from omegaconf import DictConfig, OmegaConf from unilab.base.backend.base import BackendPlayRenderPlan, normalize_play_render_mode +if TYPE_CHECKING: + from unilab.utils.nan_guard import NanGuardCfg + _TEST_LOG_ROOT_ENV = "UNILAB_TEST_LOG_ROOT" +def build_run_dir_name(timestamp: str, sim_backend: str, *, world_size: int = 1) -> str: + """Return the canonical run directory name shared by all training entries.""" + gpu_suffix = f"_gpux{world_size}" if world_size > 1 else "" + return f"{timestamp}_{sim_backend}{gpu_suffix}" + + +def algo_config_dict(cfg: DictConfig) -> dict[str, Any]: + """Resolve the composed ``cfg.algo`` subtree into a plain mutable dict.""" + raw = OmegaConf.to_container(cfg.algo, resolve=True) + if not isinstance(raw, dict): + raise TypeError("cfg.algo must resolve to a dict") + return cast(dict[str, Any], raw) + + +def format_play_checkpoint_error( + cfg: DictConfig, + *, + task_log_root: Path, + load_path: Path | None, + load_path_dir: Path | None, +) -> str: + """Build the user-facing diagnostic for an unresolvable play checkpoint.""" + selected_checkpoint = OmegaConf.select(cfg, "algo.checkpoint", default=-1) + checkpoint_hint = ( + f" algo.checkpoint={selected_checkpoint!r}" + if selected_checkpoint not in (None, "", -1, "-1") + else "" + ) + + if load_path_dir is not None and load_path is None and checkpoint_hint: + reason = f"Requested checkpoint was not found under resolved_run={load_path_dir}." + elif not task_log_root.exists(): + reason = "Task log root does not exist." + else: + latest_run = get_latest_run(task_log_root) + if latest_run is None: + reason = "No run directories were found under the task log root." + elif get_latest_checkpoint(latest_run) is None: + reason = f"Resolved latest run has no model_*.pt checkpoint files: {latest_run}." + else: + reason = "Requested run or checkpoint could not be resolved." + + return ( + "Could not resolve a checkpoint for play mode. " + f"{reason} task={cfg.training.task_name} task_log_root={task_log_root} " + f"algo.load_run={cfg.algo.load_run!r}{checkpoint_hint}." + " Use algo.load_run= " + "and optionally algo.checkpoint=." + ) + + +def resolve_nan_guard_cfg(training_cfg: Any) -> NanGuardCfg | None: + """Build the shared ``NanGuardCfg`` from ``training.nan_guard``, or ``None``.""" + nan_guard_cfg = getattr(training_cfg, "nan_guard", None) + if nan_guard_cfg is None or not getattr(nan_guard_cfg, "enabled", False): + return None + from unilab.utils.nan_guard import NanGuardCfg + + return NanGuardCfg( + enabled=True, + buffer_size=int(getattr(nan_guard_cfg, "buffer_size", 100)), + max_envs_to_dump=int(getattr(nan_guard_cfg, "max_envs_to_dump", 5)), + output_dir=getattr(nan_guard_cfg, "output_dir", None), + ) + + +def apply_env_nan_guard(env: Any, training_cfg: Any) -> None: + """Attach a ``NanGuard`` to ``env`` when ``training.nan_guard`` is enabled.""" + nan_guard_cfg = resolve_nan_guard_cfg(training_cfg) + if nan_guard_cfg is None: + return + from unilab.utils.nan_guard import NanGuard + + env.set_nan_guard( + NanGuard( + nan_guard_cfg, + num_envs=env.num_envs, + supports_state_playback=env.play_capabilities.supports_physics_state_playback, + ) + ) + + def should_run_playback(*, play_only: bool, no_play: bool, play_render_mode: str | None) -> bool: """Return whether train/eval should enter playback for the configured mode.""" if normalize_play_render_mode(play_render_mode) == "none": diff --git a/tests/scripts/test_train_scripts.py b/tests/scripts/test_train_scripts.py index 48a8db802..83ba1398d 100644 --- a/tests/scripts/test_train_scripts.py +++ b/tests/scripts/test_train_scripts.py @@ -1096,7 +1096,7 @@ def create_env(*args: Any, **kwargs: Any) -> FakeEnv: return FakeEnv() monkeypatch.setattr(mod, "create_env", create_env) - monkeypatch.setattr(mod, "_algo_config_dict", lambda _cfg: {}) + monkeypatch.setattr(mod, "algo_config_dict", lambda _cfg: {}) monkeypatch.setattr(mod, "_resolve_ppo_wrapper_cls", lambda _rl_cfg: FakeWrapper) monkeypatch.setattr(mod, "normalize_ppo_train_cfg", lambda _rl_cfg: {}) monkeypatch.setattr(mod, "patch_rsl_rl_resume_state", lambda: None) @@ -3308,6 +3308,7 @@ def _fail_bootstrap() -> None: raise RuntimeError("bootstrap failed") training_mod.ensure_registries = _fail_bootstrap + training_mod.algo_config_dict = lambda cfg: {} training_mod.get_entrypoint_log_root = lambda *args, **kwargs: Path("/tmp") training_mod.resolve_task_checkpoint_path = lambda *args, **kwargs: (None, None) monkeypatch.setitem(sys.modules, "unilab.training", training_mod)