Kac-Ward Exact Teacher for Autoregressive Samplers / kw_test.py
Mechanism confirmed, baseline not beaten
1import numpy as np, math
2from itertools import product
3from experiment import lattice
4
5def calc(L,K,mode):
6 e=lattice(L); d=[]
7 for u,v in e:d += [(u,v),(v,u)]
8 m=len(d); T=np.zeros((m,m),complex); pos=lambda x:np.array([x%L,x//L],float)
9 for a,(u,v) in enumerate(d):
10 for b,(v2,w) in enumerate(d):
11 if v2!=v or w==u:continue
12 if mode[0]=='rev': a1=pos(u)-pos(v)
13 else:a1=pos(v)-pos(u)
14 a2=pos(w)-pos(v)
15 th=math.atan2(a1[0]*a2[1]-a1[1]*a2[0],a1.dot(a2))
16 if mode[1]=='half':th/=2
17 k=K[e.index(tuple(sorted((u,v))))]
18 T[a,b]=np.tanh(k)*np.exp(1j*th)
19 sg,ld=np.linalg.slogdet(np.eye(m)-T)
20 return (2**(L*L))*np.prod(np.cosh(K))*np.sqrt(sg*np.exp(ld))
21for mode in [('rev','full'),('rev','half'),('fwd','full'),('fwd','half')]:
22 for L in [2,3]:
23 e=lattice(L); K=np.array([.17+.04*(k%3-1) for k in range(len(e))]); ss=np.array(list(product([-1,1],repeat=L*L)))
24 z=np.exp(np.sum([K[k]*ss[:,u]*ss[:,v] for k,(u,v) in enumerate(e)],0)).sum()
25 print(mode,L,z,calc(L,K,mode),abs(z-calc(L,K,mode))/z)