diff options
| author | Yuren Hao <yurenh2@illinois.edu> | 2026-07-09 22:28:07 -0500 |
|---|---|---|
| committer | Yuren Hao <yurenh2@illinois.edu> | 2026-07-09 22:28:07 -0500 |
| commit | 3605c2cd994643391ebfd0780e15397403dc4144 (patch) | |
| tree | af416d1d6b9cc68471b6a10ade381b2e2efd9832 /ep_run | |
| parent | 654dfb94d727f7514470a7ca909fc865e34636d8 (diff) | |
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 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn
Diffstat (limited to 'ep_run')
| -rw-r--r-- | ep_run/casc_bp_train.py | 10 | ||||
| -rw-r--r-- | ep_run/casc_eq_train.py | 10 | ||||
| -rw-r--r-- | ep_run/cascade_probe.py | 22 | ||||
| -rw-r--r-- | ep_run/muon.py | 64 |
4 files changed, 99 insertions, 7 deletions
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] diff --git a/ep_run/casc_eq_train.py b/ep_run/casc_eq_train.py index 10206f6..4eba267 100644 --- a/ep_run/casc_eq_train.py +++ b/ep_run/casc_eq_train.py @@ -21,6 +21,8 @@ ap.add_argument('--wandb', default=''); ap.add_argument('--wandb_run', default=' ap.add_argument('--kmax', type=int, default=8) # adaptive fb rounds cap ap.add_argument('--noguard', action='store_true') # diagnosis: skip only non-finite grads ap.add_argument('--untie', action='store_true') # separate readout matrix (untied from tok) +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 with this std (GPT-standard 0.02) ap.add_argument('--compile', action='store_true') # torch.compile each block (free speed where supported) ap.add_argument('--sig_every', type=int, default=25) # tok-sigma refresh interval (amortized) @@ -66,8 +68,12 @@ mask = torch.triu(torch.full((args.T, args.T), float('-inf'), device=dev), 1) W_out = nn.Parameter(torch.randn(vocab, args.C, device=dev) * 0.02) if args.untie else None readout = (lambda z: z @ W_out.t()) if args.untie else (lambda z: z @ tok.weight.t()) all_params = list(tok.parameters()) + list(pos.parameters()) + list(blocks.parameters()) + ([W_out] if args.untie else []) -opt = torch.optim.AdamW(all_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, all_params, args.lr, args.muon_lr, args.warmup) +else: + opt = torch.optim.AdamW(all_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))) NBT = args.B * args.T def free_states_graphed(x): diff --git a/ep_run/cascade_probe.py b/ep_run/cascade_probe.py index d6d9fea..a6bb3aa 100644 --- a/ep_run/cascade_probe.py +++ b/ep_run/cascade_probe.py @@ -26,11 +26,18 @@ ap.add_argument('--ckpt', type=str, default='') # gate at a casc_bp_trai ap.add_argument('--sumscale', action='store_true') # relax in SUM units (gamma=1 Z-IL semantics) ap.add_argument('--init_sweep', action='store_true') # one reverse gamma=1 sweep as state INIT (readout still at equilibrium = clean EP) ap.add_argument('--sopt', choices=['sgd', 'adam'], default='sgd') # state-relaxation optimizer +ap.add_argument('--prec', choices=['fp32', 'tf32', 'bf16'], default='fp32') # A0.4 precision gate (EP side only; BP ref stays fp32) ap.add_argument('--geta_auto', type=float, default=0.0) # >0: per-layer gamma_l = c/(1+sigma_l+1^2), c=this; sigma via power-iter args = ap.parse_args() torch.manual_seed(args.seed) dev = 'cuda' if torch.cuda.is_available() else 'cpu' -torch.set_float32_matmul_precision('highest') +if args.prec == 'fp32': + torch.set_float32_matmul_precision('highest') +elif args.prec == 'tf32': + torch.backends.cuda.matmul.allow_tf32 = True + torch.backends.cudnn.allow_tf32 = True + torch.set_float32_matmul_precision('high') +_BF16 = (args.prec == 'bf16') DD = Path('/home/yurenh2/ept/ep_run/data/tinystories_bpe') vocab = pickle.load(open(DD / 'meta.pkl', 'rb'))['vocab_size'] @@ -237,7 +244,11 @@ for bi in range(args.batches): ce = F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1)) gbp = list(torch.autograd.grad(ce, gate_params, allow_unused=True)) gbp = [g if g is not None else torch.zeros(1, device=dev) for g in gbp] - # two-phase (or single-phase interleaved zil) + # EP side: pure-bf16 MODEL when --prec bf16 (BP reference above stays fp32); cast back after + if _BF16: + sd32 = [ (m, {k: v.clone() for k, v in m.state_dict().items()}) for m in [tok, pos, blocks] ] + for m in [tok, pos, blocks]: m.bfloat16() + z0 = z0.bfloat16(); mask = mask.bfloat16() if args.scheme == 'zil': gep = zil_grads(x, y, args.beta) else: @@ -245,6 +256,11 @@ for bi in range(args.batches): gp = dFdtheta(zp, z0, y, +args.beta, gate_params, x) gm = dFdtheta(zm, z0, y, -args.beta, gate_params, x) gep = [(a - b) / (2 * args.beta) for a, b in zip(gp, gm)] + gep = [g.float() for g in gep] + if _BF16: + for m, sd in sd32: + m.float(); m.load_state_dict(sd) + mask = mask.float() cos_b.append(F.cosine_similarity(flat(gep), flat(gbp), dim=0).item()) shr_b.append((flat(gep).norm() / flat(gbp).norm()).item()) for l in range(args.L): @@ -257,6 +273,6 @@ print('per-block cos: ' + ' '.join(f'{l}:{np.mean(v):.4f}' for l, v in enumerat if args.include_io: idx = [i for i, n in enumerate(names) if n.startswith('io.')] print(f'io cos (last batch): {F.cosine_similarity(flat([gep[i] for i in idx]), flat([gbp[i] for i in idx]), dim=0).item():.4f}') -print(f'SUMMARY scheme={args.scheme}{"+init" if args.init_sweep else ""}+{args.sopt} L={args.L} C={args.C} K={args.K} eta={args.eta} beta={args.beta} ' +print(f'SUMMARY prec={args.prec} scheme={args.scheme}{"+init" if args.init_sweep else ""}+{args.sopt} L={args.L} C={args.C} K={args.K} eta={args.eta} beta={args.beta} ' f'io={int(args.include_io)} warp={args.warp} ckpt={args.ckpt or "-"} ' f'cos={cm:.4f} cosmin={cmin:.4f} shrink={sm:.3f} t={time.time()-t0:.1f}s') diff --git a/ep_run/muon.py b/ep_run/muon.py new file mode 100644 index 0000000..5ac6fc6 --- /dev/null +++ b/ep_run/muon.py @@ -0,0 +1,64 @@ +"""Muon optimizer (Newton-Schulz orthogonalized momentum) + hybrid helpers. +Convention: Muon on 2D hidden matrices, AdamW on everything else (emb/pos/LN/bias).""" +import torch + + +def newton_schulz(G, steps=5, eps=1e-7): + """approximate polar factor of G via the quintic NS iteration (Keller Jordan coefficients).""" + a, b, c = 3.4445, -4.7750, 2.0315 + X = G / (G.norm() + eps) + transposed = X.size(0) > X.size(1) + if transposed: X = X.T + for _ in range(steps): + A = X @ X.T + B = b * A + c * (A @ A) + X = a * X + B @ X + return X.T if transposed else X + + +class Muon(torch.optim.Optimizer): + def __init__(self, params, lr=0.02, momentum=0.95, ns_steps=5, nesterov=True): + super().__init__(params, dict(lr=lr, momentum=momentum, ns_steps=ns_steps, nesterov=nesterov)) + + @torch.no_grad() + def step(self, closure=None): + for group in self.param_groups: + for p in group['params']: + if p.grad is None: continue + g = p.grad + st = self.state[p] + if 'mom' not in st: st['mom'] = torch.zeros_like(g) + buf = st['mom'] + buf.mul_(group['momentum']).add_(g) + u = g.add(buf, alpha=group['momentum']) if group['nesterov'] else buf + if u.ndim == 2: + u = newton_schulz(u, group['ns_steps']) + u = u * max(1.0, u.size(0) / u.size(1)) ** 0.5 # rms-matched scaling + p.add_(u, alpha=-group['lr']) + + +class MultiOpt: + """duck-typed bundle of optimizers (step/zero_grad API-compatible).""" + def __init__(self, opts): self.optimizers = opts + def step(self): + for o in self.optimizers: o.step() + def zero_grad(self, set_to_none=True): + for o in self.optimizers: o.zero_grad(set_to_none=set_to_none) + + +class MultiSched: + def __init__(self, scheds): self.scheds = scheds + def step(self): + for s in self.scheds: s.step() + + +def build_hybrid(blocks, other_params, lr_adamw, lr_muon, warmup): + """Muon(2D block matrices) + AdamW(everything else), with linear-warmup scheds for both.""" + mats = [p for p in blocks.parameters() if p.ndim == 2] + mat_ids = {id(p) for p in mats} + rest = [p for p in other_params if id(p) not in mat_ids] + om = Muon(mats, lr=lr_muon) + oa = torch.optim.AdamW(rest, lr=lr_adamw, weight_decay=1e-4) + fn = lambda s: min(1.0, (s + 1) / max(warmup, 1)) + scheds = [torch.optim.lr_scheduler.LambdaLR(om, fn), torch.optim.lr_scheduler.LambdaLR(oa, fn)] + return MultiOpt([om, oa]), MultiSched(scheds) |
