Download training_code/launch.sbatch from SleepMastger/pusht-flashwam: direct link, hf CLI and curl.
- Browser
- Download file 6.37 kB
-
https://huggingface.co/SleepMastger/pusht-flashwam/resolve/main/training_code/launch.sbatch
- Command line
-
hf download hf://SleepMastger/pusht-flashwam/training_code/launch.sbatch
-
curl -L -o launch.sbatch https://huggingface.co/SleepMastger/pusht-flashwam/resolve/main/training_code/launch.sbatch
6.37 kB
| #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." | |