Collision-Aware Subset Attention / collision_subset_attention.py

Mechanism confirmed, baseline not beaten

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