Affine-symmetry-free GMM latent prior / experiment.py

Unverified

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