File size: 6,372 Bytes
2622c40
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env bash
#SBATCH --job-name=pusht-flashwam-scratch
#SBATCH --partition=gpu
#SBATCH --nodelist=gpu-h200-103
#SBATCH --gres=gpu:4
#SBATCH --cpus-per-task=64
#SBATCH --mem=500G
#SBATCH --time=48:00:00
#SBATCH --chdir=/shared_work/george/real_world_wam/pusht_train
#SBATCH --output=/shared_work/george/real_world_wam/pusht_train/slurm_logs/%x-%j.out
#
# FlashWAM (M1_FusedKV_RopeFixed) — decoupled MoT, 1-layer action expert
# FROM SCRATCH on the pusht dataset (100 real-Franka Push-T demos, 32,131
# frames @ 10 Hz; HF SleepMastger/pusht-manipulation).
# Recipe matches the dish_utensil runs: 4 GPU x batch 8 x accum 1 = global 32,
# 30 epochs -> 1,005 steps/epoch, 30,150 steps, 6 checkpoints.
#
# Submitted INDEPENDENTLY of its twin — no --dependency, per the user's
# standing preference; Slurm decides whether they overlap. If they DO land
# concurrently (8 GPUs = all of node 103), the malloc tuning in run_one.sh is
# what keeps host RAM under earlyoom's ~200GB-free trigger. Note gpu-h200-103's
# gres accounting has double-booked GPUs before — treat "free GPUs" on 103 as
# unreliable while other users have jobs pending.
#
# FAILURE POLICY: killed by cluster manager => do NOT resubmit; own-error
# failure => fix root cause, resubmit once.
#
# Usage: sbatch train_pusht_flashwam_scratch_n103.sbatch
set -euo pipefail

RW=/shared_work/george/real_world_wam
PKG="${RW}/pusht_train"
DATASET_SRC="${RW}/datasets/pusht_lerobot_v21"
CACHE_SRC="${PKG}/text_embeds_cache"
PY=/shared_work/environments/miniconda3/envs/fastwam-robotwin-rui/bin/python

echo "host=$(hostname) date=$(date -u +%FT%TZ) CUDA_VISIBLE_DEVICES=${CUDA_VISIBLE_DEVICES:-unset} SLURM_JOB_ID=${SLURM_JOB_ID:-none}"
nvidia-smi --query-gpu=index,memory.used,utilization.gpu --format=csv

if [ ! -d "${DATASET_SRC}" ]; then
    echo "ERROR: dataset not found at ${DATASET_SRC}." >&2
    exit 1
fi

# --- fail fast on the things that silently corrupt a 12h run -----------------
N_EPS=$("${PY}" -c "import json; print(json.load(open('${DATASET_SRC}/meta/info.json'))['total_episodes'])")
if [ "${N_EPS}" -ne 100 ]; then
    echo "ERROR: ${DATASET_SRC} has total_episodes=${N_EPS}, expected 100." >&2
    exit 1
fi

FPS=$("${PY}" -c "import json; print(json.load(open('${DATASET_SRC}/meta/info.json'))['fps'])")
if [ "${FPS}" -ne 10 ]; then
    echo "ERROR: ${DATASET_SRC} has fps=${FPS}, expected 10 (pusht was recorded at 10 Hz)." >&2
    exit 1
fi

# Exact-hash text-cache check against THIS dataset's tasks.jsonl. The
# dataloader hard-fails on a cache miss, so catch it here rather than 20 min in.
"${PY}" - "${DATASET_SRC}" "${CACHE_SRC}" <<'EOF'
import hashlib, json, pathlib, sys
dataset, cache = sys.argv[1], sys.argv[2]
task = json.loads(pathlib.Path(dataset, "meta/tasks.jsonl").read_text().splitlines()[0])["task"]
prompt = f"A video recorded from a robot's point of view executing the following instruction: {task}"
f = pathlib.Path(cache) / f"{hashlib.sha256(prompt.encode()).hexdigest()}.t5_len128.wan22ti2v5b.pt"
if not f.exists():
    sys.exit(f"ERROR: missing text-embed cache file {f}\n       task string: {task!r}\n"
             "       Run precompute_pusht_text_embeds.sh with this exact string.")
print(f"text-embed cache OK: {f.name} (task {task!r})")
EOF

# Assert the normalizer's degenerate-channel guard actually covers this
# dataset's 6 constant channels (action drx/dry/drz/gripper, state
# gripL/gripR). SingleFieldLinearNormalizer maps a channel to a constant 0
# when its range < range_tol=1e-4, and to inf if it did not. The tightest
# constant channel here (gripper state, 3.2e-05) clears range_tol by only ~3x,
# so verify it rather than assume it — a silent inf would poison the run.
"${PY}" - "${DATASET_SRC}" <<'EOF'
import json, pathlib, sys
root = pathlib.Path(sys.argv[1])
path = root / "meta" / "episodes_stats.jsonl"
if not path.exists():
    sys.exit(f"ERROR: {path} missing; cannot verify normalization ranges.")
TOL = 1e-4
# range_tol classifies on the DATASET-wide min/max, so aggregate across episodes.
agg = {}
with path.open() as f:
    for line in f:
        line = line.strip()
        if not line:
            continue
        st = json.loads(line)["stats"]
        for key in ("action", "observation.state"):
            if key not in st:
                continue
            lo, hi = st[key]["min"], st[key]["max"]
            if key not in agg:
                agg[key] = [list(lo), list(hi)]
            else:
                a = agg[key]
                a[0] = [min(x, y) for x, y in zip(a[0], lo)]
                a[1] = [max(x, y) for x, y in zip(a[1], hi)]
if not agg:
    sys.exit("ERROR: no action/observation.state stats found.")
EXPECT_DEGENERATE = {"action": [3, 4, 5, 6], "observation.state": [6, 7]}
bad = []
for key, (lo, hi) in sorted(agg.items()):
    rng = [h - l for l, h in zip(lo, hi)]
    deg = [i for i, r in enumerate(rng) if r < TOL]
    live = [i for i, r in enumerate(rng) if r >= TOL]
    print(f"{key}: degenerate dims {deg} (-> normalized to 0), live dims {live}")
    want = EXPECT_DEGENERATE[key]
    if deg != want:
        bad.append(f"{key}: degenerate dims {deg}, expected {want} "
                   f"(ranges: {['%.2e' % r for r in rng]})")
    for i in deg:
        if rng[i] > TOL * 0.5:
            bad.append(f"{key}[{i}] range {rng[i]:.2e} is within 2x of range_tol={TOL}")
if bad:
    sys.exit("ERROR: degenerate-channel contract violated:\n  " + "\n  ".join(bad))
print("degenerate-channel guard OK (4 action + 2 state dims safely under range_tol)")
EOF

bash "${PKG}/stage_pusht_local.sh"

TOTAL_FRAMES=$("${PY}" -c "import json; print(json.load(open('${DATASET_SRC}/meta/info.json'))['total_frames'])")
STEPS_PER_EPOCH=$(( (TOTAL_FRAMES + 31) / 32 ))
SAVE_EVERY=$(( STEPS_PER_EPOCH * 5 ))
echo "[$(date)] pusht: total_frames=${TOTAL_FRAMES}, steps/epoch=${STEPS_PER_EPOCH}, save_every=${SAVE_EVERY}"

# decode-fix: bump dataloader MAX_GETITEM_ATTEMPT 5 -> 500 (pyav decode crash,
# "Failed to load valid sample after 5 attempts").
export PYTHONPATH="${RW}/place_cube_train/decode_fix:${PYTHONPATH:-}"

# NB: do NOT set PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True (nan loss).
echo "[$(date)] ==> training pusht_flashwam_scratch on $(hostname), port=29566"
NUM_GPUS=4 bash "${PKG}/run_one.sh" pusht_flashwam_scratch 29566 \
    "save_every=${SAVE_EVERY}"

echo "[$(date)] ==> done."