Instructions to use AlexWortega/tinyvla with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- LeRobot
How to use AlexWortega/tinyvla with LeRobot:
- Notebooks
- Google Colab
- Kaggle
File size: 6,145 Bytes
1601a2f | 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 | #!/usr/bin/env python
"""Stage 3: spatial distillation labels from a Qwen3-VL teacher (XS-VLA style).
For every Nth frame of each episode, ask the teacher for the bounding box of
the task-relevant object; quantize the box center onto a GRID x GRID map and
store the cell index. The student later learns a linear classifier over its
fast-path spatial tokens with CE x 0.15 on labeled frames.
Output: parquet with columns (dataset, episode_index, frame_index, cell,
cx, cy, confidence_ok) at ~/tinyvla_data/spatial_labels/<dataset>.parquet
Usage:
python scripts/label_spatial.py --teacher Qwen/Qwen3-VL-4B-Instruct \
--frame-stride 10 --episode-frac 0.4 [--datasets-limit 2] [--pilot 20]
"""
from __future__ import annotations
import argparse
import json
import re
from pathlib import Path
import torch
GRID = 32
OUT_DIR = Path.home() / "tinyvla_data" / "spatial_labels"
POINT_RE = re.compile(r"\[?\s*(\d+)\s*,\s*(\d+)\s*\]?")
PROMPT = (
"Task: {task}\n"
"Look at the image. Where is the single object the robot must interact with next "
"to accomplish this task? Answer with ONLY its center point as [x, y] in 0-1000 "
"normalized coordinates. If unsure, output [0,0]."
)
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--teacher", default="Qwen/Qwen3.5-4B")
parser.add_argument("--data-root", type=Path, default=Path.home() / "tinyvla_data/so101_v3")
parser.add_argument("--frame-stride", type=int, default=10)
parser.add_argument("--episode-frac", type=float, default=0.4)
parser.add_argument("--datasets-limit", type=int, default=None)
parser.add_argument("--pilot", type=int, default=None, help="label only N frames total, print results")
parser.add_argument("--batch-size", type=int, default=16)
args = parser.parse_args()
import pyarrow as pa
import pyarrow.parquet as pq
from lerobot.datasets.lerobot_dataset import LeRobotDataset
from transformers import AutoModelForImageTextToText, AutoProcessor
model = AutoModelForImageTextToText.from_pretrained(
args.teacher, dtype=torch.bfloat16, device_map="cuda"
)
proc = AutoProcessor.from_pretrained(args.teacher)
OUT_DIR.mkdir(parents=True, exist_ok=True)
roots = sorted(args.data_root.iterdir())
if args.datasets_limit:
roots = roots[: args.datasets_limit]
total_done = 0
for root in roots:
if not (root / "meta/info.json").exists():
continue
out_path = OUT_DIR / f"{root.name}.parquet"
if out_path.exists() and not args.pilot:
continue
ds = LeRobotDataset(root.name, root=root, video_backend="torchcodec")
image_key = sorted(k for k in ds.meta.features if k.startswith("observation.images"))[0]
n_eps = max(1, int(ds.num_episodes * args.episode_frac))
rows = []
pending = [] # (ep, fi, image_pil, task)
def flush():
nonlocal total_done
if not pending:
return
msgs = [
[{"role": "user", "content": [
{"type": "image", "image": img},
{"type": "text", "text": PROMPT.format(task=task)},
]}]
for _, _, img, task in pending
]
texts = [
proc.apply_chat_template(
m, tokenize=False, add_generation_prompt=True, enable_thinking=False
)
for m in msgs
]
images = [[p[2]] for p in pending]
inputs = proc(text=texts, images=images, return_tensors="pt", padding=True).to("cuda")
with torch.no_grad():
out = model.generate(**inputs, max_new_tokens=16, do_sample=False)
answers = proc.batch_decode(out[:, inputs["input_ids"].shape[1]:], skip_special_tokens=True)
for (ep, fi, _, _), ans in zip(pending, answers):
m = POINT_RE.search(ans)
ok = False
cell, cx, cy = -1, -1.0, -1.0
if m:
px, py = (int(g) for g in m.groups())
if 0 < px <= 1000 and 0 < py <= 1000:
cx, cy = px / 1000.0, py / 1000.0
gx, gy = min(GRID - 1, int(cx * GRID)), min(GRID - 1, int(cy * GRID))
cell = gy * GRID + gx
ok = True
rows.append({"dataset": root.name, "episode_index": ep, "frame_index": fi,
"cell": cell, "cx": cx, "cy": cy, "confidence_ok": ok})
if args.pilot:
print(f"ep{ep} f{fi}: '{ans.strip()[:60]}' -> cell {cell} ({cx:.2f},{cy:.2f})")
total_done += len(pending)
pending.clear()
from torchvision.transforms.functional import to_pil_image
import torch.nn.functional as F
for ep in range(n_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, args.frame_stride):
item = ds[idx]
# label on the same 256^2 view the student sees; ~8x fewer
# teacher vision tokens than full res
small = F.interpolate(
item[image_key][None].clamp(0, 1), size=(256, 256),
mode="bilinear", align_corners=False,
)[0]
img = to_pil_image(small)
pending.append((ep, idx - start, img, item.get("task") or ""))
if len(pending) >= args.batch_size:
flush()
if args.pilot and total_done + len(pending) >= args.pilot:
flush()
print(f"pilot done: {total_done} frames")
return
flush()
pq.write_table(pa.Table.from_pylist(rows), out_path)
ok_rate = sum(r["confidence_ok"] for r in rows) / max(len(rows), 1)
print(f"{root.name}: {len(rows)} labels -> {out_path} (ok {ok_rate:.1%})")
if __name__ == "__main__":
main()
|