State-Dependent Temperature Langevin / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import json, math, time
 2from pathlib import Path
 3import numpy as np
 4
 5SEED = 2845
 6NU = 5.0
 7EPS = 2e-3
 8
 9# Student-t target: U(x)=(nu+1)/2 log(1+x^2/nu), so Var[X]=nu/(nu-2).
10def U(x):
11    return 0.5*(NU+1.0)*np.log1p(x*x/NU)
12
13def grad_U(x):
14    return (NU+1.0)*x/(NU+x*x)
15
16def sigma(x, alpha):
17    r = np.sqrt(x*x + 1e-8)
18    return 1.0 + alpha*np.log1p(r)
19
20def dsigma2(x, alpha):
21    r = np.sqrt(x*x + 1e-8)
22    s = 1.0 + alpha*np.log1p(r)
23    return 2.0*s*alpha*x/(r*(1.0+r))
24
25def drift(x, alpha, correction=True):
26    s = sigma(x, alpha)
27    return (dsigma2(x, alpha) if correction else 0.0) - s*s*grad_U(x)
28
29def current_residual(x, alpha):
30    # J=b*pi-d(a*pi)/dx; analytic derivative gives d(a*pi)=(a'-a U')pi.
31    pi=np.exp(-U(x)); a=sigma(x,alpha)**2
32    return drift(x,alpha,True)*pi - (dsigma2(x,alpha)-a*grad_U(x))*pi
33
34def run(alpha, correction=True, nsteps=50000, nchains=128, burn=8000, seed=0):
35    rng=np.random.default_rng(seed)
36    x=rng.normal(0,2,nchains)
37    samples=[]; grad_calls=0
38    for k in range(nsteps):
39        s=sigma(x,alpha)
40        x=x+EPS*drift(x,alpha,correction)+np.sqrt(2*EPS)*s*rng.normal(size=nchains)
41        grad_calls += nchains
42        if k>=burn and k%5==0: samples.append(x.copy())
43    z=np.concatenate(samples)
44    # Batch means gives a conservative ESS estimate.
45    flat=z.reshape(-1); m=100
46    nb=len(flat)//m
47    bm=flat[:nb*m].reshape(nb,m).mean(1)
48    var=np.var(flat); ess=min(len(flat), nb*m*var/(m*np.var(bm)+1e-30))
49    return {'mean':float(np.mean(flat)), 'second':float(np.mean(flat**2)),
50            'q95':float(np.quantile(flat,.95)), 'ess':float(ess),
51            'ess_per_grad':float(ess/grad_calls), 'max_abs':float(np.max(np.abs(flat)))}
52
53def main():
54    # Prediction 1: divergence correction makes stationary current identically zero.
55    xs=np.linspace(-12,12,10001)
56    current={str(a):float(np.max(np.abs(current_residual(xs,a)))) for a in [0,.25,.5,1.0,2.0]}
57
58    # Prediction 2: sigma(1,alpha)-1 is exactly linear in alpha.
59    alphas=np.array([0,.25,.5,1.,2.])
60    sig1=np.array([sigma(np.array(1.),a) for a in alphas])
61    slope=float(np.polyfit(alphas,sig1,1)[0])
62    linear_maxerr=float(np.max(np.abs(sig1-(1+alphas*np.log(2)))))
63
64    # Prediction 3: correction magnitude is alpha*(1+alpha*c), hence quadratic
65    # coefficient is predicted from analytic expansion at x=2.
66    x0=2.; h=1e-4
67    vals=[]
68    for a in alphas:
69        vals.append(abs(dsigma2(np.array(x0),a)))
70    vals=np.array(vals)
71    fit=np.polyfit(alphas,vals,2)
72    c=np.log1p(abs(x0)); predicted_quad=2*c*(x0/(abs(x0)*(1+abs(x0))))
73    # Direct finite-difference confirmation of the derivative correction.
74    fd=[]
75    for a in [.25,.5,1.0]:
76        f=lambda q:sigma(np.array(q),a)**2
77        fd.append(abs((f(x0+h)-f(x0-h))/(2*h)-dsigma2(np.array(x0),a)))
78    fd_max=float(max(fd))
79
80    # Mechanism sweep and baseline: same ULA update, with/without correction.
81    rows=[]
82    for a in [0.,.5,1.0,2.0]:
83        idea=run(a,True,seed=SEED+int(100*a))
84        base=run(a,False,seed=SEED+int(100*a))
85        rows.append({'alpha':a,'corrected':idea,'uncorrected':base})
86
87    out={'target_variance':NU/(NU-2), 'predictions':{
88        'zero_current_max_abs':current,
89        'linear_sigma_at_x1':{'predicted_slope':math.log(2),'observed_slope':slope,'max_error':linear_maxerr},
90        'quadratic_correction_at_x2':{'fit_coefficients_constant_linear_quadratic':fit.tolist(), 'predicted_quadratic_coefficient':predicted_quad, 'fd_max_error':fd_max}},
91        'sweep':rows}
92    Path('results.json').write_text(json.dumps(out,indent=2))
93    print(json.dumps(out,indent=2))
94
95if __name__=='__main__': main()