File size: 4,858 Bytes
07f85c4 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 | #!/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()
|