File size: 4,495 Bytes
5a2e445
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python
"""Offline eval for canonical-schema checkpoints (stage2+).

Reports action-chunk MSE (normalized space) on held-out episodes, the
per-timestep error curve, and stale-latent degradation.

Usage:
    python scripts/eval_offline.py \
        --checkpoint outputs/stage2_mixture/final \
        --repo-id VoicAndrei__so100_kitchen \
        --root ~/tinyvla_data/so101_v3/VoicAndrei__so100_kitchen \
        --episodes 8 --stale-s 0 1 2 [--embodiment-id 0] [--no-latent]
"""

from __future__ import annotations

import argparse
from pathlib import Path

import torch


@torch.no_grad()
def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--checkpoint", type=Path, required=True)
    parser.add_argument("--repo-id", required=True)
    parser.add_argument("--root", default=None)
    parser.add_argument("--episodes", type=int, default=8)
    parser.add_argument("--stride", type=int, default=30)
    parser.add_argument("--stale-s", type=float, nargs="*", default=[0.0, 1.0, 2.0])
    parser.add_argument("--embodiment-id", type=int, default=0)
    parser.add_argument("--split", choices=["first", "last"], default="last")
    args = parser.parse_args()

    from lerobot.datasets.lerobot_dataset import LeRobotDataset, LeRobotDatasetMetadata
    from transformers import AutoTokenizer
    from tinyvla.data.mixture import CanonicalSource
    from tinyvla.modeling_tinyvla import TinyVLAPolicy

    policy = TinyVLAPolicy.from_pretrained(args.checkpoint).cuda().eval()
    cfg = policy.config
    chunk = cfg.chunk_size
    tok = AutoTokenizer.from_pretrained(cfg.lm_model_name)

    meta = LeRobotDatasetMetadata(args.repo_id, root=args.root)
    ds = LeRobotDataset(
        args.repo_id,
        root=args.root,
        delta_timestamps={"action": [t / meta.fps for t in range(chunk)]},
        video_backend="torchcodec",
    )
    src = CanonicalSource(
        ds, args.embodiment_id, cfg.image_size, cfg.max_state_dim, cfg.max_action_dim
    )

    if args.split == "first":
        eps = list(range(args.episodes))
    else:
        eps = list(range(ds.num_episodes - args.episodes, ds.num_episodes))

    def to_batch(item):
        t = tok([item.pop("task")], padding=True, truncation=True,
                max_length=cfg.tokenizer_max_length, return_tensors="pt")
        b = {k: v[None].cuda() if torch.is_tensor(v) else v for k, v in item.items()}
        b["observation.language.tokens"] = t["input_ids"].cuda()
        b["observation.language.attention_mask"] = t["attention_mask"].bool().cuda()
        return b

    results = {s: [] for s in args.stale_s}
    per_t = torch.zeros(chunk)
    n_chunks = 0

    for ep in eps:
        start = int(ds.meta.episodes["dataset_from_index"][ep])
        end = int(ds.meta.episodes["dataset_to_index"][ep])
        for idx in range(start, end - 1, args.stride):
            if idx >= len(src):
                break
            item = src[idx - 0]
            gt = item["action"].clone()  # (chunk, A) normalized
            mask = item["action_dim_mask"].clone()
            pad = item.get("action_is_pad")
            batch = to_batch(dict(item))
            for stale_s in args.stale_s:
                b = dict(batch)
                if stale_s > 0:
                    stale_idx = max(start, idx - int(stale_s * ds.fps))
                    stale_item = src[stale_idx]
                    sb = to_batch(dict(stale_item))
                    b["semantic_latent"] = policy._semantic_latent(sb)
                pred = policy.predict_action_chunk(b)[0].cpu()  # (chunk, A) normalized
                err = (pred[:, mask] - gt[:, mask]) ** 2
                if pad is not None:
                    err = err[~pad]
                mse = err.mean().item()
                results[stale_s].append(mse)
                if stale_s == 0:
                    e = ((pred - gt) ** 2)[:, mask].mean(dim=-1)
                    if pad is not None:
                        e = e * (~pad).float()
                    per_t += e
                    n_chunks += 1

    print(f"\n=== {args.repo_id} | {len(eps)} {args.split} episodes | {n_chunks} chunks | ckpt {args.checkpoint} ===")
    for s, vals in results.items():
        print(f"stale {s:.0f}s: normalized chunk MSE {sum(vals)/len(vals):.4f}")
    curve = (per_t / max(n_chunks, 1)).sqrt()
    print("per-timestep normalized RMSE (t=0,10,25,49):",
          [round(curve[i].item(), 3) for i in (0, 10, 25, 49)])


if __name__ == "__main__":
    main()