tinyvla / tinyvla2 /scripts /eval_lekiwi_groot.py
AlexWortega's picture
Upload tinyvla2/scripts/eval_lekiwi_groot.py with huggingface_hub
e22723b verified
Raw
History Blame Contribute Delete
8.28 kB
#!/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()