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
|
"""Cascade-EP trainer (zil mode): train L DISTINCT standard transformer blocks with the
interleaved reverse-sweep energy rule — per layer: update z_l (gamma=1, SUM units), then
read that layer's theta-grad from its LOCAL energy term. Numerically == BP restructured
as local two-factor rules; inference = plain forward (standard LLM).
Twin of casc_bp_train.py (same seed => same init & data stream)."""
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_ep6')
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('--save_every', type=int, default=1000); ap.add_argument('--log', type=int, default=200)
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()
io_params = list(tok.parameters()) + list(pos.parameters())
all_params = io_params + list(blocks.parameters())
blk_params = [list(b.parameters()) for b in blocks]
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 zil_step(x, y):
"""One cascade-EP (zil) gradient computation. Returns free-phase CE for logging;
leaves grads in p.grad (normalized to CE_mean units)."""
z0g = tok(x) + pos(torch.arange(args.T, device=dev))[None]
z0 = z0g.detach()
with torch.no_grad():
zs, z = [], z0
for b in blocks:
z = b(z, mask); zs.append(z)
free_ce = F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1)).item()
for p in all_params: p.grad = None
for l in range(args.L - 1, -1, -1):
zl = zs[l].detach().requires_grad_(True)
# 1) state update: gamma=1 in SUM units (top gets the beta-nudge)
if l + 1 < args.L:
obj_u = 0.5 * ((zs[l + 1].detach() - blocks[l + 1](zl, mask)) ** 2).sum()
else:
obj_u = args.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()
# 2) immediate local read: this block's params (+ io at the ends)
prev = z0g if l == 0 else zs[l - 1].detach() # l=0 keeps emb graph
params_l = blk_params[l] + (io_params if l == 0 else [])
obj_r = 0.5 * ((zs[l] - blocks[l](prev, mask)) ** 2).sum() / NBT
if l + 1 == args.L:
obj_r = obj_r + args.beta * F.cross_entropy(readout(zs[l]).reshape(-1, vocab), y.reshape(-1))
params_l = params_l + [tok.weight]
seen, uniq = set(), []
for p in params_l:
if id(p) not in seen: seen.add(id(p)); uniq.append(p)
gs = torch.autograd.grad(obj_r, uniq, allow_unused=True)
for p, gg in zip(uniq, gs):
if gg is None: continue
p.grad = gg if p.grad is None else p.grad + gg
for p in all_params:
if p.grad is not None: p.grad /= 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(zil) L{args.L} C{args.C} H{args.H} T{args.T} beta={args.beta} | {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 = zil_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):.2f} 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 reference: casc_bp6)', flush=True)
if wb is not None:
try: wb.summary['best_val_ce'] = best; wb.finish()
except Exception: pass
|