Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
62 changes: 34 additions & 28 deletions src/unilab/scripts/train_appo.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
ensure_registries,
get_log_root,
is_viser_play_render_mode,
nonfatal_play_step,
resolve_nan_guard_cfg,
should_run_playback,
)
Expand Down Expand Up @@ -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
Expand All @@ -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.")
Expand Down
68 changes: 37 additions & 31 deletions src/unilab/scripts/train_offpolicy.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@
ensure_registries,
get_log_root,
is_viser_play_render_mode,
nonfatal_play_step,
resolve_nan_guard_cfg,
should_run_playback,
)
Expand Down Expand Up @@ -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.")

Expand All @@ -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.")
Expand Down
47 changes: 26 additions & 21 deletions src/unilab/scripts/train_rsl_rl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
2 changes: 2 additions & 0 deletions src/unilab/training/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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",
Expand Down
8 changes: 6 additions & 2 deletions src/unilab/training/onnx_export.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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"]
Expand Down
14 changes: 14 additions & 0 deletions src/unilab/training/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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")
Expand Down
Loading
Loading