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

Mechanism confirmed, baseline not beaten

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