Sublinear-expander sparse attention / bench_expander.py
Unverified
1import os, sys, math, json, random
2import numpy as np
3import torch
4import torch.nn as nn
5
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
8
9SEEDS = tuple(range(8))
10# Union is used on both sides: baseline and idea each see all three learning rates.
11LRS = (0.0015, 0.003, 0.006)
12EPOCHS = 10
13NTRAIN, NTEST = 400, 400
14D_MODEL, HEADS, DEPTH = 64, 2, 2
15DEGREE = 8
16
17
18def seed_all(seed):
19 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
20 if torch.cuda.is_available():
21 torch.cuda.manual_seed_all(seed)
22
23
24def random_regular(n, d, seed=17):
25 # Fixed-degree undirected circulant with a random relabeling; add self edges.
26 if d % 2 or d >= n: raise ValueError('degree must be even and < n')
27 rng = np.random.RandomState(seed)
28 base = [set((i+z) % n for z in range(1, d//2+1)) |
29 set((i-z) % n for z in range(1, d//2+1)) for i in range(n)]
30 p = rng.permutation(n); adj = [set() for _ in range(n)]
31 for i in range(n):
32 for j in base[i]: adj[p[i]].add(int(p[j]))
33 # query attends to self plus graph neighbors; degree below means non-self edges.
34 return [sorted([i] + list(adj[i])) for i in range(n)]
35
36
37class Attention(nn.Module):
38 def __init__(self, d, heads, adj=None):
39 super().__init__(); assert d % heads == 0
40 self.d, self.h, self.dk, self.adj = d, heads, d//heads, adj
41 self.q = nn.Linear(d, d); self.k = nn.Linear(d, d)
42 self.v = nn.Linear(d, d); self.o = nn.Linear(d, d)
43 self.last_weights = None
44 def forward(self, x):
45 b,n,d = x.shape
46 q = self.q(x).view(b,n,self.h,self.dk).transpose(1,2)
47 k = self.k(x).view(b,n,self.h,self.dk).transpose(1,2)
48 v = self.v(x).view(b,n,self.h,self.dk).transpose(1,2)
49 if self.adj is None:
50 scores = torch.matmul(q, k.transpose(-2,-1)) / math.sqrt(self.dk)
51 w = scores.softmax(-1); z = torch.matmul(w,v)
52 else:
53 idx = torch.as_tensor(self.adj, device=x.device, dtype=torch.long)
54 kk = k[:,:,idx,:] # B,H,N,K,dk
55 vv = v[:,:,idx,:]
56 scores = (q.unsqueeze(3) * kk).sum(-1) / math.sqrt(self.dk)
57 w = scores.softmax(-1); z = (w.unsqueeze(-1)*vv).sum(3)
58 self.last_weights = w.detach()
59 z = z.transpose(1,2).contiguous().view(b,n,d)
60 return self.o(z)
61
62
63class Block(nn.Module):
64 def __init__(self, d, heads, adj):
65 super().__init__(); self.n1=nn.LayerNorm(d); self.attn=Attention(d,heads,adj)
66 self.n2=nn.LayerNorm(d); self.ff=nn.Sequential(nn.Linear(d,128),nn.ReLU(),nn.Linear(128,d))
67 def forward(self,x):
68 x=x+self.attn(self.n1(x)); return x+self.ff(self.n2(x))
69
70
71class Net(nn.Module):
72 def __init__(self, win, sparse):
73 super().__init__(); self.inp=nn.Linear(1,D_MODEL)
74 self.pos=nn.Parameter(torch.zeros(1,win,D_MODEL)); nn.init.normal_(self.pos,std=.02)
75 adj=random_regular(win,DEGREE,17) if sparse else None
76 self.blocks=nn.ModuleList([Block(D_MODEL,HEADS,adj) for _ in range(DEPTH)])
77 self.head=nn.Linear(win*D_MODEL,1); self.adj=adj
78 def forward_features(self,x):
79 z=self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]
80 for block in self.blocks: z=block(z)
81 return z
82 def forward(self,x):
83 z=self.forward_features(x)
84 return self.head(z.reshape(z.shape[0],-1))
85
86
87def train_one(seed, lr, sparse, return_model=False):
88 seed_all(seed); ds=get_dataset('sequence',seed,n_train=NTRAIN,n_test=NTEST)
89 model=Net(ds['input_shape'][0],sparse)
90 net, metric, hist=train_model(model,ds,epochs=EPOCHS,lr=lr,batch=128)
91 if return_model: return metric, net, ds
92 return metric
93
94
95def make_fn(sparse, lr):
96 return lambda seed: train_one(seed, lr, sparse)
97
98
99def signature():
100 # Re-test route growth on trained sparse models: gradient support is measured
101 # after training, while predicted support is graph reachability from token 0.
102 pred=[]; obs=[]
103 metric, net, ds = train_one(0,0.003,True,True)
104 net.eval(); device=next(net.parameters()).device
105 x=ds['xte'][:1].to(device).clone().requires_grad_(True)
106 z=net.forward_features(x); z[0,0,0].backward(); g=x.grad.detach().abs()[0].cpu().numpy()
107 # observed number of input positions with meaningful trained-model sensitivity
108 threshold=max(float(g.max())*1e-3,1e-12)
109 observed=int((g>threshold).sum())
110 seen={0}; frontier={0}; counts=[1]
111 for _ in range(DEPTH):
112 nxt=set()
113 for i in frontier: nxt.update(net.adj[i])
114 nxt-=seen; seen |= nxt; frontier=nxt; counts.append(len(seen))
115 predicted=counts[-1]
116 return {'layers':DEPTH,'degree_nonself':DEGREE,'predicted_reachable_tokens':predicted,
117 'observed_gradient_sensitive_tokens':observed,'predicted_growth_by_layer':counts,
118 'gradient_threshold':threshold,'measurement':'gradient of trained final-layer token-0 latent wrt input positions',
119 'confirmed': bool(abs(observed-predicted) <= max(1,int(.1*predicted)))}
120
121
122def main():
123 os.environ.setdefault('CUDA_VISIBLE_DEVICES','0')
124 # Baseline sweep includes every lr tested by the idea (search-space parity).
125 grid=[{'lr':lr,'epochs':EPOCHS,'degree':DEGREE} for lr in LRS]
126 base=sweep_baseline(lambda cfg: make_fn(False,cfg['lr']),grid,seeds=(0,1,2,3))
127 idea_runs=[]
128 for lr in LRS:
129 r=evaluate(make_fn(True,lr),seeds=SEEDS)
130 idea_runs.append({'cfg':{'lr':lr,'epochs':EPOCHS,'degree':DEGREE},'result':r})
131 best=min(idea_runs,key=lambda q:q['result']['mean'])
132 rep=make_report('sequence','transformer_tiny',base,best['result'],
133 {'graph':'random_regular_plus_self','degree_nonself':DEGREE,
134 'attention_layers':DEPTH,'idea_lr_sweep':idea_runs,
135 'route_growth':signature()})
136 rep['protocol_notes']={'dataset_sizes':[NTRAIN,NTEST],'matched_seeds':list(SEEDS),
137 'structural_match':'sequence forecast requires multi-token correlations; attention is the sole architectural change',
138 'baseline_sweep_union_lrs':list(LRS),'idea_best_cfg':best['cfg']}
139 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
140 print(json.dumps(rep,indent=2))
141
142if __name__=='__main__': main()