Kac-Ward Exact Teacher for Autoregressive Samplers / kw_test.py

Mechanism confirmed, baseline not beaten

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