File size: 6,372 Bytes
2622c40 | 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 | #!/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."
|