Collision-Aware Subset Attention / collision_subset_attention.py
Mechanism confirmed, baseline not beaten
1import json, math, random
2import torch.nn.functional as F
3import numpy as np
4import torch
5from torch import nn
6
7SEED=283
8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
9torch.set_num_threads(4)
10device='cuda' if torch.cuda.is_available() else 'cpu'
11try:
12 if device=='cuda': torch.cuda.get_device_name(0)
13except Exception:
14 device='cpu'
15
16class SubsetRouter(nn.Module):
17 def __init__(self,d,n):
18 super().__init__(); self.n=n
19 self.score=nn.Sequential(nn.Linear(2*d,32),nn.Tanh(),nn.Linear(32,1))
20 self.pair=nn.Sequential(nn.Linear(2*d,32),nn.Tanh(),nn.Linear(32,1))
21 self.decoder=nn.Sequential(nn.Linear(2*d,32),nn.Tanh(),nn.Linear(32,d))
22 self.register_buffer('masks',torch.tensor([[((m>>i)&1) for i in range(n)] for m in range(1,1<<n)],dtype=torch.float32))
23 def forward(self,x,z):
24 B,N,D=x.shape; M=self.masks.to(x)
25 s=self.score(torch.cat((x,z[:,None,:].expand(-1,N,-1)),-1)).squeeze(-1)
26 xa=x[:,:,None,:].expand(-1,-1,N,-1); xb=x[:,None,:,:].expand(-1,N,-1,-1)
27 qlog=self.pair(torch.cat((xa,xb),-1)).squeeze(-1); qlog=(qlog+qlog.transpose(1,2))/2
28 # log sigmoid is numerically stable and equals log Q
29 logq=F.logsigmoid(qlog)
30 agg=torch.einsum('kn,bnd->bkd',M,x)
31 dec=self.decoder(torch.cat((agg,z[:,None,:].expand(-1,agg.size(1),-1)),-1))
32 logw=torch.einsum('bn,kn->bk',s,M)+torch.einsum('bij,kij->bk',logq,torch.einsum('ki,kj->kij',M,M)* (1-torch.eye(N,device=x.device))) / 2
33 pi=torch.softmax(logw,-1)
34 r=torch.einsum('bk,kn->bn',pi,M)
35 # normalized reconstruction from subset decoded observations
36 recon=torch.einsum('bk,bkd->bd',pi,dec)
37 return r,recon,pi,logw
38
39class IndependentRouter(nn.Module):
40 def __init__(self,d):
41 super().__init__(); self.net=nn.Sequential(nn.Linear(2*d,32),nn.Tanh(),nn.Linear(32,1)); self.decoder=nn.Sequential(nn.Linear(d,32),nn.Tanh(),nn.Linear(32,d))
42 def forward(self,x,z):
43 logits=self.net(torch.cat((x,z[:,None,:].expand_as(x)),-1)).squeeze(-1)
44 r=torch.sigmoid(logits)
45 recon=self.decoder((r[:,:,None]*x).sum(1))
46 return r,recon
47
48def make_data(num,n,d):
49 # Each sample has 1-3 mutually similar members of a latent collision group.
50 xs=[]; zs=[]; ys=[]
51 for _ in range(num):
52 k=random.randint(1,3); center=torch.randn(d); y=torch.zeros(n); y[:k]=1
53 x=torch.randn(n,d)*1.0; x[:k]=center+0.12*torch.randn(k,d)
54 perm=torch.randperm(n); x=x[perm]; y=y[perm]
55 z=center+0.08*torch.randn(d)
56 xs.append(x); zs.append(z); ys.append(y)
57 return torch.stack(xs),torch.stack(zs),torch.stack(ys)
58
59def f1(r,y):
60 pred=(r>0.5).float(); tp=(pred*y).sum(); return float((2*tp/(pred.sum()+y.sum()+1e-8)).item())
61
62def train(model,train,test,subset=False,steps=500):
63 opt=torch.optim.Adam(model.parameters(),lr=3e-3)
64 x,z,y=train; xt,zt,yt=test; model.to(device); x=x.to(device); z=z.to(device); y=y.to(device)
65 for step in range(steps):
66 ix=torch.randint(0,len(x),(64,),device=device); xb,zb,yb=x[ix],z[ix],y[ix]
67 out=model(xb,zb); r,recon=out[:2]
68 # responsibility supervision plus observation reconstruction
69 loss=nn.functional.binary_cross_entropy(r,yb)+0.5*nn.functional.mse_loss(recon,zb)
70 opt.zero_grad(); loss.backward(); opt.step()
71 model.eval()
72 with torch.no_grad():
73 out=model(xt.to(device),zt.to(device)); r,recon=out[:2]
74 return f1(r.cpu(),yt),float(nn.functional.mse_loss(recon,zt.to(device)).item())
75
76def math_check():
77 torch.manual_seed(SEED); n=5; d=4
78 x=torch.randn(1,n,d); z=torch.randn(1,d)
79 m=SubsetRouter(d,n).eval()
80 with torch.no_grad(): r,recon,pi,logw=m(x,z)
81 # Exact claims: posterior sums to one; marginal is subset inclusion sum;
82 # positive pair compatibility increases mass of subsets containing the pair.
83 masks=m.masks
84 direct=torch.einsum('bk,kn->bn',pi,masks)
85 before=float(pi.sum()); err=float((direct-r).abs().max())
86 # isolate a pair by changing only its symmetric pair logit via log weights
87 pair=(0,1); boost=torch.zeros_like(logw)
88 boost[0]=masks[:,pair[0]]*masks[:,pair[1]]*2.0
89 p2=torch.softmax(logw+boost,1)
90 pairmass_before=float((pi[0]*(masks[:,0]*masks[:,1])).sum())
91 pairmass_after=float((p2[0]*(masks[:,0]*masks[:,1])).sum())
92 return {'posterior_sum':before,'marginal_max_error':err,'pair_mass_before':pairmass_before,'pair_mass_after_boost':pairmass_after}
93
94def main():
95 trainset=make_data(768,6,8); testset=make_data(256,6,8)
96 # same initialization seed and optimizer budget; subset has exact local enumeration (63 subsets)
97 torch.manual_seed(SEED); base=IndependentRouter(8)
98 b=train(base,trainset,testset,steps=500)
99 torch.manual_seed(SEED); idea=SubsetRouter(8,6)
100 s=train(idea,trainset,testset,subset=True,steps=500)
101 result={'device':device,'math_check':math_check(),'baseline':{'f1':b[0],'reconstruction_mse':b[1]},'subset_attention':{'f1':s[0],'reconstruction_mse':s[1]},'subsets_per_neighborhood':63,'steps':500,'seed':SEED}
102 print(json.dumps(result,indent=2))
103 with open('results.json','w') as f: json.dump(result,f,indent=2)
104if __name__=='__main__': main()