"""Cascade-EP gradient gate v2 — L distinct standard transformer blocks, layered energy E = sum_l 0.5||z_l - f_l(z_{l-1})||^2. Free equilibrium == the standard forward pass. Two-phase (+-beta) relaxation of states -> local weight gradient; gate vs true BP. v2: relaxation schemes (jacobi = physical parallel | gsf/gsr = Gauss-Seidel fwd/reverse sweeps), --include_io (emb/pos/tied-readout grads via the full EP formula), multi-batch gates, machine-parseable SUMMARY line. Usage examples: python3 cascade_probe.py --L 12 --K 480 python3 cascade_probe.py --L 6 --scheme gsr --K 12 --include_io """ import argparse, math, pickle, time import numpy as np, torch, torch.nn as nn, torch.nn.functional as F from pathlib import Path ap = argparse.ArgumentParser() ap.add_argument('--L', type=int, default=3); ap.add_argument('--C', type=int, default=128) ap.add_argument('--H', type=int, default=4); ap.add_argument('--T', type=int, default=32) ap.add_argument('--B', type=int, default=4); ap.add_argument('--beta', type=float, default=0.03) ap.add_argument('--K', type=int, default=400); ap.add_argument('--eta', type=float, default=0.3) ap.add_argument('--mom', type=float, default=0.9) ap.add_argument('--scheme', choices=['jacobi', 'gsf', 'gsr'], default='jacobi') ap.add_argument('--batches', type=int, default=4) ap.add_argument('--include_io', action='store_true') # gate emb/pos/tied-readout too ap.add_argument('--seed', type=int, default=0); ap.add_argument('--warp', type=float, default=1.0) ap.add_argument('--ckpt', type=str, default='') # gate at a casc_bp_train checkpoint args = ap.parse_args() torch.manual_seed(args.seed) dev = 'cuda' if torch.cuda.is_available() else 'cpu' torch.set_float32_matmul_precision('highest') DD = Path('/home/yurenh2/ept/ep_run/data/tinystories_bpe') vocab = pickle.load(open(DD / 'meta.pkl', 'rb'))['vocab_size'] data = np.memmap(DD / 'val.bin', dtype=np.uint16, mode='r') class Block(nn.Module): """Standard pre-LN transformer block (distinct weights per layer, no recurrence).""" def __init__(self, C, H): super().__init__() self.ln1, self.ln2 = nn.LayerNorm(C), nn.LayerNorm(C) self.attn = nn.MultiheadAttention(C, H, batch_first=True) self.ff = nn.Sequential(nn.Linear(C, 4 * C), nn.GELU(), nn.Linear(4 * C, C)) def forward(self, z, mask): h = self.ln1(z); a, _ = self.attn(h, h, h, attn_mask=mask, need_weights=False) z = z + a; return z + self.ff(self.ln2(z)) tok = nn.Embedding(vocab, args.C).to(dev) pos = nn.Embedding(args.T, args.C).to(dev) blocks = nn.ModuleList([Block(args.C, args.H) for _ in range(args.L)]).to(dev) if args.ckpt: sd = torch.load(args.ckpt, map_location=dev) tok.load_state_dict(sd['tok']); pos.load_state_dict(sd['pos']); blocks.load_state_dict(sd['blocks']) print(f"[gate] loaded {args.ckpt} (step {sd.get('step')}, val {sd.get('val', float('nan')):.4f})") elif args.warp != 1.0: with torch.no_grad(): for p in blocks.parameters(): p.mul_(args.warp) mask = torch.triu(torch.full((args.T, args.T), float('-inf'), device=dev), 1) readout = lambda z: z @ tok.weight.t() io_params = list(tok.parameters()) + list(pos.parameters()) blk_params = list(blocks.parameters()) gate_params = blk_params + (io_params if args.include_io else []) NBT = args.B * args.T def get_batch(i): g = torch.Generator().manual_seed(1000 + i) ix = torch.randint(len(data) - args.T - 1, (args.B,), generator=g) x = torch.stack([torch.from_numpy(data[j:j + args.T].astype(np.int64)) for j in ix]).to(dev) y = torch.stack([torch.from_numpy(data[j + 1:j + 1 + args.T].astype(np.int64)) for j in ix]).to(dev) return x, y def fwd_states(z0): zs = [z0] for b in blocks: zs.append(b(zs[-1], mask)) return zs def local_grad(zs, z0, l, beta, y): """dF/dz_l with neighbors fixed (zs: list of L free states, 0-indexed block l).""" zl = zs[l].detach().requires_grad_(True) prev = z0 if l == 0 else zs[l - 1].detach() obj = 0.5 * ((zl - blocks[l](prev, mask)) ** 2).sum() / NBT if l + 1 < args.L: obj = obj + 0.5 * ((zs[l + 1].detach() - blocks[l + 1](zl, mask)) ** 2).sum() / NBT else: obj = obj + beta * F.cross_entropy(readout(zl).reshape(-1, vocab), y.reshape(-1)) return torch.autograd.grad(obj, zl)[0] def relax(z0, y, beta): """minimize F_beta over states; free init = forward pass. Returns relaxed states.""" with torch.no_grad(): zs = [z.detach().clone() for z in fwd_states(z0)[1:]] if args.scheme == 'jacobi': for z in zs: z.requires_grad_(True) opt = torch.optim.SGD(zs, lr=args.eta, momentum=args.mom) for _ in range(args.K): opt.zero_grad() E = 0.0; prev = z0 for z, b in zip(zs, blocks): E = E + 0.5 * ((z - b(prev, mask)) ** 2).sum(); prev = z (E / NBT + beta * F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1))).backward() opt.step() return [z.detach() for z in zs] order = range(args.L) if args.scheme == 'gsf' else range(args.L - 1, -1, -1) bufs = [torch.zeros_like(z) for z in zs] for _ in range(args.K): for l in order: g = local_grad(zs, z0, l, beta, y) bufs[l].mul_(args.mom).add_(g) zs[l] = (zs[l] - args.eta * bufs[l]).detach() return zs def dFdtheta(zs, z0, y, beta, params): """dF/dtheta at FIXED states (E-term local reads + beta*CE readout term for io params).""" E = 0.0; prev = z0 for z, b in zip(zs, blocks): E = E + 0.5 * ((z - b(prev, mask)) ** 2).sum(); prev = z obj = E / NBT + beta * F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1)) gs = torch.autograd.grad(obj, params, allow_unused=True) return [g if g is not None else torch.zeros(1, device=dev) for g in gs] flat = lambda gs: torch.cat([g.reshape(-1) for g in gs]) names = ([f'blk.{n}' for n, _ in blocks.named_parameters()] + ([f'io.{n}' for n, _ in list(tok.named_parameters()) + list(pos.named_parameters())] if args.include_io else [])) cos_b, shr_b, blkcos = [], [], [[] for _ in range(args.L)] t0 = time.time() for bi in range(args.batches): x, y = get_batch(bi) z0 = (tok(x) + pos(torch.arange(args.T, device=dev))[None]).detach() # BP reference (embedding z0 grad flows for io gate) z0ref = tok(x) + pos(torch.arange(args.T, device=dev))[None] zs = [z0ref] for b in blocks: zs.append(b(zs[-1], mask)) 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 zp = relax(z0, y, +args.beta); zm = relax(z0, y, -args.beta) gp = dFdtheta(zp, z0, y, +args.beta, gate_params) gm = dFdtheta(zm, z0, y, -args.beta, gate_params) gep = [(a - b) / (2 * args.beta) for a, b in zip(gp, gm)] 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): idx = [i for i, n in enumerate(names) if n.startswith(f'blk.{l}.')] blkcos[l].append(F.cosine_similarity(flat([gep[i] for i in idx]), flat([gbp[i] for i in idx]), dim=0).item()) cm, cmin = float(np.mean(cos_b)), float(np.min(cos_b)) sm = float(np.mean(shr_b)) print('per-block cos: ' + ' '.join(f'{l}:{np.mean(v):.4f}' for l, v in enumerate(blkcos))) 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} 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')