summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--docs/campaign/FW135M_BP_BASELINE.md5
-rw-r--r--docs/campaign/FW135M_BP_HANDOFF.md35
-rw-r--r--ep_run/baseline_configs/fw135m_matched_bp.json10
-rw-r--r--ep_run/casc_bp_train.py48
-rw-r--r--ep_run/fw135m_baseline.py6
-rw-r--r--ep_run/prepare_fineweb.py5
-rw-r--r--ep_run/runs/fw135m_bp_sweep.sh48
-rw-r--r--sbatch/fw135m_full.sbatch54
-rw-r--r--sbatch/fw135m_lr_sweep.sbatch67
-rw-r--r--sbatch/fw135m_smoke.sbatch57
10 files changed, 313 insertions, 22 deletions
diff --git a/docs/campaign/FW135M_BP_BASELINE.md b/docs/campaign/FW135M_BP_BASELINE.md
index 3f81ba8..2aade2c 100644
--- a/docs/campaign/FW135M_BP_BASELINE.md
+++ b/docs/campaign/FW135M_BP_BASELINE.md
@@ -12,8 +12,8 @@ Status: commands configured on branch `xiang`; no training has started.
- Data: existing FineWeb-Edu 32k bins
- Batch: `B24`
- Target: `20N = 2,706,078,720` tokens
-- Complete-batch exposure: `2,706,081,792` tokens
-- Trainer argument: `--steps 440442` (440,443 inclusive updates)
+- EP-matched complete-batch exposure: `2,703,366,144` tokens
+- Trainer argument: `--steps 440000` (440,001 inclusive updates, matching `fw135m_bsign`)
## What changes from 72M
@@ -29,6 +29,7 @@ Status: commands configured on branch `xiang`; no training has started.
- FineWeb-Edu data and 32k tokenizer
- OLMo2-style model implementation
- Muon hybrid optimizer and cosine schedule
+- 1,000-step warmup, matching `fw135m_bsign`
- Weight decay 0.1 and BF16 autocast
- Seed list and validation cadence
- Existing `casc_bp_train.py` code path
diff --git a/docs/campaign/FW135M_BP_HANDOFF.md b/docs/campaign/FW135M_BP_HANDOFF.md
new file mode 100644
index 0000000..1467ac0
--- /dev/null
+++ b/docs/campaign/FW135M_BP_HANDOFF.md
@@ -0,0 +1,35 @@
+# FW135M BP Sweep Handoff
+
+The 135M BP smoke test completed on an RTX A6000:
+
+- `fw135m_bp_smoke_s1`
+- 400/400 steps completed; best validation CE `5.7126`
+- W&B: `eqprop-llm-training/ept-fineweb-135M`
+
+The BP sweep is matched to the active `fw135m_bsign` EP run in all shared
+settings: `L12/C768/H12/T256/B24`, FineWeb-Edu 32k, OLMo2, Muon
+(`--muon_lr 0.02`), BF16, weight decay `0.1`, cosine to `0.1×`, 440,000
+trainer steps, 1,000 warmup steps, and NCCL data parallelism. `B24` is per
+rank, so four A6000s use effective batch `96`. The intentional difference is
+BP versus EP's `--bsign_rand --beta 0.003`.
+
+## Run one candidate
+
+From `ep_run/`, each candidate should use all four A6000s:
+
+```bash
+CUDA_VISIBLE_DEVICES=0,1,2,3 GPUS=4 bash runs/fw135m_bp_sweep.sh 7e-4 1
+CUDA_VISIBLE_DEVICES=0,1,2,3 GPUS=4 bash runs/fw135m_bp_sweep.sh 1e-3 1
+CUDA_VISIBLE_DEVICES=0,1,2,3 GPUS=4 bash runs/fw135m_bp_sweep.sh 1.4e-3 1
+```
+
+Run the three commands sequentially when only four GPUs are available. Each
+candidate gets a distinct W&B run name:
+`fw135m_bp_lr7em4_s1`, `fw135m_bp_lr1em3_s1`, or `fw135m_bp_lr1p4em3_s1`.
+
+The launcher defaults data to `ep_run/data/fineweb_edu`. Set
+`EPT_DATA_ROOT=/path/to/data` only when FineWeb data is stored elsewhere.
+
+Choose the LR by best validation CE, final validation CE, and tail-median CE.
+Then run the selected BP setting for seeds 1 and 2 before reporting a BP/EP
+comparison.
diff --git a/ep_run/baseline_configs/fw135m_matched_bp.json b/ep_run/baseline_configs/fw135m_matched_bp.json
index 0e22686..98fa032 100644
--- a/ep_run/baseline_configs/fw135m_matched_bp.json
+++ b/ep_run/baseline_configs/fw135m_matched_bp.json
@@ -5,7 +5,7 @@
"width": "512 -> 768",
"heads": "8 -> 12",
"parameters": "72,114,688 -> 135,303,936",
- "training_updates": "440,443 updates to preserve approximately 20 tokens/parameter"
+ "training_updates": "440,001 updates to match the fw135m_bsign EP trainer argument"
},
"held_fixed_from_72m": {
"layers": 12,
@@ -19,7 +19,7 @@
"muon_lr": 0.02,
"adam_side_lr_center": 0.001,
"weight_decay": 0.1,
- "warmup_steps": 500,
+ "warmup_steps": 1000,
"schedule": "cosine to 0.1 of peak",
"precision": "bf16 autocast with fp32 parameters/states"
},
@@ -33,9 +33,9 @@
"batch": 24,
"parameters": 135303936,
"target_tokens_20N": 2706078720,
- "trainer_steps_argument": 440442,
- "actual_updates": 440443,
- "actual_tokens": 2706081792
+ "trainer_steps_argument": 440000,
+ "actual_updates": 440001,
+ "actual_tokens": 2703366144
},
"bp_lr_sweep_smallest_rung_only": [0.0007, 0.001, 0.0014],
"minimum_seeds_before_reporting": 2
diff --git a/ep_run/casc_bp_train.py b/ep_run/casc_bp_train.py
index 0515ec1..39a69d2 100644
--- a/ep_run/casc_bp_train.py
+++ b/ep_run/casc_bp_train.py
@@ -1,8 +1,9 @@
"""BP-train a small cascade-form standard transformer (L distinct blocks), saving ckpts
every --save_every for the A0.2 on-trajectory gradient gate (cascade_probe.py --ckpt).
Plain LLM training — this is also the BP twin for the C-tier money runs."""
-import argparse, math, pickle, time, json
+import argparse, math, os, pickle, time, json
import numpy as np, torch, torch.nn as nn, torch.nn.functional as F
+import torch.distributed as dist
from pathlib import Path
ap = argparse.ArgumentParser()
@@ -30,17 +31,28 @@ ap.add_argument('--zloss', type=float, default=0.0) # z-loss coefficient; 0
ap.add_argument('--qup_bits', type=int, default=0) # STAGE-0 mirror: naked resident-cell writes
ap.add_argument('--qcomp_bits', type=int, default=0) # STAGE-0 mirror: compute on DAC grid, fp32 master
ap.add_argument('--data', default='tinystories_bpe') # dataset dir under ep_run/data
+ap.add_argument('--ddp_backend', default='nccl', choices=['nccl', 'gloo'])
args = ap.parse_args()
if args.olmo2 and args.tok_init <= 0: args.tok_init = 0.02
-torch.manual_seed(args.seed)
+
+DDP = int(os.environ.get('WORLD_SIZE', '1')) > 1
+if DDP:
+ dist.init_process_group(args.ddp_backend)
+ RANK, WORLD = dist.get_rank(), dist.get_world_size()
+ torch.cuda.set_device(int(os.environ['LOCAL_RANK']) % max(torch.cuda.device_count(), 1))
+else:
+ RANK, WORLD = 0, 1
+torch.manual_seed(args.seed) # identical initialization on every rank
+DGEN = torch.Generator().manual_seed(args.seed * 7919 + RANK * 104729 + 11)
dev = 'cuda' if torch.cuda.is_available() else 'cpu'
-DD = Path('/home/yurenh2/ept/ep_run/data') / args.data
+DATA_ROOT = Path(os.environ.get('EPT_DATA_ROOT', Path(__file__).resolve().parent / 'data'))
+DD = DATA_ROOT / args.data
vocab = pickle.load(open(DD / 'meta.pkl', 'rb'))['vocab_size']
def get_batch(split):
data = np.memmap(DD / ('train.bin' if split == 'train' else 'val.bin'), dtype=np.uint16, mode='r')
- ix = torch.randint(len(data) - args.T - 1, (args.B,))
+ ix = torch.randint(len(data) - args.T - 1, (args.B,), generator=DGEN)
x = torch.stack([torch.from_numpy(data[i:i + args.T].astype(np.int64)) for i in ix])
y = torch.stack([torch.from_numpy(data[i + 1:i + 1 + args.T].astype(np.int64)) for i in ix])
return x.to(dev), y.to(dev)
@@ -148,7 +160,13 @@ if args.resume:
with torch.no_grad(): W_out.copy_(_ck['wout'].to(dev))
if _ck.get('lnf') is not None and not isinstance(ln_f, nn.Identity): ln_f.load_state_dict(_ck['lnf'])
start_step = int(_ck.get('step', 0))
- print(f'[resume] loaded {args.resume} at step {start_step}', flush=True)
+ if RANK == 0: print(f'[resume] loaded {args.resume} at step {start_step}', flush=True)
+if DDP:
+ with torch.no_grad():
+ for p in params:
+ dist.broadcast(p.data, 0)
+ if RANK == 0:
+ print(f'[ddp] world={WORLD} backend={args.ddp_backend}; effective batch {args.B}x{WORLD}={args.B * WORLD}', flush=True)
if args.opt == 'muon':
from muon import build_hybrid
opt, sched = build_hybrid(blocks, params, args.lr, args.muon_lr, args.warmup,
@@ -187,7 +205,7 @@ def evaluate(nb=6):
wb = None
if args.wandb == 'auto':
args.wandb = 'ept-fineweb-72m' if 'fineweb' in args.data else 'ept-tinystories-42m'
-if args.wandb:
+if args.wandb and RANK == 0:
try:
import wandb as _w
wb = _w.init(entity='eqprop-llm-training', project=args.wandb, name=args.wandb_run or args.tag, id=args.wandb_run or args.tag,
@@ -196,7 +214,8 @@ if args.wandb:
print(f'[wandb] disabled ({e})', flush=True)
n = sum(p.numel() for p in params)
-print(f'[{args.tag}] cascade-BP L{args.L} C{args.C} H{args.H} T{args.T} | {n/1e6:.2f}M params | {dev}', flush=True)
+if RANK == 0:
+ print(f'[{args.tag}] cascade-BP L{args.L} C{args.C} H{args.H} T{args.T} | {n/1e6:.2f}M params | {dev}', flush=True)
best, t0 = 1e9, time.time()
outdir = Path('runs'); outdir.mkdir(exist_ok=True)
for _ in range(start_step): sched.step() # advance LR schedule to the resumed step
@@ -216,6 +235,12 @@ for step in range(start_step, args.steps + 1):
if args.zloss > 0:
loss = loss + args.zloss * (torch.logsumexp(logits.float(), -1) ** 2).mean()
opt.zero_grad(set_to_none=True); loss.backward()
+ if DDP:
+ for p in params:
+ if p.grad is None:
+ p.grad = torch.zeros_like(p)
+ dist.all_reduce(p.grad, op=dist.ReduceOp.SUM)
+ p.grad.div_(WORLD)
if args.qcomp_bits > 0:
with torch.no_grad():
for p, q in zip(params, QSAVE): p.copy_(q)
@@ -231,19 +256,22 @@ for step in range(start_step, args.steps + 1):
q = p / g_
fl = q.floor()
p.copy_((fl + (torch.rand_like(p) < (q - fl)).float()) * g_)
- if step % args.log == 0:
+ if step % args.log == 0 and RANK == 0:
val = evaluate(); best = min(best, val)
print(f'step {step:5d}/{args.steps} | train {loss.item():.4f} val {val:.4f} (best {best:.4f}) '
f'| {step/max(time.time()-t0,1e-9):.2f} it/s', flush=True)
if wb is not None:
try: wb.log({'train_ce': loss.item(), 'val_ce': val, 'best': best}, step=step)
except Exception: pass
- if step % args.save_every == 0 or step == args.steps:
+ if (step % args.save_every == 0 or step == args.steps) and RANK == 0:
torch.save({'tok': tok.state_dict(), 'pos': pos.state_dict(), 'blocks': blocks.state_dict(),
'wout': (W_out.detach().cpu() if args.olmo2 else None),
'lnf': (ln_f.state_dict() if not isinstance(ln_f, nn.Identity) else None),
'step': step, 'val': best, 'config': vars(args)}, outdir / f'{args.tag}_s{step}.pt')
-print(f'[{args.tag}] DONE best val CE {best:.4f} (random ln({vocab})={math.log(vocab):.3f})', flush=True)
+if RANK == 0:
+ print(f'[{args.tag}] DONE best val CE {best:.4f} (random ln({vocab})={math.log(vocab):.3f})', flush=True)
if wb is not None:
try: wb.summary['best_val_ce'] = best; wb.finish()
except Exception: pass
+if DDP:
+ dist.destroy_process_group()
diff --git a/ep_run/fw135m_baseline.py b/ep_run/fw135m_baseline.py
index ddcc5c0..8daf145 100644
--- a/ep_run/fw135m_baseline.py
+++ b/ep_run/fw135m_baseline.py
@@ -13,7 +13,7 @@ HEADS = 12
CONTEXT = 256
BATCH = 24
PARAMETERS = 135303936
-FULL_STEPS_ARG = 440442 # casc_bp_train.py loops inclusively: 440,443 updates
+FULL_STEPS_ARG = 440000 # Match fw135m_bsign; casc_bp_train.py loops inclusively: 440,001 updates.
LR_SWEEP = ("7e-4", "1e-3", "1.4e-3")
@@ -94,12 +94,12 @@ def main():
f"fw135m_bp_lr{lr.replace('-', 'm').replace('.', 'p')}_s{args.seed}",
lr,
args.sweep_steps,
- 500,
+ 1000,
)
for lr in LR_SWEEP
]
else:
- runs = [(f"fw135m_bp_s{args.seed}", args.lr, FULL_STEPS_ARG, 500)]
+ runs = [(f"fw135m_bp_s{args.seed}", args.lr, FULL_STEPS_ARG, 1000)]
print(
f"# L{LAYERS} C{WIDTH} H{HEADS} T{CONTEXT} B{BATCH} "
diff --git a/ep_run/prepare_fineweb.py b/ep_run/prepare_fineweb.py
index 33fce59..afd2e45 100644
--- a/ep_run/prepare_fineweb.py
+++ b/ep_run/prepare_fineweb.py
@@ -11,7 +11,7 @@ Phases (all resumable-ish, markers for the watcher):
Docs are joined with a <|eot|> separator (id 0). vocab 32768 fits uint16.
NFS note: peak disk = raw parquet ~28GB + bins ~20GB; keep raw/ for tokenizer reruns.
"""
-import pickle, time
+import os, pickle, time
from pathlib import Path
import numpy as np
import pyarrow.parquet as pq
@@ -22,7 +22,8 @@ from tokenizers.trainers import BpeTrainer
from tokenizers.pre_tokenizers import ByteLevel
from tokenizers.decoders import ByteLevel as ByteLevelDec
-D = Path('/home/yurenh2/ept/ep_run/data/fineweb_edu')
+DATA_ROOT = Path(os.environ.get('EPT_DATA_ROOT', Path(__file__).resolve().parent / 'data'))
+D = DATA_ROOT / 'fineweb_edu'
RAW = D / 'raw'
D.mkdir(parents=True, exist_ok=True)
VOCAB = 32768
diff --git a/ep_run/runs/fw135m_bp_sweep.sh b/ep_run/runs/fw135m_bp_sweep.sh
new file mode 100644
index 0000000..e9715b4
--- /dev/null
+++ b/ep_run/runs/fw135m_bp_sweep.sh
@@ -0,0 +1,48 @@
+#!/usr/bin/env bash
+# Run one EP-matched 135M BP LR-sweep candidate.
+#
+# Usage:
+# ./runs/fw135m_bp_sweep.sh <7e-4|1e-3|1.4e-3> [seed]
+# Set GPUS=4 to launch a four-A6000 NCCL data-parallel run with torchrun.
+
+set -euo pipefail
+
+LR="${1:?Usage: $0 <7e-4|1e-3|1.4e-3> [seed]}"
+SEED="${2:-1}"
+case "${LR}" in
+ 7e-4|1e-3|1.4e-3) ;;
+ *) echo "Unsupported LR: ${LR}" >&2; exit 2 ;;
+esac
+
+RUN_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
+PYTHON_BIN="${PYTHON_BIN:-python}"
+WAND_PROJECT="${WAND_PROJECT:-ept-fineweb-135M}"
+GPUS="${GPUS:-1}"
+export EPT_DATA_ROOT="${EPT_DATA_ROOT:-${RUN_DIR}/data}"
+
+LR_TAG="${LR//-/m}"
+LR_TAG="${LR_TAG//./p}"
+TAG="fw135m_bp_lr${LR_TAG}_s${SEED}"
+
+cd "${RUN_DIR}"
+if (( GPUS > 1 )); then
+ LAUNCH=(torchrun --standalone --nproc_per_node="${GPUS}")
+else
+ LAUNCH=("${PYTHON_BIN}")
+fi
+
+exec "${LAUNCH[@]}" casc_bp_train.py \
+ --tag "${TAG}" \
+ --L 12 --C 768 --H 12 --T 256 --B 24 \
+ --steps 440000 \
+ --lr "${LR}" \
+ --warmup 1000 \
+ --amp --olmo2 \
+ --wd 0.1 \
+ --opt muon --muon_lr 0.02 \
+ --cosine --lr_min_ratio 0.1 \
+ --data fineweb_edu \
+ --seed "${SEED}" \
+ --save_every 5000 --log 100 \
+ --wandb "${WAND_PROJECT}" \
+ --wandb_run "${TAG}"
diff --git a/sbatch/fw135m_full.sbatch b/sbatch/fw135m_full.sbatch
new file mode 100644
index 0000000..f0a8b79
--- /dev/null
+++ b/sbatch/fw135m_full.sbatch
@@ -0,0 +1,54 @@
+#!/bin/bash
+#SBATCH --job-name=fw135m-full
+#SBATCH --output=/orion/u/oscarwan/ept/ep_run/runs/slurm-fw135m-full-%A_%a.out
+#SBATCH --error=/orion/u/oscarwan/ept/ep_run/runs/slurm-fw135m-full-%A_%a.err
+#SBATCH --time=72:00:00
+#SBATCH --account=orion
+#SBATCH --partition=orion
+#SBATCH --nodes=1
+#SBATCH --ntasks-per-node=1
+#SBATCH --gres=gpu:a6000:1
+#SBATCH --cpus-per-task=16
+#SBATCH --mem=64G
+#SBATCH --array=1-2
+
+set -euo pipefail
+
+ROOT_DIR="${ROOT_DIR:-/orion/u/oscarwan/ept}"
+RUN_DIR="${ROOT_DIR}/ep_run"
+VENV_DIR="${VENV_DIR:-${ROOT_DIR}/.venv}"
+PYTHON_BIN="${PYTHON_BIN:-python}"
+WAND_PROJECT="ept-fineweb-135M"
+WAND_ENTITY="eqprop-llm-training"
+DATA_DIR="${DATA_DIR:-${RUN_DIR}/data/fineweb_edu}"
+export EPT_DATA_ROOT="${EPT_DATA_ROOT:-${DATA_DIR%/fineweb_edu}}"
+SELECTED_LR="${SELECTED_LR:?Submit with --export=ALL,SELECTED_LR=<winning-LR>}"
+SEED="${SLURM_ARRAY_TASK_ID:?This script must be submitted as an array job}"
+
+cd "${RUN_DIR}"
+mkdir -p runs
+
+if [[ -f "${VENV_DIR}/bin/activate" ]]; then
+ source "${VENV_DIR}/bin/activate"
+fi
+
+if [[ ! -f "${DATA_DIR}/meta.pkl" || ! -f "${DATA_DIR}/train.bin" || ! -f "${DATA_DIR}/val.bin" ]]; then
+ echo "FineWeb-Edu data is missing at ${DATA_DIR}." >&2
+ echo "Run prepare_fineweb.py first to create the dataset." >&2
+ exit 1
+fi
+
+echo "Host: $(hostname)"
+echo "Job ID: ${SLURM_JOB_ID:-local}; seed: ${SEED}"
+echo "Git commit: $(git rev-parse --short HEAD)"
+echo "CUDA_VISIBLE_DEVICES=${CUDA_VISIBLE_DEVICES:-unset}"
+echo "W&B: ${WAND_ENTITY}/${WAND_PROJECT}"
+echo "Selected Adam-side LR: ${SELECTED_LR}; data: ${DATA_DIR}"
+nvidia-smi
+
+"${PYTHON_BIN}" fw135m_baseline.py \
+ --mode full \
+ --lr "${SELECTED_LR}" \
+ --seed "${SEED}" \
+ --wandb_project "${WAND_PROJECT}" \
+ --execute
diff --git a/sbatch/fw135m_lr_sweep.sbatch b/sbatch/fw135m_lr_sweep.sbatch
new file mode 100644
index 0000000..bdb403f
--- /dev/null
+++ b/sbatch/fw135m_lr_sweep.sbatch
@@ -0,0 +1,67 @@
+#!/bin/bash
+#SBATCH --job-name=fw135m-lr
+#SBATCH --output=/orion/u/oscarwan/ept/ep_run/runs/slurm-fw135m-lr-%A_%a.out
+#SBATCH --error=/orion/u/oscarwan/ept/ep_run/runs/slurm-fw135m-lr-%A_%a.err
+#SBATCH --time=72:00:00
+#SBATCH --account=orion
+#SBATCH --partition=orion
+#SBATCH --nodes=1
+#SBATCH --ntasks-per-node=1
+#SBATCH --gres=gpu:a6000:1
+#SBATCH --cpus-per-task=16
+#SBATCH --mem=64G
+#SBATCH --array=0-2
+
+set -euo pipefail
+
+ROOT_DIR="${ROOT_DIR:-/orion/u/oscarwan/ept}"
+RUN_DIR="${ROOT_DIR}/ep_run"
+VENV_DIR="${VENV_DIR:-${ROOT_DIR}/.venv}"
+PYTHON_BIN="${PYTHON_BIN:-python}"
+WAND_PROJECT="ept-fineweb-135M"
+WAND_ENTITY="eqprop-llm-training"
+DATA_DIR="${DATA_DIR:-${RUN_DIR}/data/fineweb_edu}"
+export EPT_DATA_ROOT="${EPT_DATA_ROOT:-${DATA_DIR%/fineweb_edu}}"
+SEED="${SEED:-1}"
+LRS=(7e-4 1e-3 1.4e-3)
+LR="${LRS[${SLURM_ARRAY_TASK_ID:?This script must be submitted as an array job}]}"
+LR_TAG="${LR//-/m}"
+LR_TAG="${LR_TAG//./p}"
+TAG="fw135m_bp_lr${LR_TAG}_s${SEED}"
+
+cd "${RUN_DIR}"
+mkdir -p runs
+
+if [[ -f "${VENV_DIR}/bin/activate" ]]; then
+ source "${VENV_DIR}/bin/activate"
+fi
+
+if [[ ! -f "${DATA_DIR}/meta.pkl" || ! -f "${DATA_DIR}/train.bin" || ! -f "${DATA_DIR}/val.bin" ]]; then
+ echo "FineWeb-Edu data is missing at ${DATA_DIR}." >&2
+ echo "Run prepare_fineweb.py first to create the dataset." >&2
+ exit 1
+fi
+
+echo "Host: $(hostname)"
+echo "Job ID: ${SLURM_JOB_ID:-local}; array task: ${SLURM_ARRAY_TASK_ID}"
+echo "Git commit: $(git rev-parse --short HEAD)"
+echo "CUDA_VISIBLE_DEVICES=${CUDA_VISIBLE_DEVICES:-unset}"
+echo "W&B: ${WAND_ENTITY}/${WAND_PROJECT}; run: ${TAG}"
+echo "Adam-side LR: ${LR}; data: ${DATA_DIR}"
+nvidia-smi
+
+"${PYTHON_BIN}" casc_bp_train.py \
+ --tag "${TAG}" \
+ --L 12 --C 768 --H 12 --T 256 --B 24 \
+ --steps 440000 \
+ --lr "${LR}" \
+ --warmup 1000 \
+ --amp --olmo2 \
+ --wd 0.1 \
+ --opt muon --muon_lr 0.02 \
+ --cosine --lr_min_ratio 0.1 \
+ --data fineweb_edu \
+ --seed "${SEED}" \
+ --save_every 5000 --log 100 \
+ --wandb "${WAND_PROJECT}" \
+ --wandb_run "${TAG}"
diff --git a/sbatch/fw135m_smoke.sbatch b/sbatch/fw135m_smoke.sbatch
new file mode 100644
index 0000000..8fc366f
--- /dev/null
+++ b/sbatch/fw135m_smoke.sbatch
@@ -0,0 +1,57 @@
+#!/bin/bash
+#SBATCH --job-name=fw135m-smoke
+#SBATCH --output=/orion/u/oscarwan/ept/ep_run/runs/slurm-fw135m-smoke-%j.out
+#SBATCH --error=/orion/u/oscarwan/ept/ep_run/runs/slurm-fw135m-smoke-%j.err
+#SBATCH --time=02:00:00
+#SBATCH --account=orion
+#SBATCH --partition=orion
+#SBATCH --nodes=1
+#SBATCH --ntasks-per-node=1
+#SBATCH --gres=gpu:a6000:1
+#SBATCH --cpus-per-task=16
+#SBATCH --mem=64G
+
+set -euo pipefail
+
+ROOT_DIR="${ROOT_DIR:-/orion/u/oscarwan/ept}"
+RUN_DIR="${ROOT_DIR}/ep_run"
+VENV_DIR="${VENV_DIR:-${ROOT_DIR}/.venv}"
+PYTHON_BIN="${PYTHON_BIN:-python}"
+WAND_PROJECT="ept-fineweb-135M"
+WAND_ENTITY="eqprop-llm-training"
+DATA_DIR="${DATA_DIR:-${RUN_DIR}/data/fineweb_edu}"
+export EPT_DATA_ROOT="${EPT_DATA_ROOT:-${DATA_DIR%/fineweb_edu}}"
+SEED="${SEED:-1}"
+
+cd "${RUN_DIR}"
+mkdir -p runs
+
+if [[ -f "${VENV_DIR}/bin/activate" ]]; then
+ source "${VENV_DIR}/bin/activate"
+fi
+
+if [[ ! -f "${DATA_DIR}/meta.pkl" || ! -f "${DATA_DIR}/train.bin" || ! -f "${DATA_DIR}/val.bin" ]]; then
+ echo "FineWeb-Edu data is missing at ${DATA_DIR}." >&2
+ echo "Run prepare_fineweb.py first to create the dataset." >&2
+ exit 1
+fi
+
+echo "Host: $(hostname)"
+echo "Job ID: ${SLURM_JOB_ID:-local}"
+echo "Git commit: $(git rev-parse --short HEAD)"
+echo "CUDA_VISIBLE_DEVICES=${CUDA_VISIBLE_DEVICES:-unset}"
+echo "W&B: ${WAND_ENTITY}/${WAND_PROJECT}"
+echo "Data: ${DATA_DIR}"
+nvidia-smi
+
+"${PYTHON_BIN}" - <<'PY'
+import torch, wandb
+print("torch", torch.__version__, "cuda_available", torch.cuda.is_available(), "gpu_count", torch.cuda.device_count())
+print("wandb", wandb.__version__)
+PY
+
+"${PYTHON_BIN}" fw135m_baseline.py \
+ --mode smoke \
+ --seed "${SEED}" \
+ --wandb_project "${WAND_PROJECT}" \
+ --execute