tinyvla / tinyvla2 /scripts /fewshot_grid.py
AlexWortega's picture
Upload folder using huggingface_hub
5a2e445 verified
Raw
History Blame Contribute Delete
4.09 kB
#!/usr/bin/env python
"""Few-shot adaptation grid — the DECISIVE TinyVLA-2 experiment.
Zero-shot offline is confounded by A/B normalization scale (plan Q5). Few-shot is
fair: adapt each pretrained variant on N held-out-robot episodes from an equal
start, measure endpoint error on a disjoint test split. Hypothesis: C adapts with
the fewest examples (write a descriptor), then B (learn a fresh ID row), then A.
Protocol per (variant, N):
- resume the pretrained final; fast path (expert + projectors + conditioning) is
trainable, backbone frozen (as in pretraining). A/B get a fresh embodiment id
(3, an unused row); C uses jaco's morphology descriptor (no new parameter).
- FT on jaco episodes [0:N], 5k steps, lr 5e-5.
- eval endpoint error on jaco episodes [test_start:test_start+20] (disjoint).
Usage: python scripts/fewshot_grid.py (runs the whole grid, writes results json)
"""
from __future__ import annotations
import json
import subprocess
from pathlib import Path
ROOT = Path("/home/alexw/tinyvla")
PY = str(Path.home() / "tinyvla_venv/bin/python")
JACO = "heldout_jaco_play"
JACO_ROOT = str(Path.home() / "tinyvla_data/heldout/heldout_jaco_play")
TEST_START = 956 # jaco has 976 eps; eval on last 20, FT on first N (<= 936)
SIZES = [100, 500]
VARIANTS = {
"A": {"dir": "tv2_A_id_native", "cond": "id", "space": "native", "emb": 3, "mkey": None},
"B": {"dir": "tv2_B_id_canon", "cond": "id", "space": "canonical", "emb": 3, "mkey": None},
"C": {"dir": "tv2_C_morph_canon", "cond": "morph", "space": "canonical", "emb": 0, "mkey": "jaco"},
}
OUT = ROOT / "outputs" / "fewshot_results.json"
def write_config(v, spec, n):
cfg = ROOT / "outputs" / f"fewshot_{v}_{n}.yaml"
ds_extra = f" morph_key: {spec['mkey']}\n" if spec["mkey"] else ""
cfg.write_text(f"""output_dir: {ROOT}/outputs/fewshot_{v}_{n}
wandb: null
seed: 42
morphology_descriptors: {ROOT}/configs/morphology/descriptors.yaml
resume_from: {ROOT}/outputs/{spec['dir']}/final
resume_step: 0
datasets:
- repo_id: {JACO}
root: {JACO_ROOT}
episodes: {n}
weight: 1.0
embodiment_id: {spec['emb']}
{ds_extra}policy:
chunk_size: 50
n_action_steps: 50
image_size: 256
freeze_lm: true
freeze_vision_encoder: true
num_embodiments: 16
conditioning: {spec['cond']}
action_space: {spec['space']}
batch_size: 32
grad_accum: 1
num_workers: 8
lr: 5.0e-5
warmup_steps: 100
steps: 5000
grad_clip: 10.0
log_freq: 500
save_freq: 5000
mixed_precision: bf16
""")
return cfg
def run(cmd, log):
with open(log, "w") as f:
return subprocess.run(cmd, stdout=f, stderr=subprocess.STDOUT).returncode
def main():
results = {}
for v, spec in VARIANTS.items():
for n in SIZES:
tag = f"{v}_{n}"
cfg = write_config(v, spec, n)
ft_log = ROOT / "outputs" / f"fewshot_{tag}_ft.log"
print(f"[FT] {tag} ...", flush=True)
run([PY, str(ROOT / "scripts/train.py"), "--config", str(cfg)], ft_log)
ckpt = ROOT / "outputs" / f"fewshot_{tag}" / "final"
if not ckpt.exists():
results[tag] = {"error": "FT produced no final ckpt"}
continue
ev_log = ROOT / "outputs" / f"fewshot_{tag}_eval.log"
cmd = [PY, str(ROOT / "scripts/eval_canonical.py"),
"--checkpoint", str(ckpt), "--dataset", JACO, "--root", JACO_ROOT,
"--ep-start", str(TEST_START), "--episodes", "20", "--oracle-ids", "4"]
if spec["mkey"]:
cmd += ["--morph-key", spec["mkey"]]
print(f"[EVAL] {tag} ...", flush=True)
run(cmd, ev_log)
line = [l for l in ev_log.read_text().splitlines() if "BEST" in l]
results[tag] = {"eval": line[-1] if line else "no BEST line"}
print(f" {tag}: {results[tag]['eval']}", flush=True)
OUT.write_text(json.dumps(results, indent=1))
print("\n=== FEW-SHOT GRID DONE ===")
for k, v in results.items():
print(f"{k}: {v}")
if __name__ == "__main__":
main()