1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
|
"""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')
|