Tau-leaped parallel discrete Hamiltonian sampler / tau_leap_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, math, time
  2import numpy as np
  3
  4SEED = 2485
  5D, K = 20, 3
  6rng0 = np.random.default_rng(SEED)
  7one = rng0.normal(0, 0.7, size=(D, K))
  8pair = rng0.normal(0, 0.06, size=(D-1, K, K))
  9
 10def energy(x):
 11    v = float(np.sum(one[np.arange(D), x]))
 12    for i in range(D-1): v += pair[i, x[i], x[i+1]]
 13    return v
 14
 15def rates(x, p):
 16    u0 = energy(x); q = np.zeros((D, K))
 17    for i in range(D):
 18        for z in range(K):
 19            if z != x[i]:
 20                y=x.copy(); y[i]=z
 21                q[i,z]=max(float(p[i]),0.0)*math.exp(-0.5*(energy(y)-u0))
 22    return q
 23
 24def exact_path(x, p, rng, horizon):
 25    t=0.; n=0
 26    while t < horizon:
 27        q=rates(x,p); lam=float(q.sum())
 28        if lam <= 0: break
 29        t += rng.exponential(1./lam)
 30        if t > horizon: break
 31        flat=int(rng.choice(D*K,p=(q.ravel()/lam)))
 32        i,z=divmod(flat,K); x[i]=z; n+=1
 33    return n
 34
 35def tau_path(x, p, rng, horizon, h):
 36    n=0; t=0.
 37    while t < horizon-1e-12:
 38        step=min(h,horizon-t); q=rates(x,p)
 39        active=rng.random(q.shape) < (-np.expm1(-step*q))
 40        for i in range(D):
 41            zs=np.flatnonzero(active[i])
 42            if len(zs):
 43                w=q[i,zs]; x[i]=int(zs[rng.choice(len(zs),p=w/w.sum())]); n+=1
 44        t += step
 45    return n
 46
 47def make_state(rng): return rng.integers(0,K,size=D), rng.normal(0,1,size=D)
 48
 49def poisson_checks():
 50    q=0.73; hvals=[0.05,0.2,0.8,1.5]; N=300000; rows=[]
 51    for h in hvals:
 52        s=np.random.default_rng(SEED+int(100*h)).poisson(h*q,N)
 53        rows.append({'h':h,'mean_obs':float(s.mean()),'mean_pred':h*q,
 54                     'nonzero_obs':float(np.mean(s>0)),'nonzero_pred':float(1-math.exp(-h*q))})
 55    return rows
 56
 57def bias_sweep():
 58    # Compare endpoint distributions to exact trajectories using common initial states.
 59    hs=[0.01,0.025,0.05,0.1,0.2]
 60    M=120; horizon=.25; out=[]
 61    for h in hs:
 62        exact=np.zeros((M,D),int); approx=np.zeros((M,D),int); ne=na=0
 63        for m in range(M):
 64            a,b=make_state(np.random.default_rng(SEED+m));
 65            xe=a.copy(); xa=a.copy();
 66            ne+=exact_path(xe,b,np.random.default_rng(10000+m),horizon)
 67            na+=tau_path(xa,b,np.random.default_rng(20000+m),horizon,h)
 68            exact[m]=xe; approx[m]=xa
 69        # Total variation of coordinate-0 marginals and mismatch rate.
 70        tv=0.
 71        for z in range(K):
 72            tv += abs(np.mean(exact[:,0]==z)-np.mean(approx[:,0]==z))
 73        tv*=.5
 74        out.append({'h':h,'coord0_TV':float(tv),'endpoint_mismatch':float(np.mean(np.any(exact!=approx,axis=1))),
 75                    'exact_events_per_path':ne/M,'tau_events_per_path':na/M})
 76    return out
 77
 78def throughput():
 79    M=100; horizon=.25; h=.05
 80    states=[]; moms=[]
 81    for m in range(M):
 82        x,p=make_state(np.random.default_rng(50000+m)); states.append(x); moms.append(p)
 83    t=time.perf_counter(); en=0
 84    for m in range(M): en+=exact_path(states[m].copy(),moms[m],np.random.default_rng(60000+m),horizon)
 85    te=time.perf_counter()-t
 86    t=time.perf_counter(); tn=0
 87    for m in range(M): tn+=tau_path(states[m].copy(),moms[m],np.random.default_rng(70000+m),horizon,h)
 88    tt=time.perf_counter()-t
 89    return {'M':M,'h':h,'exact_sec':te,'tau_sec':tt,'exact_paths_per_sec':M/te,'tau_paths_per_sec':M/tt,'speedup':te/tt,'exact_events':en/M,'tau_events':tn/M}
 90
 91def main():
 92    result={'seed':SEED,'dimensions':[D,K],'poisson_checks':poisson_checks(),
 93            'bias_sweep':bias_sweep(),'throughput':throughput()}
 94    # Fit log-log observed bias scaling, excluding numerical floor if necessary.
 95    hs=np.array([r['h'] for r in result['bias_sweep']]); tv=np.array([r['coord0_TV'] for r in result['bias_sweep']])
 96    result['loglog_bias_slope']=float(np.polyfit(np.log(hs),np.log(np.maximum(tv,1e-8)),1)[0])
 97    with open('results.json','w') as f: json.dump(result,f,indent=2)
 98    print(json.dumps(result,indent=2))
 99
100if __name__=='__main__': main()