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()