Wick-Matching Polynomial Interaction Layer / wick_experiment.py
Mechanism failed
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))