Boundary-Compressed Approximate Pruning / boundary_bench.py
Beats tuned baseline
1import sys, json, math, time
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, train_model, make_report, permutation_pvalue
8
9# Same architecture as bench.transformer_tiny, with a fixed structured token mask.
10class MaskedTransformer(nn.Module):
11 def __init__(self, win=32, d=64, depth=2, keep=None):
12 super().__init__(); self.win=win
13 self.inp=nn.Linear(1,d); self.pos=nn.Parameter(torch.zeros(1,win,d)); nn.init.normal_(self.pos,std=.02)
14 layer=nn.TransformerEncoderLayer(d,nhead=2,dim_feedforward=128,batch_first=True,dropout=0.0)
15 self.enc=nn.TransformerEncoder(layer,depth); self.head=nn.Linear(win*d,1)
16 self.register_buffer('mask', torch.ones(win) if keep is None else torch.tensor(keep,dtype=torch.float32))
17 def forward(self,x):
18 x=x*self.mask[:x.shape[1]]
19 h=self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]
20 return self.head(self.enc(h).reshape(x.shape[0],-1))
21
22def seed_all(seed):
23 np.random.seed(seed); torch.manual_seed(seed)
24 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
25
26def interactions(x):
27 # Local interaction graph from calibration covariance; diagonal is marginal importance.
28 a=x-x.mean(0,keepdims=True); cov=(a.T@a)/max(1,len(a)-1)
29 sd=np.sqrt(np.maximum(np.diag(cov),1e-8)); corr=cov/(sd[:,None]*sd[None,:])
30 np.fill_diagonal(corr,0); return corr
31
32def magnitude_keep(x,k):
33 # Standard activation-magnitude token pruning.
34 score=np.mean(np.abs(x),axis=0); ix=np.argsort(score)[-k:]
35 m=np.zeros(x.shape[1],dtype=np.float32); m[ix]=1; return m
36
37def boundary_keep(x,k,eta=2,threshold=0.12):
38 # Boundary-compressed DP on a chain ordering. Q support contains local, strong
39 # interactions; state key is (count,last-eta bits), and distant interactions are truncated.
40 q=np.abs(interactions(x)); n=q.shape[0]; adj=(q>=threshold).astype(np.int8); np.fill_diagonal(adj,0)
41 # retain only nearest graph-neighbor interactions, making the intended sparse graph explicit
42 for i in range(n):
43 for j in range(n):
44 if abs(i-j)>eta: adj[i,j]=0
45 states={(0,0):(0.0,[])}; peaks=[1]
46 for i in range(n):
47 nxt={}
48 for (used,hist),(cost,prefix) in states.items():
49 for bit in (0,1):
50 if used+bit>k: continue
51 # reward important tokens, penalize selected pair interactions locally
52 val=cost-bit*float(np.mean(np.abs(x[:,i])))
53 for dist in range(1,min(eta,i)+1):
54 if bit and ((hist>>(dist-1))&1) and adj[i-dist,i]:
55 val += 0.20*float(q[i-dist,i])
56 nh=((hist<<1)|bit)&((1<<eta)-1)
57 key=(used+bit,nh)
58 if key not in nxt or val<nxt[key][0]: nxt[key]=(val,prefix+[bit])
59 states=nxt; peaks.append(len(states))
60 candidates=[v for (u,_),v in states.items() if u==k]
61 _,bits=min(candidates,key=lambda z:z[0]); return np.asarray(bits,dtype=np.float32),max(peaks),int(adj.sum()/2)
62
63def run_one(seed,lr,method,epochs=10,n=1200):
64 seed_all(seed); ds=get_dataset('sequence',seed,n_train=n,n_test=400); x=ds['xtr'].numpy()
65 if method=='baseline': keep=magnitude_keep(x,16); peak=0; edges=0
66 else: keep,peak,edges=boundary_keep(x,16,eta=2,threshold=.12)
67 # train_model is the canonical benchmark path; only the fixed structured mask differs.
68 seed_all(seed)
69 net=MaskedTransformer(32,64,2,keep)
70 net,metric,hist=train_model(net,ds,epochs=epochs,lr=lr,batch=128,log=lambda *a:None)
71 return {'seed':seed,'metric':float(metric),'keep':int(keep.sum()),'peak_states':peak,'graph_edges':edges,
72 'retained_fraction':float(keep.mean()),'history_last':float(hist[-1])}
73
74def baseline_sweep(lrs,epochs=10):
75 out=[]
76 for lr in lrs:
77 vals=[run_one(s,lr,'baseline',epochs)['metric'] for s in range(4)]
78 out.append({'config':{'lr':lr,'epochs':epochs,'keep':16},'mean':float(np.mean(vals)),'per_seed':vals})
79 best=min(out,key=lambda z:z['mean'])
80 return {'grid':out,'best':best['config'],'full':{'per_seed':[run_one(s,best['config']['lr'],'baseline',epochs)['metric'] for s in range(8)]}}
81
82def main():
83 # Union parity: all idea lrs are also baseline sweep values.
84 lrs=[0.0015,0.003,0.006]; epochs=10; t=time.time()
85 base=baseline_sweep(lrs,epochs)
86 idea_runs=[]
87 for lr in lrs:
88 rr=[run_one(s,lr,'idea',epochs) for s in range(8)]
89 idea_runs.append({'config':{'lr':lr,'epochs':epochs,'eta':2,'keep':16},'mean':float(np.mean([r['metric'] for r in rr])),'runs':rr})
90 best=min(idea_runs,key=lambda z:z['mean']); idea={'config':best['config'],'per_seed':[r['metric'] for r in best['runs']], 'details':best['runs']}
91 # Behavioral signature: prediction from omitted empirical interaction mass vs observed
92 sig=[]
93 for r in best['runs']:
94 s=r['seed']; ds=get_dataset('sequence',s,n_train=1200,n_test=400); x=ds['xte'].numpy(); q=np.abs(interactions(x));
95 # Reconstruct selected masks deterministically and measure trained-model masked sensitivity.
96 keep,_,_=boundary_keep(get_dataset('sequence',s,n_train=1200,n_test=400)['xtr'].numpy(),16,2,.12)
97 pred=float(q[np.outer(1-keep,1-keep).astype(bool)].mean()) if np.any(keep==0) else 0.0
98 seed_all(s); net=MaskedTransformer(32,64,2,keep); net,_,_=train_model(net,ds,epochs=epochs,lr=best['config']['lr'],batch=128,log=lambda *a:None)
99 dev=next(net.parameters()).device
100 xt=ds['xte'].to(dev)
101 with torch.no_grad():
102 trained_mask=net.mask.detach().clone()
103 net.mask.fill_(1.0)
104 full=net(xt)
105 net.mask.copy_(trained_mask)
106 masked=net(xt)
107 obs=float(((full-masked)**2).mean()); sig.append((pred,obs))
108 pairs=np.asarray(sig); corr=float(np.corrcoef(pairs.T)[0,1]) if np.std(pairs[:,0])>0 and np.std(pairs[:,1])>0 else 0.0
109 extra={'prediction':'omitted-token interaction mass predicts output sensitivity','predicted_mean':float(pairs[:,0].mean()),'observed_mean':float(pairs[:,1].mean()),'correlation':corr,'confirmed':bool(corr>0.3)}
110 rep=make_report('sequence','transformer_tiny',base,idea,extra)
111 rep['idea_sweep']=[{'config':z['config'],'mean':z['mean'],'per_seed':[r['metric'] for r in z['runs']]} for z in idea_runs]
112 rep['runtime_sec']=time.time()-t; rep['structural_match']='multi-token sequence forecasting with token interaction graph'; rep['state_counts']=[r['peak_states'] for r in best['runs']]
113 Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
114if __name__=='__main__': main()