import json, math, random import torch.nn.functional as F import numpy as np import torch from torch import nn SEED=283 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) device='cuda' if torch.cuda.is_available() else 'cpu' try: if device=='cuda': torch.cuda.get_device_name(0) except Exception: device='cpu' class SubsetRouter(nn.Module): def __init__(self,d,n): super().__init__(); self.n=n self.score=nn.Sequential(nn.Linear(2*d,32),nn.Tanh(),nn.Linear(32,1)) self.pair=nn.Sequential(nn.Linear(2*d,32),nn.Tanh(),nn.Linear(32,1)) self.decoder=nn.Sequential(nn.Linear(2*d,32),nn.Tanh(),nn.Linear(32,d)) self.register_buffer('masks',torch.tensor([[((m>>i)&1) for i in range(n)] for m in range(1,1<bkd',M,x) dec=self.decoder(torch.cat((agg,z[:,None,:].expand(-1,agg.size(1),-1)),-1)) 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 pi=torch.softmax(logw,-1) r=torch.einsum('bk,kn->bn',pi,M) # normalized reconstruction from subset decoded observations recon=torch.einsum('bk,bkd->bd',pi,dec) return r,recon,pi,logw class IndependentRouter(nn.Module): def __init__(self,d): 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)) def forward(self,x,z): logits=self.net(torch.cat((x,z[:,None,:].expand_as(x)),-1)).squeeze(-1) r=torch.sigmoid(logits) recon=self.decoder((r[:,:,None]*x).sum(1)) return r,recon def make_data(num,n,d): # Each sample has 1-3 mutually similar members of a latent collision group. xs=[]; zs=[]; ys=[] for _ in range(num): k=random.randint(1,3); center=torch.randn(d); y=torch.zeros(n); y[:k]=1 x=torch.randn(n,d)*1.0; x[:k]=center+0.12*torch.randn(k,d) perm=torch.randperm(n); x=x[perm]; y=y[perm] z=center+0.08*torch.randn(d) xs.append(x); zs.append(z); ys.append(y) return torch.stack(xs),torch.stack(zs),torch.stack(ys) def f1(r,y): pred=(r>0.5).float(); tp=(pred*y).sum(); return float((2*tp/(pred.sum()+y.sum()+1e-8)).item()) def train(model,train,test,subset=False,steps=500): opt=torch.optim.Adam(model.parameters(),lr=3e-3) x,z,y=train; xt,zt,yt=test; model.to(device); x=x.to(device); z=z.to(device); y=y.to(device) for step in range(steps): ix=torch.randint(0,len(x),(64,),device=device); xb,zb,yb=x[ix],z[ix],y[ix] out=model(xb,zb); r,recon=out[:2] # responsibility supervision plus observation reconstruction loss=nn.functional.binary_cross_entropy(r,yb)+0.5*nn.functional.mse_loss(recon,zb) opt.zero_grad(); loss.backward(); opt.step() model.eval() with torch.no_grad(): out=model(xt.to(device),zt.to(device)); r,recon=out[:2] return f1(r.cpu(),yt),float(nn.functional.mse_loss(recon,zt.to(device)).item()) def math_check(): torch.manual_seed(SEED); n=5; d=4 x=torch.randn(1,n,d); z=torch.randn(1,d) m=SubsetRouter(d,n).eval() with torch.no_grad(): r,recon,pi,logw=m(x,z) # Exact claims: posterior sums to one; marginal is subset inclusion sum; # positive pair compatibility increases mass of subsets containing the pair. masks=m.masks direct=torch.einsum('bk,kn->bn',pi,masks) before=float(pi.sum()); err=float((direct-r).abs().max()) # isolate a pair by changing only its symmetric pair logit via log weights pair=(0,1); boost=torch.zeros_like(logw) boost[0]=masks[:,pair[0]]*masks[:,pair[1]]*2.0 p2=torch.softmax(logw+boost,1) pairmass_before=float((pi[0]*(masks[:,0]*masks[:,1])).sum()) pairmass_after=float((p2[0]*(masks[:,0]*masks[:,1])).sum()) return {'posterior_sum':before,'marginal_max_error':err,'pair_mass_before':pairmass_before,'pair_mass_after_boost':pairmass_after} def main(): trainset=make_data(768,6,8); testset=make_data(256,6,8) # same initialization seed and optimizer budget; subset has exact local enumeration (63 subsets) torch.manual_seed(SEED); base=IndependentRouter(8) b=train(base,trainset,testset,steps=500) torch.manual_seed(SEED); idea=SubsetRouter(8,6) s=train(idea,trainset,testset,subset=True,steps=500) 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} print(json.dumps(result,indent=2)) with open('results.json','w') as f: json.dump(result,f,indent=2) if __name__=='__main__': main()