| |
| """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 ( |
| A2C2RandomSampleDataset, |
| discover_episode_pairs, |
| move_batch_to_device, |
| pick_device, |
| resolve_dataset_root, |
| split_episode_pairs, |
| ) |
| from model import A2C2CorrectionHead, A2C2CorrectionHeadConfig |
|
|
|
|
| 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() |
|
|