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."
|