summaryrefslogtreecommitdiff
path: root/ep_run
diff options
context:
space:
mode:
Diffstat (limited to 'ep_run')
-rw-r--r--ep_run/casc_bp_train.py10
-rw-r--r--ep_run/casc_eq_train.py10
-rw-r--r--ep_run/cascade_probe.py22
-rw-r--r--ep_run/muon.py64
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)