a2c2 / scripts /eval.py
dennis96's picture
Upload folder using huggingface_hub
07f85c4 verified
Raw
History Blame Contribute Delete
4.86 kB
#!/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()