import json, math, random from itertools import permutations import numpy as np import torch from torch import nn SEED=1484 def seed(s=SEED): random.seed(s); np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def sym_penalty(w, mu, cov, eps=.5, lcov=1., lw=1.): K=mu.shape[0]; out=0. for p in permutations(range(K)): if list(p)==list(range(K)): continue pp=torch.tensor(p,device=mu.device) dist=torch.linalg.vector_norm(mu-mu[pp],dim=1) dist=dist+lcov*torch.linalg.matrix_norm(cov-cov[pp],dim=(1,2))+lw*torch.abs(w-w[pp]) out=out+torch.nn.functional.softplus(eps-dist).sum() return out def affine_residual(w,mu,cov): # minimum residual over all component permutations, fitting affine map to means K,d=mu.shape; best=[] X=np.concatenate([mu,np.ones((K,1))],1) for p in permutations(range(K)): if list(p)==list(range(K)): continue Y=mu[list(p)] M=np.linalg.lstsq(X,Y,rcond=None)[0]; A=M[:d].T; b=M[d] mr=np.linalg.norm(X@M-Y)/math.sqrt(K) cr=np.mean([np.linalg.norm(A@cov[k]@A.T-cov[p[k]],'fro') for k in range(K)]) wr=np.mean([abs(w[k]-w[p[k]]) for k in range(K)]) best.append((mr+cr+wr,mr,cr,wr,p)) return min(best) def toy_sweep(): # Quadratic confinement makes the otherwise unbounded separation relaxation well posed. rows=[]; eps=.5; c=.08; K=2; d=1 for eta in [0.,.03,.1,.3,1.,3.]: seed(); mu=torch.tensor([[-1e-3],[1e-3]],dtype=torch.float32,requires_grad=True) opt=torch.optim.Adam([mu],lr=.04) for _ in range(1800): w=torch.ones(K)/K; cov=torch.eye(d).repeat(K,1,1) loss=c*(mu**2).sum()+eta*sym_penalty(w,mu,cov,eps,0,0) opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): # nearest pair distance is the relevant lower signature distance D=torch.cdist(mu,mu); mind=D[~torch.eye(K,dtype=torch.bool)].min().item() # For mu=(-r/2,r/2), the objective is c*r^2/2 + 2*eta*softplus(eps-r). # Its stationary equation is c*r = 2*eta*sigmoid(eps-r). if eta == 0: pred=0. else: lo, hi = 0., max(2*eta/c, eps+2*math.log1p(2*eta/c)+2) for _ in range(100): mid=(lo+hi)/2 f=c*mid-2*eta/(1+math.exp(mid-eps)) if f > 0: hi=mid else: lo=mid pred=(lo+hi)/2 rows.append({'eta':eta,'observed_min_mean_distance':mind,'predicted_stationary_distance':pred}) # epsilon sweep tests the margin transition (distance tracks epsilon for strong eta) erows=[] for eps2 in [.1,.3,.5,.8,1.2]: seed(); mu=torch.tensor([[-1e-3],[1e-3]],dtype=torch.float32,requires_grad=True); opt=torch.optim.Adam([mu],lr=.04) for _ in range(1800): loss=c*(mu**2).sum()+1.0*sym_penalty(torch.ones(K)/K,mu,torch.eye(d).repeat(K,1,1),eps2,0,0) opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): erows.append({'epsilon':eps2,'observed_min_mean_distance':torch.cdist(mu,mu)[~torch.eye(K,dtype=torch.bool)].min().item()}) # Exact affine symmetry prediction: identical circular components have zero residual; # a generic unequal signature set should have positive residual. w0=np.ones(4)/4; m0=np.array([[-1,-1],[-1,1],[1,-1],[1,1.]],float); c0=np.tile(np.eye(2),(4,1,1)) generic=(np.array([.1,.2,.3,.4]),m0+np.array([[0,.0],[.2,.1],[-.1,.15],[.05,-.2]]),np.array([np.eye(2),[[1.3,.1],[.1,.8]],[[.7,0],[0,1.4]],[[1.1,.2],[.2,1.2]]])) return {'eta_sweep':rows,'epsilon_sweep':erows,'symmetric_residual':affine_residual(w0,m0,c0),'generic_residual':affine_residual(*generic)} class AE(nn.Module): def __init__(self,K=4,d=2): super().__init__(); self.enc=nn.Sequential(nn.Linear(2,24),nn.Tanh(),nn.Linear(24,d)); self.dec=nn.Sequential(nn.Linear(d,24),nn.Tanh(),nn.Linear(24,2)) self.logits=nn.Parameter(torch.zeros(K)); self.mu=nn.Parameter(torch.randn(K,d)*1.2); self.rawL=nn.Parameter(torch.zeros(K,d, d)); self.K=K; self.d=d def cov(self): L=torch.tril(self.rawL); L=L-torch.diag_embed(torch.diagonal(L,dim1=1,dim2=2))+torch.diag_embed(torch.exp(torch.diagonal(L,dim1=1,dim2=2))); return L@L.transpose(1,2)+.08*torch.eye(self.d,device=L.device) def nll(self,z): C=self.cov(); diff=z[:,None,:]-self.mu[None,:,:]; inv=torch.linalg.inv(C); q=torch.einsum('nkd,kde,nke->nk',diff,inv,diff); ld=torch.logdet(C); lp=torch.log_softmax(self.logits,0)-.5*(q+ld+self.d*math.log(2*math.pi)); return -torch.logsumexp(lp,1).mean() def mini(): seed(); device='cuda' if torch.cuda.is_available() else 'cpu' try: dev=torch.device(device); x=[] centers=np.array([[-2,-2],[-2,2],[2,-2],[2,2.]]) for j in range(4): x.append(centers[j]+.35*np.random.randn(160,2)) x=torch.tensor(np.concatenate(x),dtype=torch.float32,device=dev) ans={} for reg in [False,True]: seed(); model=AE().to(dev); opt=torch.optim.Adam(model.parameters(),lr=.008) for step in range(500): ix=torch.randint(len(x),(64,),device=dev); xb=x[ix]; z=model.enc(xb); recon=((model.dec(z)-xb)**2).mean(); loss=recon+0.08*model.nll(z) if reg: loss=loss+.12*sym_penalty(torch.softmax(model.logits,0),model.mu,model.cov(),.7,.4,.8) opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): z=model.enc(x); rec=((model.dec(z)-x)**2).mean().item(); nll=model.nll(z).item(); ar=affine_residual(torch.softmax(model.logits,0).cpu().numpy(),model.mu.cpu().numpy(),model.cov().cpu().numpy()) ans['regularized' if reg else 'baseline']={'reconstruction_mse':rec,'latent_mixture_nll':nll,'best_affine_residual':ar[:4]} return ans except Exception as e: return {'error':repr(e)} if __name__=='__main__': out={'toy':toy_sweep(),'mini':mini()} with open('results.json','w') as f: json.dump(out,f,indent=2,default=lambda o: o.item() if hasattr(o,'item') else list(o)) print(json.dumps(out,indent=2,default=lambda o: o.item() if hasattr(o,'item') else list(o)))