diff --git a/src/unilab/scripts/train_appo.py b/src/unilab/scripts/train_appo.py index cf2cfc060..371a6c3d7 100644 --- a/src/unilab/scripts/train_appo.py +++ b/src/unilab/scripts/train_appo.py @@ -31,6 +31,7 @@ ensure_registries, get_log_root, is_viser_play_render_mode, + nonfatal_play_step, resolve_nan_guard_cfg, should_run_playback, ) @@ -250,25 +251,26 @@ def play_appo( # Export actor to ONNX if load_path_dir is not None: - import torch.nn as nn + with nonfatal_play_step("ONNX export"): + import torch.nn as nn - class _DeterministicAPPOActor(nn.Module): - def __init__(self, mlp: nn.Module): - super().__init__() - self.mlp = mlp + class _DeterministicAPPOActor(nn.Module): + def __init__(self, mlp: nn.Module): + super().__init__() + self.mlp = mlp - def forward(self, obs: torch.Tensor) -> torch.Tensor: - return self.mlp(obs) + def forward(self, obs: torch.Tensor) -> torch.Tensor: + return self.mlp(obs) - export_module = _DeterministicAPPOActor(actor.mlp) - onnx_path = os.path.join(load_path_dir, "policy.onnx") - obs_dim = int(session.wrapped_env.num_obs) - dummy_input = torch.randn(1, obs_dim, device=device) - export_policy_onnx(export_module, onnx_path, (dummy_input,), input_names=["obs"]) + export_module = _DeterministicAPPOActor(actor.mlp) + onnx_path = os.path.join(load_path_dir, "policy.onnx") + obs_dim = int(session.wrapped_env.num_obs) + dummy_input = torch.randn(1, obs_dim, device=device) + export_policy_onnx(export_module, onnx_path, (dummy_input,), input_names=["obs"]) - # Verify ONNX output matches PyTorch - verify_input = torch.randn(1, obs_dim, device=device) - verify_policy_onnx(export_module, onnx_path, (verify_input,), input_names=["obs"]) + # Verify ONNX output matches PyTorch + verify_input = torch.randn(1, obs_dim, device=device) + verify_policy_onnx(export_module, onnx_path, (verify_input,), input_names=["obs"]) if is_viser_play_render_mode(getattr(cfg.training, "play_render_mode", "auto")): # Browser-based viser playback renders through the shared MuJoCo @@ -281,19 +283,23 @@ def forward(self, obs: torch.Tensor) -> torch.Tensor: run_viser_playback_from_cfg(session, cfg, entrypoint="train_appo play") return None - with torch.inference_mode(): - play_video_path = env.run_playback_mode( - play_render_mode=getattr(cfg.training, "play_render_mode", "auto"), - play_steps=getattr(cfg.training, "play_steps", None), - output_video=os.path.join(load_path_dir, "play_video.mp4") if load_path_dir else None, - render_spacing=float( - getattr(cfg.training, "render_spacing", getattr(env.cfg, "render_spacing", 1.0)) - ), - initialize=session.reset, - step=lambda _obs: session.step_once(), - camera_kwargs=camera_cfg_from_training(cfg.training), - on_plan=log_playback_plan, - ) + play_video_path: str | None = None + with nonfatal_play_step("video rendering"): + with torch.inference_mode(): + play_video_path = env.run_playback_mode( + play_render_mode=getattr(cfg.training, "play_render_mode", "auto"), + play_steps=getattr(cfg.training, "play_steps", None), + output_video=os.path.join(load_path_dir, "play_video.mp4") + if load_path_dir + else None, + render_spacing=float( + getattr(cfg.training, "render_spacing", getattr(env.cfg, "render_spacing", 1.0)) + ), + initialize=session.reset, + step=lambda _obs: session.step_once(), + camera_kwargs=camera_cfg_from_training(cfg.training), + on_plan=log_playback_plan, + ) if play_video_path is not None: print(f"Saving video to {play_video_path} ...") print("Done.") diff --git a/src/unilab/scripts/train_offpolicy.py b/src/unilab/scripts/train_offpolicy.py index ba8dceb79..2304744bd 100644 --- a/src/unilab/scripts/train_offpolicy.py +++ b/src/unilab/scripts/train_offpolicy.py @@ -45,6 +45,7 @@ ensure_registries, get_log_root, is_viser_play_render_mode, + nonfatal_play_step, resolve_nan_guard_cfg, should_run_playback, ) @@ -373,27 +374,28 @@ def play_offpolicy( # Export actor to ONNX if load_path_dir is not None and bool(getattr(cfg.training, "export_onnx", True)): - obs_dim, _ = resolve_play_obs_dims(env.obs_groups_spec) - onnx_path = os.path.join(load_path_dir, "policy.onnx") - dummy_input = torch.randn(1, obs_dim, device=device) - with torch.inference_mode(): - if normalizer: - dummy_input = normalizer(dummy_input, update=False) - assert actor is not None - if algo_name in ("sac", "flashsac", "warpsac"): - export_module = actor.as_export_module() - else: - export_module = actor - export_inputs = (dummy_input,) - input_names = ["obs"] - export_policy_onnx(export_module, onnx_path, export_inputs, input_names=input_names) - - # Verify ONNX output matches PyTorch - verify_input = torch.randn(1, obs_dim, device=device) - with torch.inference_mode(): - onnx_feed = normalizer(verify_input, update=False) if normalizer else verify_input - verify_inputs = (onnx_feed,) - verify_policy_onnx(export_module, onnx_path, verify_inputs, input_names=input_names) + with nonfatal_play_step("ONNX export"): + obs_dim, _ = resolve_play_obs_dims(env.obs_groups_spec) + onnx_path = os.path.join(load_path_dir, "policy.onnx") + dummy_input = torch.randn(1, obs_dim, device=device) + with torch.inference_mode(): + if normalizer: + dummy_input = normalizer(dummy_input, update=False) + assert actor is not None + if algo_name in ("sac", "flashsac", "warpsac"): + export_module = actor.as_export_module() + else: + export_module = actor + export_inputs = (dummy_input,) + input_names = ["obs"] + export_policy_onnx(export_module, onnx_path, export_inputs, input_names=input_names) + + # Verify ONNX output matches PyTorch + verify_input = torch.randn(1, obs_dim, device=device) + with torch.inference_mode(): + onnx_feed = normalizer(verify_input, update=False) if normalizer else verify_input + verify_inputs = (onnx_feed,) + verify_policy_onnx(export_module, onnx_path, verify_inputs, input_names=input_names) elif load_path_dir is not None: print("Skipping ONNX export because training.export_onnx=false.") @@ -408,16 +410,20 @@ def play_offpolicy( run_viser_playback_from_cfg(session, cfg, entrypoint=f"train_{algo_name} play") return None - with torch.inference_mode(): - play_video_path = env.run_playback_mode( - play_render_mode=getattr(cfg.training, "play_render_mode", "auto"), - play_steps=getattr(cfg.training, "play_steps", None), - output_video=os.path.join(load_path_dir, "play_video.mp4") if load_path_dir else None, - initialize=session.reset, - step=lambda _obs: session.step_once(), - camera_kwargs=camera_cfg_from_training(cfg.training), - on_plan=log_playback_plan, - ) + play_video_path: str | None = None + with nonfatal_play_step("video rendering"): + with torch.inference_mode(): + play_video_path = env.run_playback_mode( + play_render_mode=getattr(cfg.training, "play_render_mode", "auto"), + play_steps=getattr(cfg.training, "play_steps", None), + output_video=os.path.join(load_path_dir, "play_video.mp4") + if load_path_dir + else None, + initialize=session.reset, + step=lambda _obs: session.step_once(), + camera_kwargs=camera_cfg_from_training(cfg.training), + on_plan=log_playback_plan, + ) if play_video_path is not None: print(f"Saving video to {play_video_path} ...") print("Done.") diff --git a/src/unilab/scripts/train_rsl_rl.py b/src/unilab/scripts/train_rsl_rl.py index e475ceffb..993353086 100644 --- a/src/unilab/scripts/train_rsl_rl.py +++ b/src/unilab/scripts/train_rsl_rl.py @@ -46,6 +46,7 @@ format_play_checkpoint_error, get_log_root, is_viser_play_render_mode, + nonfatal_play_step, parse_checkpoint_path, should_run_playback, ) @@ -375,8 +376,9 @@ def play_rsl_rl(cfg: DictConfig, device: str) -> str | None: if EXPORT_POLICY: # The checkpoint early-returns above guarantee a loaded runner here. assert runner is not None - runner.export_policy_to_onnx(path=str(load_path_dir)) - runner.export_policy_to_jit(path=str(load_path_dir)) + with nonfatal_play_step("policy export"): + runner.export_policy_to_onnx(path=str(load_path_dir)) + runner.export_policy_to_jit(path=str(load_path_dir)) if is_viser_play_render_mode(getattr(cfg.training, "play_render_mode", "auto")): # Browser-based viser playback renders through the shared MuJoCo # playback shell; it replaces backend-native playback and records no @@ -397,25 +399,28 @@ def _log_plan(plan) -> None: playback_mode = plan.mode log_playback_plan(plan) - try: - with torch.inference_mode(): - play_video_path = env.run_playback_mode( - play_render_mode=getattr(cfg.training, "play_render_mode", "auto"), - play_steps=num_steps, - output_video=output_video, - render_spacing=float( - getattr(cfg.training, "render_spacing", getattr(env.cfg, "render_spacing", 1.0)) - ), - render_offset_mode=str(getattr(env.cfg, "render_offset_mode", "grid")), - initialize=session.reset, - step=lambda _obs: session.step_once(), - camera_kwargs=camera_cfg_from_training(cfg.training), - on_plan=_log_plan, - debug_overlay_getter=_playback_debug_overlay_getter(env), - ) - except RenderClosedError: - # Interface-level signal: the user closed the backend render window. - print("Render window closed.") + with nonfatal_play_step("video rendering"): + try: + with torch.inference_mode(): + play_video_path = env.run_playback_mode( + play_render_mode=getattr(cfg.training, "play_render_mode", "auto"), + play_steps=num_steps, + output_video=output_video, + render_spacing=float( + getattr( + cfg.training, "render_spacing", getattr(env.cfg, "render_spacing", 1.0) + ) + ), + render_offset_mode=str(getattr(env.cfg, "render_offset_mode", "grid")), + initialize=session.reset, + step=lambda _obs: session.step_once(), + camera_kwargs=camera_cfg_from_training(cfg.training), + on_plan=_log_plan, + debug_overlay_getter=_playback_debug_overlay_getter(env), + ) + except RenderClosedError: + # Interface-level signal: the user closed the backend render window. + print("Render window closed.") if playback_mode != "none" and num_steps is not None: print("Done.") return play_video_path diff --git a/src/unilab/training/__init__.py b/src/unilab/training/__init__.py index 32d60e4be..f960687fc 100644 --- a/src/unilab/training/__init__.py +++ b/src/unilab/training/__init__.py @@ -23,6 +23,7 @@ format_play_checkpoint_error, get_log_root, is_viser_play_render_mode, + nonfatal_play_step, parse_checkpoint_path, resolve_nan_guard_cfg, should_run_playback, @@ -60,6 +61,7 @@ "get_log_root", "is_viser_play_render_mode", "log_playback_plan", + "nonfatal_play_step", "parse_checkpoint_path", "resolve_checkpoint_path", "resolve_nan_guard_cfg", diff --git a/src/unilab/training/onnx_export.py b/src/unilab/training/onnx_export.py index df9e7d3ba..3c93b97b1 100644 --- a/src/unilab/training/onnx_export.py +++ b/src/unilab/training/onnx_export.py @@ -13,7 +13,7 @@ def export_policy_onnx( *, input_names: list[str], output_names: list[str] | None = None, - opset_version: int = 17, + opset_version: int = 18, ) -> None: """Export ``export_module`` to ``onnx_path`` and print the artifact path. @@ -23,7 +23,11 @@ def export_policy_onnx( export_inputs: Positional example inputs matching ``input_names``. input_names: ONNX input names, aligned positionally with ``export_inputs``. output_names: ONNX output names; defaults to ``["action"]``. - opset_version: ONNX opset version; defaults to 17. + opset_version: ONNX opset version; defaults to 18. Lower versions fail + in the dynamo exporter's InlinePass: torch's ONNX function library + (e.g. ``aten_isnan`` from ``torch.nan_to_num``) is emitted at opset + 18, and version-converting the model below that leaves a function + the inliner rejects as an opset mismatch. """ if output_names is None: output_names = ["action"] diff --git a/src/unilab/training/run.py b/src/unilab/training/run.py index 6f4829320..6dbcd1b50 100644 --- a/src/unilab/training/run.py +++ b/src/unilab/training/run.py @@ -3,6 +3,9 @@ from __future__ import annotations import os +import traceback +from collections.abc import Iterator +from contextlib import contextmanager from os import PathLike from pathlib import Path from typing import TYPE_CHECKING, Any, cast @@ -124,6 +127,17 @@ def should_run_playback(*, play_only: bool, no_play: bool, play_render_mode: str return bool(play_only) or not bool(no_play) +@contextmanager +def nonfatal_play_step(step: str) -> Iterator[None]: + """Keep independent post-training play artifacts (ONNX export, video render) + from blocking each other: log the failure and continue with the rest.""" + try: + yield + except Exception: + print(f"WARNING: {step} failed; continuing with the remaining play steps.") + traceback.print_exc() + + 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") diff --git a/tests/scripts/test_train_scripts.py b/tests/scripts/test_train_scripts.py index 26400f00c..3bb0b2b81 100644 --- a/tests/scripts/test_train_scripts.py +++ b/tests/scripts/test_train_scripts.py @@ -1866,6 +1866,178 @@ def run_playback_mode(self, **kwargs): assert not (run_dir / "policy.onnx").exists() +def _patch_offpolicy_play_session( + monkeypatch: pytest.MonkeyPatch, + mod, + cfg, + run_dir: Path, + checkpoint: Path, + captured: dict[str, Any], + fake_actor, + fake_env_cls, +) -> None: + import uni_rl.algos.common.actor_factory as actor_factory + + import unilab.utils.checkpoint as checkpoint_utils + + monkeypatch.setattr(mod, "build_offpolicy_env_cfg_override", lambda algo_name, cfg: {}) + monkeypatch.setattr(mod, "default_device", lambda torch_module, preferred=None: "cpu") + monkeypatch.setattr(mod, "create_env", lambda *args, **kwargs: fake_env_cls()) + monkeypatch.setattr( + mod, + "resolve_checkpoint_path", + lambda *args, **kwargs: (str(checkpoint), str(run_dir)), + ) + monkeypatch.setattr( + checkpoint_utils, + "resolve_offpolicy_checkpoint_path", + lambda *args, **kwargs: (str(checkpoint), str(run_dir)), + ) + monkeypatch.setattr(actor_factory, "build_actor", lambda *args, **kwargs: fake_actor) + + +class _PlayFakeActor: + def eval(self): + return self + + def load_state_dict(self, state_dict): + pass + + def as_export_module(self): + return self + + def explore(self, obs, deterministic=True): + import torch + + return torch.zeros((obs.shape[0], 2), dtype=obs.dtype, device=obs.device) + + +def _make_play_fake_env(captured: dict[str, Any], play_env_num: int): + import numpy as np + + class FakeEnv: + def __init__(self): + self.obs_groups_spec = {"obs": 4} + self.action_space = type("ActionSpace", (), {"shape": (2,)})() + self.state = None + + def init_state(self): + self.state = type( + "State", + (), + {"obs": {"obs": np.zeros((play_env_num, 4), dtype=np.float32)}}, + )() + + def reset(self, env_ids): + batch = len(env_ids) + return ({"obs": np.zeros((batch, 4), dtype=np.float32)}, {}) + + def step(self, actions): + batch = actions.shape[0] + self.state = type( + "State", + (), + { + "obs": {"obs": np.ones((batch, 4), dtype=np.float32)}, + "info": {}, + }, + )() + return self.state + + def run_playback_mode(self, **kwargs): + init_obs = kwargs["initialize"]() + kwargs["step"](init_obs) + return str(kwargs["output_video"]) + + return FakeEnv + + +def test_play_offpolicy_onnx_export_failure_still_records_video( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, capsys: pytest.CaptureFixture[str] +): + import torch + + mod = _offpolicy() + cfg = _offpolicy_cfg( + [ + "task=g1_walk_flat/mujoco", + "training.play_only=true", + "training.play_render_mode=record", + ] + ) + run_dir = tmp_path / "run" + run_dir.mkdir() + checkpoint = run_dir / "model_5000.pt" + torch.save({"actor": {}}, checkpoint) + + captured: dict[str, Any] = {} + _patch_offpolicy_play_session( + monkeypatch, + mod, + cfg, + run_dir, + checkpoint, + captured, + _PlayFakeActor(), + _make_play_fake_env(captured, cfg.training.play_env_num), + ) + monkeypatch.setattr( + mod, + "export_policy_onnx", + lambda *args, **kwargs: (_ for _ in ()).throw(RuntimeError("onnx boom")), + ) + + result = mod.play_offpolicy("sac", cfg) + out = capsys.readouterr().out + + assert result == str(run_dir / "play_video.mp4") + assert "WARNING: ONNX export failed; continuing with the remaining play steps." in out + + +def test_play_offpolicy_video_render_failure_returns_none( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, capsys: pytest.CaptureFixture[str] +): + import torch + + mod = _offpolicy() + cfg = _offpolicy_cfg( + [ + "task=g1_walk_flat/mujoco", + "training.play_only=true", + "training.play_render_mode=record", + "training.export_onnx=false", + ] + ) + run_dir = tmp_path / "run" + run_dir.mkdir() + checkpoint = run_dir / "model_5000.pt" + torch.save({"actor": {}}, checkpoint) + + captured: dict[str, Any] = {} + fake_env_cls = _make_play_fake_env(captured, cfg.training.play_env_num) + + class FailingRenderEnv(fake_env_cls): + def run_playback_mode(self, **kwargs): + raise RuntimeError("render boom") + + _patch_offpolicy_play_session( + monkeypatch, + mod, + cfg, + run_dir, + checkpoint, + captured, + _PlayFakeActor(), + FailingRenderEnv, + ) + + result = mod.play_offpolicy("sac", cfg) + out = capsys.readouterr().out + + assert result is None + assert "WARNING: video rendering failed; continuing with the remaining play steps." in out + + # --------------------------------------------------------------------------- # play_interactive.py — resolve_checkpoint() # --------------------------------------------------------------------------- diff --git a/tests/training/test_onnx_export.py b/tests/training/test_onnx_export.py index 8f5bf67cf..617781af9 100644 --- a/tests/training/test_onnx_export.py +++ b/tests/training/test_onnx_export.py @@ -41,6 +41,29 @@ def forward(self, obs: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: return action, action +class _NanGuardedActor(torch.nn.Module): + """Mimics the uni_rl SAC actor export path (``torch.nan_to_num`` → ``aten_isnan``).""" + + def __init__(self, obs_dim: int = 4, action_dim: int = 2) -> None: + super().__init__() + self.mlp = torch.nn.Linear(obs_dim, action_dim) + + def forward(self, obs: torch.Tensor) -> torch.Tensor: + return torch.nan_to_num(torch.tanh(self.mlp(obs)), nan=0.0) + + +def test_export_policy_onnx_supports_nan_guarded_actor(tmp_path): + """Regression: opset-18 function-library nodes must survive the dynamo exporter.""" + torch.manual_seed(0) + module = _NanGuardedActor().eval() + onnx_path = str(tmp_path / "policy.onnx") + + export_policy_onnx(module, onnx_path, (torch.randn(1, 4),), input_names=["obs"]) + + max_diff, _ = verify_policy_onnx(module, onnx_path, (torch.randn(1, 4),), input_names=["obs"]) + assert max_diff <= 1e-4 + + def test_export_policy_onnx_writes_expected_graph(tmp_path, capsys): torch.manual_seed(0) module = _TinyActor() @@ -53,7 +76,7 @@ def test_export_policy_onnx_writes_expected_graph(tmp_path, capsys): assert [graph_input.name for graph_input in model.graph.input] == ["obs"] assert [graph_output.name for graph_output in model.graph.output] == ["action"] opsets = {opset.version for opset in model.opset_import if opset.domain in ("", "ai.onnx")} - assert opsets == {17} + assert opsets == {18} def test_verify_policy_onnx_matches_pytorch(tmp_path, capsys): diff --git a/tests/training/test_training_helpers.py b/tests/training/test_training_helpers.py index 8e7ed3b62..fe7fd38cc 100644 --- a/tests/training/test_training_helpers.py +++ b/tests/training/test_training_helpers.py @@ -25,6 +25,7 @@ from unilab.base.scene import SceneCfg from unilab.training import ( get_log_root, + nonfatal_play_step, parse_checkpoint_path, ) from unilab.utils.checkpoint import ( @@ -84,6 +85,21 @@ def _offpolicy_cfg(overrides: list[str] | None = None, *, algo: str = "sac"): return compose("config", overrides=_normalize_overrides(overrides, offpolicy=True)) +def test_nonfatal_play_step_logs_failure_and_continues(capsys: pytest.CaptureFixture[str]): + with nonfatal_play_step("ONNX export"): + raise RuntimeError("boom") + + captured = capsys.readouterr() + assert "WARNING: ONNX export failed; continuing with the remaining play steps." in captured.out + assert "RuntimeError: boom" in captured.err + + +def test_nonfatal_play_step_does_not_swallow_keyboard_interrupt(): + with pytest.raises(KeyboardInterrupt): + with nonfatal_play_step("video rendering"): + raise KeyboardInterrupt + + def test_get_latest_run_and_checkpoint_support_shared_checkpoint_resolution(tmp_path: Path): task_dir = tmp_path / "logs" / "custom_ppo" / "MyTask" older_run = task_dir / "2024-01-01_00-00-00_mujoco"