From 3605c2cd994643391ebfd0780e15397403dc4144 Mon Sep 17 00:00:00 2001 From: Yuren Hao Date: Thu, 9 Jul 2026 22:28:07 -0500 Subject: A0.4 precision gate (TF32 harmless cos 0.9946==fp32; pure-bf16 cos 0.9427) + Muon hybrid optimizer (--opt muon) wired into both trainers; D1a flagship matrix launched (L12xC512, 3xBP + 3xEP + Muon arms) Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn --- ep_run/casc_bp_train.py | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) (limited to 'ep_run/casc_bp_train.py') diff --git a/ep_run/casc_bp_train.py b/ep_run/casc_bp_train.py index 5194c45..8047091 100644 --- a/ep_run/casc_bp_train.py +++ b/ep_run/casc_bp_train.py @@ -14,6 +14,8 @@ ap.add_argument('--lr', type=float, default=3e-4); ap.add_argument('--warmup', t 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('--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) args = ap.parse_args() torch.manual_seed(args.seed) @@ -47,8 +49,12 @@ if args.tok_init > 0: blocks = nn.ModuleList([Block(args.C, args.H) for _ in range(args.L)]).to(dev) mask = torch.triu(torch.full((args.T, args.T), float('-inf'), device=dev), 1) params = list(tok.parameters()) + list(pos.parameters()) + list(blocks.parameters()) -opt = torch.optim.AdamW(params, lr=args.lr, weight_decay=1e-4) -sched = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: min(1.0, (s + 1) / max(args.warmup, 1))) +if args.opt == 'muon': + from muon import build_hybrid + opt, sched = build_hybrid(blocks, params, args.lr, args.muon_lr, args.warmup) +else: + opt = torch.optim.AdamW(params, lr=args.lr, weight_decay=1e-4) + sched = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: min(1.0, (s + 1) / max(args.warmup, 1))) def fwd(x): z = tok(x) + pos(torch.arange(args.T, device=dev))[None] -- cgit v1.2.3