tinyvla / scripts /eval_offline.py
AlexWortega's picture
Upload folder using huggingface_hub
1601a2f verified
Raw
History Blame Contribute Delete
4.5 kB
#!/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()