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