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