Kac-Ward Exact Teacher for Autoregressive Samplers / experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, time
2from itertools import product
3import numpy as np
4
5
6def lattice(L):
7 edges=[]
8 for r in range(L):
9 for c in range(L):
10 u=r*L+c
11 if c+1<L: edges.append((u,u+1))
12 if r+1<L: edges.append((u,u+L))
13 return edges
14
15
16def energies(states, edges, J):
17 # states: [S,N], values +/-1
18 return -sum(J[k]*states[:,u]*states[:,v] for k,(u,v) in enumerate(edges))
19
20
21def exact_distribution(L, beta, J):
22 N=L*L
23 states=np.array(list(product((-1,1), repeat=N)), dtype=np.int8)
24 logw=-beta*energies(states, lattice(L), J)
25 logw-=logw.max(); w=np.exp(logw); p=w/w.sum()
26 return states,p
27
28
29def conditional_oracle(states, p, i, prefix):
30 mask=np.all(states[:,:i]==np.asarray(prefix)[None,:],axis=1)
31 vals=states[mask,i]; ww=p[mask]
32 zplus=ww[vals==1].sum(); zminus=ww[vals==-1].sum()
33 return float(zplus/(zplus+zminus))
34
35
36def all_prefix_data(states,p):
37 X=[]; I=[]; Q=[]
38 N=states.shape[1]
39 # Prefix encoding: known spins, zero for unknown; position one-hot is appended.
40 for s,ps in zip(states,p):
41 # weighted selection is not needed: include every state and position,
42 # with exact q at its realized prefix; duplicate prefixes are harmless.
43 for i in range(N):
44 pref=s[:i]
45 q=conditional_oracle(states,p,i,pref)
46 x=np.zeros(N,dtype=np.float32); x[:i]=pref
47 X.append(x); I.append(i); Q.append(q)
48 return np.asarray(X),np.asarray(I),np.asarray(Q)
49
50
51def kw_partition(L, K):
52 """Kac-Ward determinant for a square lattice embedded at integer coordinates."""
53 edges=lattice(L); und=[]
54 for u,v in edges: und += [(u,v),(v,u)]
55 m=len(und); T=np.zeros((m,m),dtype=np.complex128)
56 xy=lambda x:(x//L,x%L)
57 for a,(u,v) in enumerate(und):
58 ru,cu=xy(u); rv,cv=xy(v)
59 vin=np.array([cv-cu,rv-ru],float) # incoming direction u -> v
60 for b,(v2,w) in enumerate(und):
61 if v2!=v or w==u: continue
62 rw,cw=xy(w); vout=np.array([cw-cv,rw-rv],float)
63 cross=vin[0]*vout[1]-vin[1]*vout[0]
64 dot=vin.dot(vout)
65 theta=0.5*math.atan2(cross,dot)
66 # edge coupling belongs to the undirected edge
67 kk=K[edges.index(tuple(sorted((u,v))))]
68 T[a,b]=math.tanh(kk)*np.exp(1j*theta)
69 sign,ld=np.linalg.slogdet(np.eye(m,dtype=complex)-T)
70 # For planar ferromagnetic examples this branch is positive real.
71 z=(2.0**(L*L))*np.prod(np.cosh(K))*np.sqrt(sign*np.exp(ld))
72 return float(np.real_if_close(z).real)
73
74
75def kw_check():
76 rows=[]
77 for L in (2,3):
78 edges=lattice(L); K=np.array([0.17+0.04*((k%3)-1) for k in range(len(edges))])
79 states=np.array(list(product((-1,1),repeat=L*L)),dtype=np.int8)
80 z_enum=float(np.exp(np.logaddexp.reduce(sum(K[k]*states[:,u]*states[:,v] for k,(u,v) in enumerate(edges)))))
81 z_kw=kw_partition(L,K)
82 rows.append((L,z_enum,z_kw,abs(z_enum-z_kw)/z_enum))
83 return rows
84
85
86def train_model(X,I,Q, seed, soft, beta, steps=300, batch=128):
87 import torch
88 torch.manual_seed(seed); np.random.seed(seed)
89 dev='cuda' if torch.cuda.is_available() else 'cpu'
90 try:
91 x=torch.tensor(X); ii=torch.nn.functional.one_hot(torch.tensor(I),X.shape[1]).float();
92 inp=torch.cat([x,ii],1).to(dev); y=torch.tensor(Q,dtype=torch.float32).to(dev)
93 rng=np.random.default_rng(seed+19); yy=torch.tensor((rng.random(len(Q))<Q).astype(np.float32)).to(dev)
94 model=torch.nn.Sequential(torch.nn.Linear(inp.shape[1],48),torch.nn.Tanh(),torch.nn.Linear(48,1)).to(dev)
95 opt=torch.optim.Adam(model.parameters(),lr=.025)
96 rng=np.random.default_rng(seed+7)
97 for _ in range(steps):
98 ix=torch.tensor(rng.integers(0,len(Q),size=batch),device=dev)
99 logit=model(inp[ix]).squeeze(1); target= y[ix] if soft else yy[ix]
100 loss=torch.nn.functional.binary_cross_entropy_with_logits(logit,target)
101 opt.zero_grad(); loss.backward(); opt.step()
102 with torch.no_grad():
103 pred=torch.sigmoid(model(inp).squeeze(1)).cpu().numpy()
104 eps=1e-7
105 kl=Q*np.log((Q+eps)/(pred+eps))+(1-Q)*np.log((1-Q+eps)/(1-pred+eps))
106 ce=-(Q*np.log(pred+eps)+(1-Q)*np.log(1-pred+eps))
107 return float(np.mean(kl)),float(np.mean(ce)),dev
108 except Exception as e:
109 # retry CPU, as required for shared-GPU failures
110 torch.manual_seed(seed); dev='cpu'
111 inp=torch.tensor(np.concatenate([X,np.eye(X.shape[1])[I]],1),dtype=torch.float32)
112 y=torch.tensor(Q,dtype=torch.float32); yy=torch.tensor((np.random.default_rng(seed+19).random(len(Q))<Q).astype(np.float32))
113 model=torch.nn.Sequential(torch.nn.Linear(inp.shape[1],48),torch.nn.Tanh(),torch.nn.Linear(48,1)); opt=torch.optim.Adam(model.parameters(),lr=.025)
114 rng=np.random.default_rng(seed+7)
115 for _ in range(steps):
116 ix=torch.tensor(rng.integers(0,len(Q),size=batch)); logit=model(inp[ix]).squeeze(1); target=y[ix] if soft else yy[ix]
117 loss=torch.nn.functional.binary_cross_entropy_with_logits(logit,target); opt.zero_grad(); loss.backward(); opt.step()
118 with torch.no_grad(): pred=torch.sigmoid(model(inp).squeeze(1)).numpy()
119 eps=1e-7; kl=Q*np.log((Q+eps)/(pred+eps))+(1-Q)*np.log((1-Q+eps)/(1-pred+eps)); ce=-(Q*np.log(pred+eps)+(1-Q)*np.log(1-pred+eps))
120 return float(np.mean(kl)),float(np.mean(ce)),dev
121
122
123def main():
124 kw=kw_check()
125 # Quantitative prediction 1: KW equals enumeration to numerical precision.
126 identity=[]
127 for beta in (.2,.7,1.3):
128 L=3; edges=lattice(L); J=np.array([1 if k%2 else -1 for k in range(len(edges))],float)
129 s,p=exact_distribution(L,beta,J); X,I,Q=all_prefix_data(s,p)
130 # Conditional CE - entropy equals KL, checked with a deliberately perturbed model.
131 pred=np.clip(.15+.7*Q,.001,.999)
132 ce=np.mean(-(Q*np.log(pred)+(1-Q)*np.log(1-pred)))
133 ent=np.mean(-(Q*np.log(Q)+(1-Q)*np.log(1-Q)))
134 kl=np.mean(Q*np.log(Q/pred)+(1-Q)*np.log((1-Q)/(1-pred)))
135 # Prediction 2: sampled-label variance q(1-q), increasing near beta=0.
136 var=float(np.mean(Q*(1-Q)))
137 identity.append({'beta':beta,'identity_abs_err':abs((ce-ent)-kl),'mean_label_variance':var,'mean_q':float(Q.mean())})
138 # Fixed tiny setup, same data and training budget.
139 beta=.7; L=3; edges=lattice(L); J=np.array([1 if k%2 else -1 for k in range(len(edges))],float)
140 s,p=exact_distribution(L,beta,J); X,I,Q=all_prefix_data(s,p)
141 results=[]
142 for b in (.2,.7,1.3):
143 sb, pb=exact_distribution(L,b,np.ones(len(edges)))
144 xb,ib,qb=all_prefix_data(sb,pb)
145 out=[]
146 for soft in (False,True):
147 vals=[train_model(xb,ib,qb,seed=11+r,soft=soft,beta=b)[:2] for r in range(3)]
148 out.append({'method':'soft_oracle' if soft else 'sampled_labels','kl_mean':float(np.mean([v[0] for v in vals])),'ce_mean':float(np.mean([v[1] for v in vals]))})
149 results.append({'beta':b,'mean_q_variance':float(np.mean(qb*(1-qb))),'methods':out})
150 # Prediction 3: exact enumeration grows 2^N while autoregressive pass is N outputs.
151 timing=[]
152 for L in (2,3,4):
153 t=time.perf_counter(); exact_distribution(L,.5,np.ones(len(lattice(L)))); sec=time.perf_counter()-t
154 timing.append({'L':L,'N':L*L,'enumeration_seconds':sec,'autoregressive_outputs':L*L})
155 report={'kw_check':kw,'mechanism_checks':identity,'training_sweep':results,'scaling':timing}
156 with open('results.json','w') as f: json.dump(report,f,indent=2)
157 print(json.dumps(report,indent=2))
158
159if __name__=='__main__': main()