Sublinear-expander sparse attention / bench_expander.py

Unverified

Raw ⬇ ZIP
  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()