Collision-Aware Subset Attention / bench_collision_subset.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys,json,random
 2from pathlib import Path
 3import numpy as np, torch
 4from torch import nn
 5import torch.nn.functional as F
 6sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
 7from bench import get_dataset,make_model,train_model,sweep_baseline,make_report
 8EPOCHS=12; NTRAIN=1200; NTEST=400; LRS=[1e-3,3e-3,6e-3]
 9def seed_all(s):
10 random.seed(s);np.random.seed(s);torch.manual_seed(s)
11 if torch.cuda.is_available():
12  try: torch.cuda.manual_seed_all(s)
13  except Exception: pass
14class SubsetAttn(nn.Module):
15 def __init__(self,d=64,n=4):
16  super().__init__();self.n=n
17  self.score=nn.Linear(2*d,1);self.pair=nn.Sequential(nn.Linear(2*d,32),nn.Tanh(),nn.Linear(32,1));self.dec=nn.Sequential(nn.Linear(2*d,64),nn.Tanh(),nn.Linear(64,d))
18  self.register_buffer('masks',torch.tensor([[(m>>i)&1 for i in range(n)] for m in range(1,1<<n)],dtype=torch.float32))
19 def posterior(self,x,z,boost=0.):
20  B,n,d=x.shape;M=self.masks.to(x);zz=z[:,None,:].expand(-1,n,-1)
21  s=self.score(torch.cat((x,zz),-1)).squeeze(-1);xa=x[:,:,None,:].expand(-1,-1,n,-1);xb=x[:,None,:,:].expand(-1,n,-1,-1)
22  q=self.pair(torch.cat((xa,xb),-1)).squeeze(-1);q=(q+q.transpose(1,2))/2;logq=F.logsigmoid(q)
23  agg=torch.einsum('kn,bnd->bkd',M,x);dec=self.dec(torch.cat((agg,z[:,None,:].expand(-1, M.shape[0],-1)),-1))
24  pm=torch.einsum('ki,kj->kij',M,M)*(1-torch.eye(n,device=x.device));lw=torch.einsum('bn,kn->bk',s,M)+.5*torch.einsum('bij,kij->bk',logq,pm)
25  if boost: lw=lw+boost*M[:,0]*M[:,1]
26  pi=torch.softmax(lw,-1);r=torch.einsum('bk,kn->bn',pi,M);u=torch.einsum('bk,bkd->bd',pi,dec)
27  return r,u,(pi*M[:,0]*M[:,1]).sum(-1)
28 def forward(self,x):
29  B,T,D=x.shape;parts=[]
30  for st in range(0,T,self.n):
31   c=x[:,st:st+self.n];actual=c.shape[1]
32   if actual<self.n:c=torch.cat((c,c[:,-1:].expand(-1,self.n-actual,-1)),1)
33   z=c.mean(1);r,u,_=self.posterior(c,z);parts.append(c+r.unsqueeze(-1)*u.unsqueeze(1))
34  return torch.cat(parts,1)[:,:T]
35class MatchedTransformer(nn.Module):
36 def __init__(self,win,out_dim,idea=False,d=64,depth=2):
37  super().__init__();self.idea=idea;self.inp=nn.Linear(1,d);self.pos=nn.Parameter(torch.zeros(1,win,d));nn.init.normal_(self.pos,std=.02)
38  self.norms=nn.ModuleList([nn.LayerNorm(d) for _ in range(depth)]);self.ff=nn.ModuleList([nn.Sequential(nn.Linear(d,128),nn.GELU(),nn.Linear(128,d)) for _ in range(depth)])
39  if idea:self.attn=nn.ModuleList([SubsetAttn(d,4) for _ in range(depth)])
40  else:self.attn=nn.ModuleList([nn.MultiheadAttention(d,2,dropout=0.,batch_first=True) for _ in range(depth)])
41  self.head=nn.Linear(win*d,out_dim);self.last_pair=None
42 def forward(self,x):
43  h=self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]
44  for i in range(len(self.ff)):
45   q=self.norms[i](h)
46   if self.idea:
47    a=self.attn[i](q);h=h+a
48   else:
49    a,_=self.attn[i](q,q,q,need_weights=False);h=h+a
50   h=h+self.ff[i](self.norms[i](h))
51  return self.head(h.reshape(x.shape[0],-1))
52def train_one(kind,cfg,seed,keep=False):
53 seed_all(seed);ds=get_dataset('sequence',seed,n_train=NTRAIN,n_test=NTEST)
54 net=MatchedTransformer(ds['input_shape'][0],ds['out_dim'],kind=='idea')
55 out=train_model(net,ds,epochs=cfg['epochs'],lr=cfg['lr'],batch=128,weight_decay=0.,log=lambda *_:None)
56 return (float(out[1]),net,ds) if keep else float(out[1])
57def main():
58 grid=[{'lr':x,'epochs':EPOCHS} for x in LRS]
59 base=sweep_baseline(lambda c:lambda s:train_one('baseline',c,s),grid)
60 trials=[]
61 for c in grid:
62  vals=[train_one('idea',c,s) for s in range(8)];trials.append({'cfg':c,'mean':float(np.mean(vals)),'per_seed':vals})
63 best=min(trials,key=lambda x:x['mean']);idea={'mean':float(np.mean(best['per_seed'])),'std':float(np.std(best['per_seed'])),'per_seed':best['per_seed'],'n':8,'best_cfg':best['cfg'],'sweep':trials}
64 _,net,ds=train_one('idea',best['cfg'],0,True);dev=next(net.parameters()).device
65 with torch.no_grad():
66  h=net.inp(ds['xte'][:128].to(dev).unsqueeze(-1))+net.pos[:,:32];q=net.norms[0](h);z=q[:,:4].mean(1);_,_,b=net.attn[0].posterior(q[:,:4],z);_,_,a=net.attn[0].posterior(q[:,:4],z,2.)
67 sig={'quantity':'trained-model pair inclusion mass','predicted_direction':'compatibility boost increases joint inclusion','boost':2.,'observed_before':float(b.mean()),'observed_after':float(a.mean()),'observed_change':float((a-b).mean()),'confirmed':bool(float(a.mean())>float(b.mean()))}
68 rep=make_report('sequence','transformer_tiny',base,idea,{'mechanism_signature':sig,'selection':{'baseline_best_cfg':base['best_cfg'],'idea_best_cfg':best['cfg'],'shared_lr_union':LRS},'structural_match':'multi-token correlated sequence window'})
69 rep['protocol_note']='Eight paired seeds; three shared learning rates; reduced 1200/400 samples and 12 epochs for runtime.'
70 Path('bench_report.json').write_text(json.dumps(rep,indent=2));print(json.dumps(rep,indent=2))
71if __name__=='__main__':main()