#!/usr/bin/env python3 """Evaluate an A2C2 correction head checkpoint on cached parquet data.""" from __future__ import annotations import argparse from dataclasses import fields from pathlib import Path import sys import torch import torch.nn.functional as F from torch.utils.data import DataLoader SCRIPT_ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(SCRIPT_ROOT / "src")) from dataset import ( # noqa: E402 A2C2RandomSampleDataset, discover_episode_pairs, move_batch_to_device, pick_device, resolve_dataset_root, split_episode_pairs, ) from model import A2C2CorrectionHead, A2C2CorrectionHeadConfig # noqa: E402 def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--dataset-root", type=Path, required=True) parser.add_argument("--checkpoint", type=Path, required=True) parser.add_argument("--task-dir", default=None, help="Optional task directory filter, e.g. task-0018.") parser.add_argument("--split", choices=("train", "val", "all"), default="val") parser.add_argument("--val-ratio", type=float, default=0.05) parser.add_argument("--max-episodes", type=int, default=None) parser.add_argument("--num-samples", type=int, default=10_000) parser.add_argument("--batch-size", type=int, default=256) parser.add_argument("--num-workers", type=int, default=0) parser.add_argument("--samples-per-episode", type=int, default=512) parser.add_argument("--seed", type=int, default=123) parser.add_argument("--device", default="auto") return parser.parse_args() def config_from_checkpoint(payload: dict) -> A2C2CorrectionHeadConfig: raw = payload.get("config", {}) valid_keys = {field.name for field in fields(A2C2CorrectionHeadConfig)} filtered = {key: value for key, value in raw.items() if key in valid_keys} return A2C2CorrectionHeadConfig(**filtered) def main() -> None: args = parse_args() device = pick_device(args.device) payload = torch.load(args.checkpoint.expanduser(), map_location=device) cfg = config_from_checkpoint(payload) model = A2C2CorrectionHead(cfg).to(device) model.load_state_dict(payload["model_state_dict"]) model.eval() dataset_root = resolve_dataset_root(args.dataset_root) pairs = discover_episode_pairs(dataset_root, args.task_dir) train_pairs, val_pairs = split_episode_pairs(pairs, args.val_ratio, args.seed, args.max_episodes) if args.split == "train": eval_pairs = train_pairs elif args.split == "val": eval_pairs = val_pairs if val_pairs else train_pairs else: eval_pairs = train_pairs + val_pairs dataset = A2C2RandomSampleDataset( eval_pairs, action_horizon=cfg.action_horizon, samples_per_episode=args.samples_per_episode, seed=args.seed, total_samples=args.num_samples, ) loader = DataLoader( dataset, batch_size=args.batch_size, num_workers=args.num_workers, pin_memory=device.type == "cuda", ) total = 0 residual_mse_sum = 0.0 residual_mae_sum = 0.0 corrected_mse_sum = 0.0 base_mse_sum = 0.0 with torch.no_grad(): for batch in loader: batch = move_batch_to_device(batch, device) pred_delta = model( batch["observation_state"], batch["base_action"], batch["base_action_chunk"], batch["base_policy_z"], batch["time_feature"], batch["valid_action_mask"], ) target_delta = batch["target_delta"] base_action = batch["base_action"] expert_action = batch["expert_action"] corrected_action = base_action + pred_delta batch_size = target_delta.shape[0] total += batch_size residual_mse_sum += F.mse_loss(pred_delta, target_delta, reduction="sum").item() residual_mae_sum += F.l1_loss(pred_delta, target_delta, reduction="sum").item() corrected_mse_sum += F.mse_loss(corrected_action, expert_action, reduction="sum").item() base_mse_sum += F.mse_loss(base_action, expert_action, reduction="sum").item() if total >= args.num_samples: break denom = max(total * cfg.action_dim, 1) print(f"dataset_root: {dataset_root}") print(f"checkpoint: {args.checkpoint}") print(f"split: {args.split}") print(f"episodes: {len(eval_pairs)}") print(f"samples: {total}") print(f"residual_mse: {residual_mse_sum / denom:.8f}") print(f"residual_mae: {residual_mae_sum / denom:.8f}") print(f"corrected_action_mse: {corrected_mse_sum / denom:.8f}") print(f"base_action_mse: {base_mse_sum / denom:.8f}") if __name__ == "__main__": main()