summaryrefslogtreecommitdiff
path: root/ep_run
diff options
context:
space:
mode:
authorOscar Wan <oscarwan@stanford.edu>2026-07-21 10:48:52 -0700
committerOscar Wan <oscarwan@stanford.edu>2026-07-21 10:48:52 -0700
commitc4a6b459cdd5b183220cf62996f61efc8b56620d (patch)
tree78da709a17b24f274ff5f89ad7faf34bacefce4c /ep_run
parentca209ee5df16174d2137b5d5922581fc2f6d7e88 (diff)
add 135M baseline
Diffstat (limited to 'ep_run')
-rw-r--r--ep_run/baseline_configs/fw135m_matched_bp.json42
-rw-r--r--ep_run/fw135m_baseline.py117
2 files changed, 159 insertions, 0 deletions
diff --git a/ep_run/baseline_configs/fw135m_matched_bp.json b/ep_run/baseline_configs/fw135m_matched_bp.json
new file mode 100644
index 0000000..0e22686
--- /dev/null
+++ b/ep_run/baseline_configs/fw135m_matched_bp.json
@@ -0,0 +1,42 @@
+{
+ "source_of_truth": "docs/BASELINE_SPEC.md",
+ "purpose": "First approximately 2x BP twin rung after the completed 72M FineWeb model",
+ "changed_from_72m": {
+ "width": "512 -> 768",
+ "heads": "8 -> 12",
+ "parameters": "72,114,688 -> 135,303,936",
+ "training_updates": "440,443 updates to preserve approximately 20 tokens/parameter"
+ },
+ "held_fixed_from_72m": {
+ "layers": 12,
+ "head_dim": 64,
+ "context": 256,
+ "global_sequence_batch": 24,
+ "data": "fineweb_edu",
+ "vocab": 32768,
+ "architecture": "local OLMo2-style",
+ "optimizer": "Muon on hidden matrices plus AdamW on remaining parameters",
+ "muon_lr": 0.02,
+ "adam_side_lr_center": 0.001,
+ "weight_decay": 0.1,
+ "warmup_steps": 500,
+ "schedule": "cosine to 0.1 of peak",
+ "precision": "bf16 autocast with fp32 parameters/states"
+ },
+ "locked_run": {
+ "tag": "fw135m_bp",
+ "layers": 12,
+ "width": 768,
+ "heads": 12,
+ "head_dim": 64,
+ "context": 256,
+ "batch": 24,
+ "parameters": 135303936,
+ "target_tokens_20N": 2706078720,
+ "trainer_steps_argument": 440442,
+ "actual_updates": 440443,
+ "actual_tokens": 2706081792
+ },
+ "bp_lr_sweep_smallest_rung_only": [0.0007, 0.001, 0.0014],
+ "minimum_seeds_before_reporting": 2
+}
diff --git a/ep_run/fw135m_baseline.py b/ep_run/fw135m_baseline.py
new file mode 100644
index 0000000..ddcc5c0
--- /dev/null
+++ b/ep_run/fw135m_baseline.py
@@ -0,0 +1,117 @@
+"""Dry-run-by-default launcher for BASELINE_SPEC.md's first BP twin rung."""
+from __future__ import annotations
+
+import argparse
+import subprocess
+import sys
+
+
+VOCAB = 32768
+LAYERS = 12
+WIDTH = 768
+HEADS = 12
+CONTEXT = 256
+BATCH = 24
+PARAMETERS = 135303936
+FULL_STEPS_ARG = 440442 # casc_bp_train.py loops inclusively: 440,443 updates
+LR_SWEEP = ("7e-4", "1e-3", "1.4e-3")
+
+
+def parameter_count(vocab=VOCAB, layers=LAYERS, width=WIDTH):
+ hidden = ((8 * width // 3) + 63) // 64 * 64
+ return 2 * vocab * width + width + layers * (
+ 4 * width * width + 3 * width * hidden + 4 * width
+ )
+
+
+def command(tag, lr, steps, warmup, seed, wandb_project):
+ return [
+ sys.executable,
+ "casc_bp_train.py",
+ "--tag",
+ tag,
+ "--L",
+ str(LAYERS),
+ "--C",
+ str(WIDTH),
+ "--H",
+ str(HEADS),
+ "--T",
+ str(CONTEXT),
+ "--B",
+ str(BATCH),
+ "--steps",
+ str(steps),
+ "--lr",
+ lr,
+ "--warmup",
+ str(warmup),
+ "--amp",
+ "--olmo2",
+ "--wd",
+ "0.1",
+ "--opt",
+ "muon",
+ "--muon_lr",
+ "0.02",
+ "--cosine",
+ "--lr_min_ratio",
+ "0.1",
+ "--data",
+ "fineweb_edu",
+ "--seed",
+ str(seed),
+ "--save_every",
+ "5000",
+ "--log",
+ "100",
+ "--wandb",
+ wandb_project,
+ "--wandb_run",
+ tag if wandb_project else "",
+ ]
+
+
+def main():
+ ap = argparse.ArgumentParser()
+ ap.add_argument("--mode", choices=("smoke", "sweep", "full"), default="smoke")
+ ap.add_argument("--execute", action="store_true", help="run commands; default only prints")
+ ap.add_argument("--lr", default="1e-3", help="selected LR for --mode full")
+ ap.add_argument("--seed", type=int, default=1)
+ ap.add_argument("--sweep_steps", type=int, default=FULL_STEPS_ARG)
+ ap.add_argument("--wandb_project", default="")
+ args = ap.parse_args()
+
+ count = parameter_count()
+ if count != PARAMETERS:
+ raise RuntimeError(f"parameter formula returned {count}, expected {PARAMETERS}")
+
+ if args.mode == "smoke":
+ runs = [(f"fw135m_bp_smoke_s{args.seed}", args.lr, 400, 50)]
+ elif args.mode == "sweep":
+ runs = [
+ (
+ f"fw135m_bp_lr{lr.replace('-', 'm').replace('.', 'p')}_s{args.seed}",
+ lr,
+ args.sweep_steps,
+ 500,
+ )
+ for lr in LR_SWEEP
+ ]
+ else:
+ runs = [(f"fw135m_bp_s{args.seed}", args.lr, FULL_STEPS_ARG, 500)]
+
+ print(
+ f"# L{LAYERS} C{WIDTH} H{HEADS} T{CONTEXT} B{BATCH} "
+ f"| {PARAMETERS:,} params | target 20N={20 * PARAMETERS:,} tokens",
+ flush=True,
+ )
+ for tag, lr, steps, warmup in runs:
+ cmd = command(tag, lr, steps, warmup, args.seed, args.wandb_project)
+ print(" ".join(f'"{item}"' if item == "" else item for item in cmd), flush=True)
+ if args.execute:
+ subprocess.run(cmd, check=True)
+
+
+if __name__ == "__main__":
+ main()