File size: 8,275 Bytes
e22723b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
#!/usr/bin/env python
"""Run the GR00T (physical_ai_ft) checkpoint on held-out LeKiwi.

This is a deliberately unfair test, and the point is to measure HOW unfair.

The checkpoint is `conditioning: id` + `action_space: native`. Its action head was
trained on six GR00T families whose native layouts are 44/14/24/12/23/26-43 dims,
each addressed by an embodiment id in 10..15. LeKiwi (SO-100 arm on a 3-wheel kiwi
base, 9 dims: 6 joint targets in degrees + x_mm/y_mm/theta) was never assigned an
id and its dim k has no reason to mean what any GR00T robot's dim k means.

So the model is given every advantage available:
  - each trained id 10..15 is tried, and the BEST is reported (oracle id selection)
  - id 0 is reported too: rows 0..9 are the stale C-scaled ids this run never
    touched, so it is the "no matching identity" reference
  - the same 4-seed ODE averaging and the same per-dim ABS/DELTA baseline search
    used for the in-distribution evaluation

Preprocessing mirrors HubEpisodeStream: decimate to 10 Hz, per-dataset z-score,
zero-pad state to 256 and actions to 64.
"""

from __future__ import annotations

import sys
from pathlib import Path

import numpy as np
import torch

sys.path.insert(0, str(Path(__file__).parent))
from eval_physical_ai import bootstrap_ci, chunk_err, classify_dims, predict_mean  # noqa: E402

CKPT = "/home/alexw/tinyvla/outputs/physical_ai_ft_fixed"
NAME, ROOT = "lekiwi_cleanup", str(Path("~/tinyvla_data/lekiwi/lekiwi_cleanup").expanduser())
TARGET_HZ = 10.0
IDS = [10, 11, 12, 13, 14, 15, 0]


@torch.no_grad()
def main():
    from lerobot.datasets.lerobot_dataset import LeRobotDataset, LeRobotDatasetMetadata
    from safetensors.torch import load_file
    from transformers import AutoTokenizer

    from tinyvla.modeling_tinyvla import TinyVLAPolicy

    pol = TinyVLAPolicy.from_pretrained(CKPT)
    pol.load_state_dict(load_file(f"{CKPT}/model.safetensors"), strict=True)  # never strict=False
    cfg = pol.config
    pol = pol.cuda().eval()
    tok = AutoTokenizer.from_pretrained(cfg.lm_model_name)
    print(f"GATE  strict load OK | action_dim={pol.action_dim} "
          f"max_state={cfg.max_state_dim} max_action={cfg.max_action_dim}")

    meta = LeRobotDatasetMetadata(NAME, root=ROOT)
    stride = max(1, round(meta.fps / TARGET_HZ))
    ds = LeRobotDataset(NAME, root=ROOT, video_backend="torchcodec",
                        delta_timestamps={"action": [t / TARGET_HZ for t in range(cfg.chunk_size)]})
    sm = np.asarray(meta.stats["observation.state"]["mean"], dtype=np.float64)
    ss = np.maximum(np.asarray(meta.stats["observation.state"]["std"], dtype=np.float64), 1e-6)
    am = np.asarray(meta.stats["action"]["mean"], dtype=np.float64)
    as_ = np.maximum(np.asarray(meta.stats["action"]["std"], dtype=np.float64), 1e-6)
    ad = len(am)
    print(f"LeKiwi: {meta.total_episodes} eps, {meta.fps} fps -> stride {stride} ({TARGET_HZ} Hz), "
          f"action_dim {ad}")

    # per-dim ABS/DELTA classification, on RAW units, exactly as for the GR00T families
    hf = ds.reader.hf_dataset.with_format("numpy")
    s_all = np.asarray(hf["observation.state"], dtype=np.float64)[::stride]
    a_all = np.asarray(hf["action"], dtype=np.float64)[::stride]
    kind, amap, lag, std_a = classify_dims(s_all[:4000], a_all[:4000])
    names = meta.features["action"].get("names") or list(range(ad))
    names = names.get("motors", names) if isinstance(names, dict) else names
    print("dim classification:",
          ", ".join(f"{names[j]}={kind[j]}" + (f"<-s{amap[j]}" if kind[j] == "abs" else "")
                    for j in range(ad)))

    imk = sorted(k for k in meta.features if k.startswith("observation.images"))
    prim = next((k for k in imk if "wrist" not in k), imk[0])
    wrist = next((k for k in imk if "wrist" in k), None)

    def rs(x):
        return torch.nn.functional.interpolate(x[None].float(), size=(cfg.image_size, cfg.image_size),
                                               mode="bilinear", align_corners=False)[0]

    # held-out tail episodes
    test_eps = range(max(0, meta.total_episodes - 12), meta.total_episodes)
    samples = []
    for ep in test_eps:
        s = int(meta.episodes["dataset_from_index"][ep])
        e = int(meta.episodes["dataset_to_index"][ep])
        for idx in range(s, max(s + 1, e - cfg.chunk_size * stride), 90):
            item = ds[idx]
            A = item["action"].numpy().astype(np.float64)[:, :ad]
            if A.shape[0] < cfg.chunk_size:
                continue
            st_raw = item["observation.state"].numpy().astype(np.float64)[:ad]
            st_n = (st_raw - sm) / ss
            samples.append({
                "ep": ep,
                "cam0": rs(item[prim]),
                "cam1": rs(item[wrist]) if wrist else torch.zeros(3, cfg.image_size, cfg.image_size),
                "state": torch.nn.functional.pad(torch.tensor(st_n, dtype=torch.float32),
                                                 (0, cfg.max_state_dim - len(st_n))),
                "A": (A - am) / as_,
                "ref": np.array([(st_raw[amap[j]] - am[j]) / as_[j] if kind[j] == "abs" and amap[j] >= 0
                                 else 0.0 for j in range(ad)]),
                "task": item.get("task") or "",
            })
    print(f"held-out samples: {len(samples)} from {len(list(test_eps))} episodes\n")

    def run(emb_id):
        preds = []
        for k in range(0, len(samples), 8):
            chunk = samples[k:k + 8]
            t = tok([c["task"] for c in chunk], padding="max_length", truncation=True,
                    max_length=48, return_tensors="pt")
            b = {"observation.images.cam0": torch.stack([c["cam0"] for c in chunk]).cuda(),
                 "observation.images.cam1": torch.stack([c["cam1"] for c in chunk]).cuda(),
                 "observation.state": torch.stack([c["state"] for c in chunk]).cuda(),
                 "observation.language.tokens": t["input_ids"].cuda(),
                 "observation.language.attention_mask": t["attention_mask"].bool().cuda(),
                 "embodiment_id": torch.full((len(chunk),), emb_id, dtype=torch.long).cuda()}
            preds.append(predict_mean(pol, b, 4))
        return np.concatenate(preds, 0)

    print(f"{'emb id':>8}  {'subset':7}{'dims':>5}{'err':>9}{'B0':>9}{'B1':>9}{'B2':>9}{'ratio':>8}"
          f"   95% CI")
    best = {}
    for emb in IDS:
        pred = run(emb)
        for subset in ("abs", "delta"):
            sel = (kind == subset)
            if not sel.any():
                continue
            per_ep = {}
            for i, c in enumerate(samples):
                P = pred[i][:, :ad].astype(np.float64)
                e, _ = chunk_err(P, c["A"], kind, amap, c["ref"], sel)
                b0, _ = chunk_err(np.zeros_like(c["A"]), c["A"], kind, amap, c["ref"], sel)
                b1, _ = chunk_err(np.tile(c["ref"], (len(c["A"]), 1)), c["A"], kind, amap, c["ref"], sel)
                b2, _ = chunk_err(np.tile(c["A"][0], (len(c["A"]), 1)), c["A"], kind, amap, c["ref"], sel)
                per_ep.setdefault(c["ep"], []).append((e, b0, b1, b2))
            rows = np.array([np.mean(v, axis=0) for v in per_ep.values()])
            err, b0, b1, b2 = rows.mean(0)
            floor = min(b0, b1)
            ratios = rows[:, 0] / np.maximum(rows[:, 1:3].min(1), 1e-9)
            lo, hi = bootstrap_ci(ratios)
            tag = "  (untrained id row)" if emb == 0 else ""
            print(f"{emb:>8}  {subset:7}{int(sel.sum()):>5}{err:>9.3f}{b0:>9.3f}{b1:>9.3f}{b2:>9.3f}"
                  f"{err/floor:>8.2f}   [{lo:.2f}, {hi:.2f}]{tag}")
            k = subset
            if k not in best or err / floor < best[k][1]:
                best[k] = (emb, err / floor, b2 / floor)

    print("\nORACLE id selection (best of 10..15, i.e. maximally generous to the checkpoint):")
    for subset, (emb, r, b2r) in best.items():
        note = "  <- oracle repeat-A0 is BETTER, so the model loses to a dynamics-free predictor" \
            if b2r < r else ""
        print(f"  {subset:6} best id {emb}: ratio {r:.2f}   (oracle repeat-A0 {b2r:.2f}){note}")
    print("\nratio < 1.0 beats the trivial baseline (hold still for abs dims, zero motion for delta)")


if __name__ == "__main__":
    main()