Affine-symmetry-free GMM latent prior / experiment.py
Unverified
1import json, math, random
2from itertools import permutations
3import numpy as np
4import torch
5from torch import nn
6
7SEED=1484
8
9def seed(s=SEED):
10 random.seed(s); np.random.seed(s); torch.manual_seed(s)
11 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
12
13def sym_penalty(w, mu, cov, eps=.5, lcov=1., lw=1.):
14 K=mu.shape[0]; out=0.
15 for p in permutations(range(K)):
16 if list(p)==list(range(K)): continue
17 pp=torch.tensor(p,device=mu.device)
18 dist=torch.linalg.vector_norm(mu-mu[pp],dim=1)
19 dist=dist+lcov*torch.linalg.matrix_norm(cov-cov[pp],dim=(1,2))+lw*torch.abs(w-w[pp])
20 out=out+torch.nn.functional.softplus(eps-dist).sum()
21 return out
22
23def affine_residual(w,mu,cov):
24 # minimum residual over all component permutations, fitting affine map to means
25 K,d=mu.shape; best=[]
26 X=np.concatenate([mu,np.ones((K,1))],1)
27 for p in permutations(range(K)):
28 if list(p)==list(range(K)): continue
29 Y=mu[list(p)]
30 M=np.linalg.lstsq(X,Y,rcond=None)[0]; A=M[:d].T; b=M[d]
31 mr=np.linalg.norm(X@M-Y)/math.sqrt(K)
32 cr=np.mean([np.linalg.norm(A@cov[k]@A.T-cov[p[k]],'fro') for k in range(K)])
33 wr=np.mean([abs(w[k]-w[p[k]]) for k in range(K)])
34 best.append((mr+cr+wr,mr,cr,wr,p))
35 return min(best)
36
37def toy_sweep():
38 # Quadratic confinement makes the otherwise unbounded separation relaxation well posed.
39 rows=[]; eps=.5; c=.08; K=2; d=1
40 for eta in [0.,.03,.1,.3,1.,3.]:
41 seed(); mu=torch.tensor([[-1e-3],[1e-3]],dtype=torch.float32,requires_grad=True)
42 opt=torch.optim.Adam([mu],lr=.04)
43 for _ in range(1800):
44 w=torch.ones(K)/K; cov=torch.eye(d).repeat(K,1,1)
45 loss=c*(mu**2).sum()+eta*sym_penalty(w,mu,cov,eps,0,0)
46 opt.zero_grad(); loss.backward(); opt.step()
47 with torch.no_grad():
48 # nearest pair distance is the relevant lower signature distance
49 D=torch.cdist(mu,mu); mind=D[~torch.eye(K,dtype=torch.bool)].min().item()
50 # For mu=(-r/2,r/2), the objective is c*r^2/2 + 2*eta*softplus(eps-r).
51 # Its stationary equation is c*r = 2*eta*sigmoid(eps-r).
52 if eta == 0: pred=0.
53 else:
54 lo, hi = 0., max(2*eta/c, eps+2*math.log1p(2*eta/c)+2)
55 for _ in range(100):
56 mid=(lo+hi)/2
57 f=c*mid-2*eta/(1+math.exp(mid-eps))
58 if f > 0: hi=mid
59 else: lo=mid
60 pred=(lo+hi)/2
61 rows.append({'eta':eta,'observed_min_mean_distance':mind,'predicted_stationary_distance':pred})
62 # epsilon sweep tests the margin transition (distance tracks epsilon for strong eta)
63 erows=[]
64 for eps2 in [.1,.3,.5,.8,1.2]:
65 seed(); mu=torch.tensor([[-1e-3],[1e-3]],dtype=torch.float32,requires_grad=True); opt=torch.optim.Adam([mu],lr=.04)
66 for _ in range(1800):
67 loss=c*(mu**2).sum()+1.0*sym_penalty(torch.ones(K)/K,mu,torch.eye(d).repeat(K,1,1),eps2,0,0)
68 opt.zero_grad(); loss.backward(); opt.step()
69 with torch.no_grad(): erows.append({'epsilon':eps2,'observed_min_mean_distance':torch.cdist(mu,mu)[~torch.eye(K,dtype=torch.bool)].min().item()})
70 # Exact affine symmetry prediction: identical circular components have zero residual;
71 # a generic unequal signature set should have positive residual.
72 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))
73 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]]]))
74 return {'eta_sweep':rows,'epsilon_sweep':erows,'symmetric_residual':affine_residual(w0,m0,c0),'generic_residual':affine_residual(*generic)}
75
76class AE(nn.Module):
77 def __init__(self,K=4,d=2):
78 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))
79 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
80 def cov(self):
81 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)
82 def nll(self,z):
83 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()
84
85def mini():
86 seed(); device='cuda' if torch.cuda.is_available() else 'cpu'
87 try:
88 dev=torch.device(device); x=[]
89 centers=np.array([[-2,-2],[-2,2],[2,-2],[2,2.]])
90 for j in range(4): x.append(centers[j]+.35*np.random.randn(160,2))
91 x=torch.tensor(np.concatenate(x),dtype=torch.float32,device=dev)
92 ans={}
93 for reg in [False,True]:
94 seed(); model=AE().to(dev); opt=torch.optim.Adam(model.parameters(),lr=.008)
95 for step in range(500):
96 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)
97 if reg: loss=loss+.12*sym_penalty(torch.softmax(model.logits,0),model.mu,model.cov(),.7,.4,.8)
98 opt.zero_grad(); loss.backward(); opt.step()
99 with torch.no_grad():
100 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())
101 ans['regularized' if reg else 'baseline']={'reconstruction_mse':rec,'latent_mixture_nll':nll,'best_affine_residual':ar[:4]}
102 return ans
103 except Exception as e:
104 return {'error':repr(e)}
105
106if __name__=='__main__':
107 out={'toy':toy_sweep(),'mini':mini()}
108 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))
109 print(json.dumps(out,indent=2,default=lambda o: o.item() if hasattr(o,'item') else list(o)))