File size: 4,087 Bytes
5a2e445
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/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()