summaryrefslogtreecommitdiff
path: root/ep_run/casc_bp_train.py
diff options
context:
space:
mode:
authorYuren Hao <yurenh2@illinois.edu>2026-07-13 13:56:16 -0500
committerYuren Hao <yurenh2@illinois.edu>2026-07-13 13:56:16 -0500
commitd15ae909dfb9c8870a79e6aea87e8085fda53460 (patch)
tree18d92b7b6ebef37b0344a3d7b631e2904a3896e3 /ep_run/casc_bp_train.py
parent5cdab992b203991161d9ec5235cab64f9c705ad6 (diff)
wandb ON by default in both trainers (project ept-cascade); BP twin gains --amp; fw72m_bp launched GPU1
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn
Diffstat (limited to 'ep_run/casc_bp_train.py')
-rw-r--r--ep_run/casc_bp_train.py12
1 files changed, 7 insertions, 5 deletions
diff --git a/ep_run/casc_bp_train.py b/ep_run/casc_bp_train.py
index ebd7536..6f09d0c 100644
--- a/ep_run/casc_bp_train.py
+++ b/ep_run/casc_bp_train.py
@@ -13,7 +13,8 @@ ap.add_argument('--B', type=int, default=24); ap.add_argument('--steps', type=in
ap.add_argument('--lr', type=float, default=3e-4); ap.add_argument('--warmup', type=int, default=200)
ap.add_argument('--seed', type=int, default=0)
ap.add_argument('--save_every', type=int, default=500); ap.add_argument('--log', type=int, default=200)
-ap.add_argument('--wandb', default=''); ap.add_argument('--wandb_run', default='')
+ap.add_argument('--wandb', default='ept-cascade') # ON BY DEFAULT (user directive 07-13); pass --wandb '' to disable; ap.add_argument('--wandb_run', default='')
+ap.add_argument('--amp', action='store_true') # bf16 autocast fwd/loss (no scaler needed for bf16)
ap.add_argument('--opt', choices=['adamw', 'muon'], default='adamw')
ap.add_argument('--muon_lr', type=float, default=0.02)
ap.add_argument('--tok_init', type=float, default=0.0) # >0: init tok/pos std (GPT-standard 0.02)
@@ -196,10 +197,11 @@ outdir = Path('runs'); outdir.mkdir(exist_ok=True)
for _ in range(start_step): sched.step() # advance LR schedule to the resumed step
for step in range(start_step, args.steps + 1):
x, y = get_batch('train')
- logits = fwd(x).reshape(-1, vocab)
- loss = F.cross_entropy(logits, y.reshape(-1))
- if args.zloss > 0:
- loss = loss + args.zloss * (torch.logsumexp(logits.float(), -1) ** 2).mean()
+ with torch.autocast('cuda', dtype=torch.bfloat16, enabled=args.amp):
+ logits = fwd(x).reshape(-1, vocab)
+ loss = F.cross_entropy(logits, y.reshape(-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()
torch.nn.utils.clip_grad_norm_(params, 1.0)
opt.step(); sched.step()