agent-parkour / infer.py
anngo-1's picture
Deploy Agent Parkour app
386f43a verified
Raw
History Blame Contribute Delete
4 kB
from __future__ import annotations
import argparse
from dataclasses import replace
from pathlib import Path
from typing import Any
import torch
from env import FastFloatingBeaconEnv
from replay import V2RunWriter
from ppo import evaluate, replay_payload
from runner import BeaconController, LocalBeaconPlannerController, TokenAttentionController
from runtime import normalize_config, resolve_device
from mapgen import configure_motif_mapgen, install_motif_mapgen
from settings import ACTIVE_CHECKPOINT, MAPGEN_KWARGS
def planner_mode_for_state(state: dict[str, torch.Tensor]) -> str:
return "local-beacon" if any(key.startswith("runner.") for key in state) else "direct"
def policy_arch_for_state(state: dict[str, torch.Tensor]) -> str:
return "attention" if any(key.startswith("base_encoder.") for key in state) else "mlp"
def load_model(checkpoint: dict[str, Any], cfg: Any, *, depth: int) -> torch.nn.Module:
probe = FastFloatingBeaconEnv(cfg, envs=1, seed=int(cfg.seed))
state = checkpoint["model"]
mode = planner_mode_for_state(state)
policy_arch = policy_arch_for_state(state)
if mode == "local-beacon":
model: torch.nn.Module = LocalBeaconPlannerController(
probe.obs_size,
hidden=int(cfg.hidden),
depth=int(depth),
planner_hidden=int(cfg.hidden),
planner_depth=2,
)
elif policy_arch == "attention":
model = TokenAttentionController(probe.obs_size, hidden=int(cfg.hidden), depth=int(depth))
else:
model = BeaconController(probe.obs_size, hidden=int(cfg.hidden), depth=int(depth))
model.to(resolve_device(cfg.device))
model.load_state_dict(state, strict=True)
model.eval()
return model
def main() -> None:
parser = argparse.ArgumentParser(description="Evaluate the active checkpoint and write replay artifacts.")
parser.add_argument("--checkpoint", default=ACTIVE_CHECKPOINT)
parser.add_argument("--run-name", default="active_eval_replay")
parser.add_argument("--run-dir", default="data/runs")
parser.add_argument("--device", default="cuda")
parser.add_argument("--eval-envs", type=int, default=8192)
parser.add_argument("--route-jumps", type=int, default=None)
parser.add_argument("--distractors", type=int, default=None)
parser.add_argument("--capture", type=int, default=8)
parser.add_argument("--depth", type=int, default=3)
args = parser.parse_args()
configure_motif_mapgen(**MAPGEN_KWARGS)
install_motif_mapgen()
path = Path(args.checkpoint)
device = resolve_device(str(args.device))
checkpoint = torch.load(path, map_location="cpu", weights_only=False)
cfg = normalize_config(checkpoint.get("stage_config") or checkpoint["config"])
cfg = replace(
cfg,
device=str(device),
eval_envs=int(args.eval_envs),
run_dir=str(args.run_dir),
route_jumps=int(args.route_jumps if args.route_jumps is not None else cfg.route_jumps),
distractors=int(args.distractors if args.distractors is not None else cfg.distractors),
)
model = load_model(checkpoint, cfg, depth=int(args.depth))
stats, frames, successes = evaluate(model, cfg, capture=int(args.capture))
maps, rollouts = replay_payload(frames, successes)
writer = V2RunWriter(str(args.run_dir), str(args.run_name), cfg, algorithm="fast_beacon_planner")
writer.write_replay(
1,
f"eval_g{cfg.route_jumps}_d{cfg.distractors}",
cfg,
maps,
rollouts,
stats,
meta_extra={
"controller": "egocentric goal-beacon runner PPO",
"planner": "none" if planner_mode_for_state(checkpoint["model"]) == "direct" else "learned local-beacon",
"teacher": "none",
"route_observation": "none",
"checkpoint": str(path),
},
)
print({"run": writer.id, "stats": stats, "path": str(writer.path)}, flush=True)
if __name__ == "__main__":
main()