import sys, json, math, time from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, train_model, make_report, permutation_pvalue # Same architecture as bench.transformer_tiny, with a fixed structured token mask. class MaskedTransformer(nn.Module): def __init__(self, win=32, d=64, depth=2, keep=None): super().__init__(); self.win=win self.inp=nn.Linear(1,d); self.pos=nn.Parameter(torch.zeros(1,win,d)); nn.init.normal_(self.pos,std=.02) layer=nn.TransformerEncoderLayer(d,nhead=2,dim_feedforward=128,batch_first=True,dropout=0.0) self.enc=nn.TransformerEncoder(layer,depth); self.head=nn.Linear(win*d,1) self.register_buffer('mask', torch.ones(win) if keep is None else torch.tensor(keep,dtype=torch.float32)) def forward(self,x): x=x*self.mask[:x.shape[1]] h=self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]] return self.head(self.enc(h).reshape(x.shape[0],-1)) def seed_all(seed): np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def interactions(x): # Local interaction graph from calibration covariance; diagonal is marginal importance. a=x-x.mean(0,keepdims=True); cov=(a.T@a)/max(1,len(a)-1) sd=np.sqrt(np.maximum(np.diag(cov),1e-8)); corr=cov/(sd[:,None]*sd[None,:]) np.fill_diagonal(corr,0); return corr def magnitude_keep(x,k): # Standard activation-magnitude token pruning. score=np.mean(np.abs(x),axis=0); ix=np.argsort(score)[-k:] m=np.zeros(x.shape[1],dtype=np.float32); m[ix]=1; return m def boundary_keep(x,k,eta=2,threshold=0.12): # Boundary-compressed DP on a chain ordering. Q support contains local, strong # interactions; state key is (count,last-eta bits), and distant interactions are truncated. q=np.abs(interactions(x)); n=q.shape[0]; adj=(q>=threshold).astype(np.int8); np.fill_diagonal(adj,0) # retain only nearest graph-neighbor interactions, making the intended sparse graph explicit for i in range(n): for j in range(n): if abs(i-j)>eta: adj[i,j]=0 states={(0,0):(0.0,[])}; peaks=[1] for i in range(n): nxt={} for (used,hist),(cost,prefix) in states.items(): for bit in (0,1): if used+bit>k: continue # reward important tokens, penalize selected pair interactions locally val=cost-bit*float(np.mean(np.abs(x[:,i]))) for dist in range(1,min(eta,i)+1): if bit and ((hist>>(dist-1))&1) and adj[i-dist,i]: val += 0.20*float(q[i-dist,i]) nh=((hist<<1)|bit)&((1<0 and np.std(pairs[:,1])>0 else 0.0 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)} rep=make_report('sequence','transformer_tiny',base,idea,extra) rep['idea_sweep']=[{'config':z['config'],'mean':z['mean'],'per_seed':[r['metric'] for r in z['runs']]} for z in idea_runs] 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']] Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2)) if __name__=='__main__': main()