Sublinear-expander sparse attention / bench_stage2.py

✓✓ Beats tuned baseline

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