Sublinear-expander sparse attention / bench_stage2.py
Beats tuned baseline
1import sys, json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
9
10SEEDS = tuple(range(8))
11# Same union on both sides: baseline is evaluated at every idea LR.
12GRID = [{'lr': 0.0015, 'weight_decay': 0.0},
13 {'lr': 0.0030, 'weight_decay': 0.0},
14 {'lr': 0.0060, 'weight_decay': 0.0}]
15EPOCHS = 10
16BATCH = 128
17DEGREE = 8
18
19
20def graph(n, degree, seed):
21 """Randomly relabelled circulant, exactly degree regular (including self)."""
22 if degree % 2 or degree >= n: raise ValueError('degree must be even and < n')
23 rng = np.random.default_rng(seed)
24 adj = [set() for _ in range(n)]
25 for i in range(n):
26 for z in range(1, degree//2 + 1):
27 adj[i].add((i-z) % n); adj[i].add((i+z) % n)
28 p = rng.permutation(n); out = [None] * n
29 for i in range(n): out[int(p[i])] = sorted(int(p[j]) for j in adj[i])
30 return out
31
32
33class SparseSelfAttention(nn.Module):
34 def __init__(self, d=64, heads=2, degree=8, n=32, seed=0):
35 super().__init__(); assert d % heads == 0
36 self.d, self.h, self.dk, self.degree = d, heads, d//heads, degree
37 a = graph(n, degree, seed)
38 self.register_buffer('idx', torch.tensor(a, dtype=torch.long), persistent=False)
39 self.qkv = nn.Linear(d, 3*d); self.proj = nn.Linear(d, d)
40 self.last_attention = None
41
42 def forward(self, x):
43 b,n,d = x.shape; qkv = self.qkv(x).view(b,n,3,self.h,self.dk)
44 q,k,v = qkv[:,:,0],qkv[:,:,1],qkv[:,:,2]
45 jj = self.idx[:n]
46 kg = k[:,jj,:]; vg = v[:,jj,:]
47 scores = (q.unsqueeze(2) * kg).sum(-1) / math.sqrt(self.dk)
48 att = torch.softmax(scores, dim=2)
49 self.last_attention = att.detach()
50 y = (att.unsqueeze(-1) * vg).sum(2).transpose(1,2).reshape(b,n,d)
51 return self.proj(y)
52
53
54class SparseLayer(nn.Module):
55 def __init__(self, seed):
56 super().__init__(); self.attn=SparseSelfAttention(seed=seed, degree=DEGREE)
57 self.n1=nn.LayerNorm(64); self.ff=nn.Sequential(nn.Linear(64,128),nn.ReLU(),nn.Linear(128,64)); self.n2=nn.LayerNorm(64)
58 def forward(self,x):
59 x=self.n1(x+self.attn(x)); return self.n2(x+self.ff(x))
60
61
62class SparseTransformer(nn.Module):
63 def __init__(self, n=32, out_dim=1, seed=0):
64 super().__init__(); self.inp=nn.Linear(1,64); self.pos=nn.Parameter(torch.zeros(1,n,64)); nn.init.normal_(self.pos,std=.02)
65 self.enc=nn.ModuleList([SparseLayer(seed+i*7919) for i in range(2)]); self.head=nn.Linear(n*64,out_dim)
66 def forward(self,x):
67 h=self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]
68 for layer in self.enc: h=layer(h)
69 return self.head(h.reshape(x.shape[0],-1))
70
71
72def baseline_fn(cfg):
73 def run(seed):
74 torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
75 d=get_dataset('sequence',seed,n_train=400,n_test=200)
76 m=make_model('transformer_tiny',d['input_shape'],d['out_dim'])
77 _, metric, _=train_model(m,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,weight_decay=cfg['weight_decay'],log=lambda *_:None)
78 return float(metric)
79 return run
80
81
82def idea_fn(cfg):
83 def run(seed):
84 torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
85 d=get_dataset('sequence',seed,n_train=400,n_test=200)
86 m=SparseTransformer(n=d['input_shape'][0],out_dim=d['out_dim'],seed=seed)
87 _, metric, _=train_model(m,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,weight_decay=cfg['weight_decay'],log=lambda *_:None)
88 return float(metric)
89 return run
90
91
92def signature(seed, cfg):
93 """Measured on a trained sparse model: weighted support and boundary growth."""
94 d=get_dataset('sequence',seed,n_train=400,n_test=200)
95 m=SparseTransformer(n=32,out_dim=1,seed=seed)
96 m,_,_=train_model(m,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,weight_decay=cfg['weight_decay'],log=lambda *_:None)
97 m.eval()
98 dev=next(m.parameters()).device
99 with torch.no_grad(): m(d['xte'][:64].to(dev))
100 # Average effective outgoing support (mass >= 1% of a row) and external boundary.
101 U=set(range(4)); layer=m.enc[-1]; att=layer.attn.last_attention.mean((0,3)).cpu().numpy()
102 support=[]
103 for i in U: support.append({int(layer.attn.idx[i,j]) for j in range(DEGREE) if att[i,j] >= .01})
104 ext=set().union(*support)-U
105 observed=len(ext)/len(U)
106 eps=1.0; predicted=eps/(math.log(3*len(U)/2)**2)
107 return {'subset_size':4,'k':2,'epsilon':eps,'predicted_external_ratio':predicted,
108 'observed_effective_external_ratio':float(observed),
109 'predicted_growth_factor':1+predicted,'observed_growth_factor':1+observed,
110 'confirmed': bool(observed+1e-9 >= predicted*.8)}
111
112
113def main():
114 # Baseline sweep uses 4 seeds; final best config is reevaluated on all 8 by harness.
115 base=sweep_baseline(baseline_fn,GRID,seeds=(0,1,2,3))
116 # Equal-sized idea sweep, selecting on the same four seeds and then full 8.
117 tried=[]
118 for cfg in GRID:
119 r=evaluate(idea_fn(cfg),seeds=(0,1,2,3)); tried.append({'cfg':cfg,'mean':r['mean']})
120 best=min(GRID,key=lambda c: next(x['mean'] for x in tried if x['cfg']==c))
121 idea=evaluate(idea_fn(best),seeds=SEEDS)
122 base['idea_union_sweep']=tried
123 sig=signature(0,best)
124 rep=make_report('sequence','transformer_tiny',base,idea,{
125 'mechanism_signature':sig,
126 'implementation':{'degree':DEGREE,'epochs':EPOCHS,'n_train':400,'n_test':200,'matched_lr_grid':GRID},
127 'structural_match':'sequence-level multi-token forecast; sparse self-attention is the sole architectural intervention'
128 })
129 Path('bench_report.json').write_text(json.dumps(rep,indent=2))
130 print(json.dumps(rep,indent=2))
131
132if __name__=='__main__': main()