Wick-Matching Polynomial Interaction Layer / wick_experiment.py

Mechanism failed

Raw ⬇ ZIP
  1import itertools, math, json, random
  2import numpy as np
  3import torch
  4from torch import nn
  5
  6SEED=180
  7np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
  8
  9def partitions(n, mx=None):
 10    if n==0: return [()]
 11    if mx is None or mx>n: mx=n
 12    out=[]
 13    for k in range(mx,0,-1):
 14        for rest in partitions(n-k,k): out.append((k,)+rest)
 15    return out
 16
 17def matchings(items):
 18    items=list(items)
 19    if not items: yield (); return
 20    a=items[0]
 21    for j in range(1,len(items)):
 22        b=items[j]
 23        for rest in matchings(items[1:j]+items[j+1:]): yield ((a,b),)+rest
 24
 25def delta_for_partition(lam,n):
 26    pairs=[]; pos=0
 27    for k in lam:
 28        cyc=list(range(pos,pos+k)); pos+=k
 29        for j,i in enumerate(cyc): pairs.append((i,n+cyc[(j+1)%k]))
 30    return tuple(pairs)
 31
 32def epsilon(n): return tuple((i,n+i) for i in range(n))
 33
 34def component_type(m1,m2,n):
 35    adj=[[] for _ in range(2*n)]
 36    for a,b in list(m1)+list(m2): adj[a].append(b); adj[b].append(a)
 37    seen=set(); sizes=[]
 38    for s in range(2*n):
 39        if s not in seen:
 40            stack=[s]; seen.add(s); c=0
 41            while stack:
 42                u=stack.pop(); c+=1
 43                for v in adj[u]:
 44                    if v not in seen: seen.add(v); stack.append(v)
 45            sizes.append(c//2)
 46    return tuple(sorted(sizes,reverse=True))
 47
 48def coefficient_table(n):
 49    ps=partitions(n); idx={p:i for i,p in enumerate(ps)}
 50    tab=np.zeros((len(ps),len(ps),len(ps)),dtype=np.int64)
 51    eps=epsilon(n); ms=list(matchings(range(2*n)))
 52    for li,lam in enumerate(ps):
 53        dl=delta_for_partition(lam,n)
 54        for d in ms:
 55            tab[li,idx[component_type(d,dl,n)],idx[component_type(d,eps,n)]]+=1
 56    return ps,tab,len(ms)
 57
 58def traces_products(M, ps):
 59    vals={}; P=torch.eye(M.shape[-1],device=M.device,dtype=M.dtype)
 60    for k in range(1,max(max(p) for p in ps)+1):
 61        P=P@M; vals[k]=torch.diagonal(P,dim1=-2,dim2=-1).sum(-1)
 62    return torch.stack([torch.prod(torch.stack([vals[k] for k in p],-1),-1) for p in ps],-1)
 63
 64class WickLayer(nn.Module):
 65    def __init__(self,d,q,n,hidden=24):
 66        super().__init__(); self.q=q
 67        self.a=nn.Sequential(nn.Linear(d,hidden),nn.Tanh(),nn.Linear(hidden,q*q))
 68        self.b=nn.Sequential(nn.Linear(d,hidden),nn.Tanh(),nn.Linear(hidden,q*q))
 69        ps,tab,_=coefficient_table(n); self.ps=ps
 70        self.register_buffer('coef',torch.tensor(tab,dtype=torch.float32))
 71        self.head=nn.Sequential(nn.LayerNorm(len(ps)),nn.Linear(len(ps),16),nn.Tanh(),nn.Linear(16,1))
 72    def forward(self,h):
 73        B=h.shape[0]; q=self.q
 74        U=self.a(h).reshape(B,q,q); V=self.b(h).reshape(B,q,q)
 75        eye=torch.eye(q,device=h.device,dtype=h.dtype)
 76        A=U@U.transpose(-1,-2)+.1*eye; C=V@V.transpose(-1,-2)+.1*eye
 77        pa=traces_products(A,self.ps); pb=traces_products(C,self.ps)
 78        F=torch.einsum('lmn,bm,bn->bl',self.coef,pa,pb)
 79        return self.head(F).squeeze(-1)
 80
 81class DeepSets(nn.Module):
 82    def __init__(self,d,width=32):
 83        super().__init__()
 84        self.phi=nn.Sequential(nn.Linear(d,width),nn.Tanh(),nn.Linear(width,width),nn.Tanh())
 85        self.rho=nn.Sequential(nn.Linear(width,16),nn.Tanh(),nn.Linear(16,1))
 86    def forward(self,x): return self.rho(self.phi(x).mean(1)).squeeze(-1)
 87
 88class WickModel(nn.Module):
 89    def __init__(self,d):
 90        super().__init__(); self.enc=nn.Sequential(nn.Linear(d,24),nn.Tanh(),nn.Linear(24,16),nn.Tanh()); self.w=WickLayer(16,8,3)
 91    def forward(self,x): return self.w(self.enc(x).mean(1))
 92
 93def pmat(M,p): return np.prod([np.trace(np.linalg.matrix_power(M,k)) for k in p])
 94
 95def math_check():
 96    n=3; ps,tab,nm=coefficient_table(n)
 97    A=np.array([[1.2,.2,-.1],[.2,.8,.15],[-.1,.15,1.1]])
 98    B=np.array([[.9,-.1,.2],[-.1,1.3,.05],[.2,.05,.7]])
 99    exact=np.array([sum(tab[li,mi,ni]*pmat(A,mu)*pmat(B,nu) for mi,mu in enumerate(ps) for ni,nu in enumerate(ps)) for li in range(len(ps))])
100    rng=np.random.default_rng(SEED); vals=[]
101    for _ in range(30000):
102        Z=rng.standard_normal((3,3)); M=A@Z@[email protected]
103        vals.append([pmat(M,p) for p in ps])
104    mc=np.mean(vals,0); rel=np.max(np.abs(mc-exact)/(1+np.abs(exact)))
105    return {'degree':n,'partitions':ps,'matchings':nm,'exact':exact.tolist(),'mc':mc.tolist(),'max_relative_error':float(rel)}
106
107def train(model,x,y,steps,device):
108    model.to(device); opt=torch.optim.Adam(model.parameters(),lr=3e-3); lossfn=nn.MSELoss(); model.train()
109    for _ in range(steps):
110        opt.zero_grad(); loss=lossfn(model(x),y); loss.backward(); opt.step()
111    return model
112
113def benchmark():
114    rng=np.random.default_rng(SEED); N=600; sets=5; d=4
115    x=rng.normal(size=(N,sets,d)).astype('float32'); s=x.sum(1)
116    target=(.7*s[:,0]**2+.35*s[:,1]*s[:,2]+.25*np.sum(x[:,:,3]**2,1)+.15*np.prod(s[:,:3],1)).astype('float32')
117    target += .03*rng.normal(size=N).astype('float32')
118    perm=rng.permutation(N); tr=perm[:60]; te=perm[60:]
119    device='cuda' if torch.cuda.is_available() else 'cpu'
120    try:
121        xt=torch.tensor(x); yt=torch.tensor(target); trainx=xt[tr].to(device); trainy=yt[tr].to(device); testx=xt[te].to(device); testy=yt[te].to(device)
122        torch.manual_seed(SEED); base=train(DeepSets(d,32),trainx,trainy,700,device)
123        torch.manual_seed(SEED); idea=train(WickModel(d),trainx,trainy,700,device)
124    except Exception as e:
125        if device!='cuda': raise
126        device='cpu'; xt=torch.tensor(x); yt=torch.tensor(target); trainx=xt[tr]; trainy=yt[tr]; testx=xt[te]; testy=yt[te]
127        torch.manual_seed(SEED); base=train(DeepSets(d,32),trainx,trainy,700,device)
128        torch.manual_seed(SEED); idea=train(WickModel(d),trainx,trainy,700,device)
129    base.eval(); idea.eval()
130    with torch.no_grad():
131        bm=float(torch.mean((base(testx)-testy)**2).sqrt()); im=float(torch.mean((idea(testx)-testy)**2).sqrt())
132        xx=testx[:8]; p=torch.randperm(sets,device=testx.device)
133        bi=float(torch.max(torch.abs(base(xx)-base(xx[:,p])))); ii=float(torch.max(torch.abs(idea(xx)-idea(xx[:,p]))))
134    return {'device':device,'train_examples':len(tr),'test_rmse_deepsets':bm,'test_rmse_wick':im,'perm_diff_deepsets':bi,'perm_diff_wick':ii,'params_deepsets':sum(p.numel() for p in base.parameters()),'params_wick':sum(p.numel() for p in idea.parameters())}
135
136if __name__=='__main__': print(json.dumps({'math':math_check(),'benchmark':benchmark()},indent=2))