File size: 7,095 Bytes
222aa51
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
150
151
152
153
#!/usr/bin/env bash
#SBATCH --job-name=pusht-fastwam-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
#
# ORIGINAL FastWAM (full 30/30 MoT, 6.0B, non-decoupled)
# 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.
#
# SUBMIT THIS WITH A DEPENDENCY ON THE FLASHWAM TWIN:
#   sbatch --dependency=afterany:<flashwam jobid> \
#       train_pusht_fastwam_scratch_n103.sbatch
#
# User decision 2026-09-01, after the first attempt proved concurrency does not
# work for this pair — an explicit exception to the usual no-chaining rule.
# History: jobs 2652 (flashwam) and 2653 (fastwam) were submitted independently
# and Slurm ran both at once, 4+4 GPUs on node 103. Job 2653 was killed at
# 16:44:55, ~2h09m in at step 1950/30150, with "DataLoader worker ... killed by
# signal: Terminated" — the earlyoom signature: node 103 SIGTERMs the
# biggest-RSS process when MemAvailable falls under ~10% (~200GB), and two
# dataloader-heavy 4-GPU jobs overflow its 2TB even with run_one.sh's malloc
# tuning. The surviving twin also went ~3x FASTER once its rival died
# (1.98 s/step contended -> 0.65 s/step alone), so serialising costs almost no
# wall clock: two serial runs at 1.5 step/s beat two contended ones.
#
# 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 --dependency=afterany:<flashwam jobid> train_pusht_fastwam_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_fastwam_scratch on $(hostname), port=29567"
NUM_GPUS=4 bash "${PKG}/run_one.sh" pusht_fastwam_scratch 29567 \
    "save_every=${SAVE_EVERY}"

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