summaryrefslogtreecommitdiff
path: root/ep_run/cascade_probe.py
blob: 572097c18a0603c82032e7f46e9175af10ca9d4f (plain)
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
"""Cascade-EP gradient gate: 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
(inference = plain LLM).  Training = two-phase (±beta nudge) relaxation of z, weight
grad = (1/2beta)[dE/dtheta|+ - dE/dtheta|-].  Gate: cosine vs true BP gradient."""
import argparse, math, pickle
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('--seed', type=int, default=0); ap.add_argument('--trained_warp', type=float, default=1.0)
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')
ix = torch.randint(len(data) - args.T - 1, (args.B,))
x = torch.stack([torch.from_numpy(data[i:i + args.T].astype(np.int64)) for i in ix]).to(dev)
y = torch.stack([torch.from_numpy(data[i + 1:i + 1 + args.T].astype(np.int64)) for i in ix]).to(dev)

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.trained_warp != 1.0:                       # crude "not-at-init" landscape
    with torch.no_grad():
        for p in blocks.parameters(): p.mul_(args.trained_warp)
mask = torch.triu(torch.full((args.T, args.T), float('-inf'), device=dev), 1)
readout = lambda z: z @ tok.weight.t()

def fwd_states(z0):
    zs = [z0]
    for b in blocks: zs.append(b(zs[-1], mask))
    return zs

z0 = (tok(x) + pos(torch.arange(args.T, device=dev))[None]).detach()

# ---------------- true BP reference ----------------
for p in blocks.parameters(): p.requires_grad_(True)
zs = fwd_states(z0)
ce = F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1))
gbp = torch.autograd.grad(ce, list(blocks.parameters()), allow_unused=True)
gbp = [g if g is not None else torch.zeros(1, device=dev) for g in gbp]
print(f'BP ref: CE {ce.item():.4f}')

# ---------------- two-phase cascade EP ----------------
def relax(beta):
    """minimize E + beta*CE over free states z_1..z_L by GD (free init = forward pass)."""
    with torch.no_grad():
        zs = [z.detach().clone() for z in fwd_states(z0)[1:]]
    for z in zs: z.requires_grad_(True)
    opt = torch.optim.SGD(zs, lr=args.eta, momentum=0.9)
    for k 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
        Fobj = E / (args.B * args.T) + beta * F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1))
        Fobj.backward()
        opt.step()
    return [z.detach() for z in zs], (E / (args.B * args.T)).item()

def dEdtheta(zs):
    """dE/dtheta at fixed states (the local EP readout)."""
    for p in blocks.parameters():
        if p.grad is not None: p.grad = None
    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 / (args.B * args.T)).backward()
    return [(p.grad.clone() if p.grad is not None else torch.zeros(1, device=dev)) for p in blocks.parameters()]

zp, Ep = relax(+args.beta)
zm, Em = relax(-args.beta)
gp, gm = dEdtheta(zp), dEdtheta(zm)
gep = [(a - b) / (2 * args.beta) for a, b in zip(gp, gm)]
print(f'nudged relax E+: {Ep:.3e}  E-: {Em:.3e}  (K={args.K}, eta={args.eta}, beta={args.beta})')

# ---------------- gate ----------------
names = [n for n, _ in blocks.named_parameters()]
flat = lambda gs: torch.cat([g.reshape(-1) for g in gs])
cos_all = F.cosine_similarity(flat(gep), flat(gbp), dim=0).item()
print(f'\nGATE cos(cascadeEP, BP) overall: {cos_all:.4f}')
for l in range(args.L):
    idx = [i for i, n in enumerate(names) if n.startswith(f'{l}.')]
    c = F.cosine_similarity(flat([gep[i] for i in idx]), flat([gbp[i] for i in idx]), dim=0).item()
    r = (flat([gep[i] for i in idx]).norm() / flat([gbp[i] for i in idx]).norm()).item()
    print(f'  block {l}: cos {c:.4f}  |EP|/|BP| {r:.3f}')