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
6 changes: 2 additions & 4 deletions scripts/play_interactive.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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")
Expand Down
21 changes: 4 additions & 17 deletions scripts/play_viser.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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,
Expand Down
28 changes: 9 additions & 19 deletions scripts/train_appo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand All @@ -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


Expand Down Expand Up @@ -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)

Expand All @@ -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
Expand Down
76 changes: 12 additions & 64 deletions scripts/train_him_ppo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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,
)
Expand All @@ -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=<run-dir-or-checkpoint-path> "
"and optionally algo.checkpoint=<iteration-or-filename>."
)


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,
Expand Down Expand Up @@ -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
Expand All @@ -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":
Expand Down
17 changes: 8 additions & 9 deletions scripts/train_hora_distill.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
import datetime
import json
import sys
from pathlib import Path
from typing import Any, cast
Expand Down Expand Up @@ -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(
Expand All @@ -69,20 +68,20 @@ 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")),
"log_dir": str(log_dir),
"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]:
Expand Down
18 changes: 3 additions & 15 deletions scripts/train_offpolicy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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:
Expand All @@ -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)

Expand Down Expand Up @@ -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":
Expand Down
Loading
Loading