summaryrefslogtreecommitdiff
path: root/ep_run/cascade_probe.py
blob: dda4059416fe699670af71b80c23fe4c49006f7a (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
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
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
"""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', 'zil'], 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
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
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. sumscale: SUM-unit energy (gamma=1 == Z-IL exact sweep)."""
    sc = 1.0 if args.sumscale else 1.0 / NBT
    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() * sc
    if l + 1 < args.L:
        obj = obj + 0.5 * ((zs[l + 1].detach() - blocks[l + 1](zl, mask)) ** 2).sum() * sc
    else:
        obj = obj + beta * (NBT * sc) * 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.init_sweep:                                  # numerical warm-start of the STATE only
        sv = args.sumscale; args.sumscale = True
        for l in range(args.L - 1, -1, -1):
            g = local_grad(zs, z0, l, beta, y)
            zs[l] = (zs[l] - g).detach()
        args.sumscale = sv
    if args.scheme == 'jacobi':
        for z in zs: z.requires_grad_(True)
        opt = (torch.optim.Adam(zs, lr=args.eta) if args.sopt == 'adam'
               else 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, x=None):
    """dF/dtheta at FIXED states (E-term local reads + beta*CE readout term for io params)."""
    prev = (tok(x) + pos(torch.arange(args.T, device=dev))[None]) if (args.include_io and x is not None) else z0
    E = 0.0
    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]

def zil_grads(x, y, beta):
    """Interleaved reverse sweep (Z-IL): update z_l (gamma=1, SUM units) then read theta_l
    IMMEDIATELY (e_l = -beta*delta_l exact at the feedforward point). Single phase; /beta.
    Returns grads matching gate_params, normalized like dFdtheta (E/NBT + CE_mean)."""
    z0g = tok(x) + pos(torch.arange(args.T, device=dev))[None]
    z0 = z0g.detach()
    with torch.no_grad():
        zs = [z.detach() for z in fwd_states(z0)[1:]]
    acc = {id(p): torch.zeros_like(p) for p in gate_params}
    for l in range(args.L - 1, -1, -1):
        zl = zs[l].detach().requires_grad_(True)
        prevd = z0 if l == 0 else zs[l - 1].detach()
        if l + 1 < args.L:
            obj_u = 0.5 * ((zs[l + 1].detach() - blocks[l + 1](zl, mask)) ** 2).sum()
        else:
            obj_u = beta * NBT * F.cross_entropy(readout(zl).reshape(-1, vocab), y.reshape(-1))
        g = torch.autograd.grad(obj_u, zl)[0]
        zs[l] = (zl - g).detach()                                   # gamma=1, SUM units
        prev_read = (z0g if (l == 0 and args.include_io) else prevd)
        obj_r = 0.5 * ((zs[l] - blocks[l](prev_read, mask)) ** 2).sum() / NBT
        if l + 1 == args.L:
            obj_r = obj_r + beta * F.cross_entropy(readout(zs[l]).reshape(-1, vocab), y.reshape(-1))
        gs = torch.autograd.grad(obj_r, gate_params, allow_unused=True)
        for p, gg in zip(gate_params, gs):
            if gg is not None: acc[id(p)] += gg
    return [acc[id(p)] / beta for p in gate_params]

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 (or single-phase interleaved zil)
    if args.scheme == 'zil':
        gep = zil_grads(x, y, args.beta)
    else:
        zp = relax(z0, y, +args.beta); zm = relax(z0, y, -args.beta)
        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)]
    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}{"+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')