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