Boundary-Compressed Approximate Pruning / boundary_bench.py

✓✓ Beats tuned baseline

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