pusht-fastwam / training_code /launch.sbatch
SleepMastger's picture
add model card, conditioning, and training-time processing
222aa51 verified
Raw History Blame Contribute Delete
7.1 kB
#!/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."