#!/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."