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
|
"""Cascade-EP trainer — EQUILIBRIUM MODE (the true-EP route).
Two-phase (+-beta) relaxation of all layer states to the nudged equilibria via
Gauss-Seidel reverse sweeps (solver choice only; readout is taken AT the relaxed
states with the standard EP formula), weight grad = (1/2beta)[dF/dtheta|+ - dF/dtheta|-].
Inference = plain forward (standard LLM). Twin of casc_bp_train.py (same seed/data)."""
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('--tag', default='casc_eq6')
ap.add_argument('--L', type=int, default=6); ap.add_argument('--C', type=int, default=256)
ap.add_argument('--H', type=int, default=8); ap.add_argument('--T', type=int, default=256)
ap.add_argument('--B', type=int, default=24); ap.add_argument('--steps', type=int, default=4000)
ap.add_argument('--lr', type=float, default=3e-4); ap.add_argument('--warmup', type=int, default=200)
ap.add_argument('--beta', type=float, default=0.03); ap.add_argument('--seed', type=int, default=0)
ap.add_argument('--K', type=int, default=20) # GS reverse sweeps per phase
ap.add_argument('--geta', type=float, default=0.8) # GS state step (sum units)
ap.add_argument('--save_every', type=int, default=1000); ap.add_argument('--log', type=int, default=100)
ap.add_argument('--wandb', default=''); ap.add_argument('--wandb_run', default='')
args = ap.parse_args()
torch.manual_seed(args.seed)
dev = 'cuda' if torch.cuda.is_available() else 'cpu'
DD = Path('/home/yurenh2/ept/ep_run/data/tinystories_bpe')
vocab = pickle.load(open(DD / 'meta.pkl', 'rb'))['vocab_size']
def get_batch(split):
data = np.memmap(DD / ('train.bin' if split == 'train' else '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])
y = torch.stack([torch.from_numpy(data[i + 1:i + 1 + args.T].astype(np.int64)) for i in ix])
return x.to(dev), y.to(dev)
class Block(nn.Module):
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)
mask = torch.triu(torch.full((args.T, args.T), float('-inf'), device=dev), 1)
readout = lambda z: z @ tok.weight.t()
all_params = list(tok.parameters()) + list(pos.parameters()) + list(blocks.parameters())
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(x):
with torch.no_grad():
z = tok(x) + pos(torch.arange(args.T, device=dev))[None]
z0 = z.clone(); zs = []
for b in blocks:
z = b(z, mask); zs.append(z)
return z0, zs
def relax(z0, zs_free, y, beta):
"""GS reverse sweeps to the nudged equilibrium (SUM-unit local grads, step geta)."""
zs = [z.clone() for z in zs_free]
for _ in range(args.K):
for l in range(args.L - 1, -1, -1):
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()
if l + 1 < args.L:
obj = obj + 0.5 * ((zs[l + 1].detach() - blocks[l + 1](zl, mask)) ** 2).sum()
else:
obj = obj + beta * NBT * F.cross_entropy(readout(zl).reshape(-1, vocab), y.reshape(-1))
g = torch.autograd.grad(obj, zl)[0]
zs[l] = (zl - args.geta * g).detach()
return zs
def dFdtheta(zs, x, y, beta):
"""dF/dtheta at fixed relaxed states (z0 rebuilt WITH graph so emb gets its E-path grad)."""
prev = tok(x) + pos(torch.arange(args.T, device=dev))[None]
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, all_params, allow_unused=True)
return [g if g is not None else None for g in gs]
def ep_step(x, y):
z0, zs_free = free_states(x)
free_ce = F.cross_entropy(readout(zs_free[-1]).reshape(-1, vocab), y.reshape(-1)).item()
zp = relax(z0, zs_free, y, +args.beta)
zm = relax(z0, zs_free, y, -args.beta)
gp = dFdtheta(zp, x, y, +args.beta)
gm = dFdtheta(zm, x, y, -args.beta)
for p, a, b in zip(all_params, gp, gm):
p.grad = None if (a is None or b is None) else (a - b) / (2 * args.beta)
return free_ce
@torch.no_grad()
def evaluate(nb=6):
tot = 0.0
for _ in range(nb):
x, y = get_batch('val')
z = tok(x) + pos(torch.arange(args.T, device=dev))[None]
for b in blocks: z = b(z, mask)
tot += F.cross_entropy(readout(z).reshape(-1, vocab), y.reshape(-1)).item()
return tot / nb
wb = None
if args.wandb:
try:
import wandb as _w
wb = _w.init(project=args.wandb, name=args.wandb_run or args.tag, id=args.wandb_run or args.tag,
resume='allow', config=vars(args))
except Exception as e:
print(f'[wandb] disabled ({e})', flush=True)
n = sum(p.numel() for p in all_params)
print(f'[{args.tag}] cascade-EP(EQUILIBRIUM) L{args.L} C{args.C} T{args.T} beta={args.beta} '
f'K={args.K} geta={args.geta} | {n/1e6:.2f}M | {dev}', flush=True)
best, t0 = 1e9, time.time()
for step in range(args.steps + 1):
x, y = get_batch('train')
ce = ep_step(x, y)
torch.nn.utils.clip_grad_norm_(all_params, 1.0)
opt.step(); sched.step(); opt.zero_grad(set_to_none=True)
if step % args.log == 0:
val = evaluate(); best = min(best, val)
print(f'step {step:5d}/{args.steps} | train {ce:.4f} val {val:.4f} (best {best:.4f}) '
f'| {step/max(time.time()-t0,1e-9):.3f} it/s', flush=True)
if wb is not None:
try: wb.log({'train_ce': ce, 'val_ce': val, 'best': best}, step=step)
except Exception: pass
if step % args.save_every == 0 and step > 0:
torch.save({'tok': tok.state_dict(), 'pos': pos.state_dict(), 'blocks': blocks.state_dict(),
'step': step, 'val': best, 'config': vars(args)}, Path('runs') / f'{args.tag}_s{step}.pt')
print(f'[{args.tag}] DONE best val CE {best:.4f} (BP twin 2.9746; zil-diagnostic 3.3236)', flush=True)
if wb is not None:
try: wb.summary['best_val_ce'] = best; wb.finish()
except Exception: pass
|