Cubic-Rate Third-Order Langevin Optimizer / verify_cubic_langevin.py

Failed on benchmark

Raw ⬇ ZIP
 1import json, math
 2from pathlib import Path
 3import numpy as np
 4
 5def cubic_rate(gamma, kappa):
 6    if kappa <= 0: return 0.0
 7    lo, hi = 0.0, kappa**(1/3) + math.sqrt(kappa/gamma) + 1.0
 8    for _ in range(100):
 9        mid=(lo+hi)/2
10        if mid**3+gamma*mid**2 < kappa: lo=mid
11        else: hi=mid
12    return (lo+hi)/2
13
14def third_order_step(x,v,a,grad,dt,gamma,eps=0.,rng=None):
15    a += dt*(-grad-gamma*a)
16    if eps: a += math.sqrt(2*gamma*eps*dt)*rng.normal()
17    v += dt*a; x += dt*v
18    return x,v,a
19
20def measure_growth(kappa,gamma,dt,T=None):
21    r=cubic_rate(gamma,kappa)
22    if T is None: T=max(60.,20./max(r,1e-5))
23    x,v,a=1e-8,0.,0.; n=int(T/dt); ts=[]; ys=[]
24    for i in range(n):
25        x,v,a=third_order_step(x,v,a,-kappa*x,dt,gamma)
26        if i>n*.45 and np.isfinite(x) and 1e-14<abs(x)<1e50:
27            ts.append(i*dt); ys.append(math.log(abs(x)))
28    if len(ts)<20: return float('nan')
29    return float(np.polyfit(ts,ys,1)[0])
30
31def double_well_escape(method,gamma=1.,dt=.01,eps=.08,trials=120,T=35.,seed=4):
32    rng=np.random.default_rng(seed); hits=[]
33    for _ in range(trials):
34        x=-1.;v=0.;a=0.;hit=None
35        for j in range(int(T/dt)):
36            grad=4*x*(x*x-1)
37            if method=='third': x,v,a=third_order_step(x,v,a,grad,dt,gamma,eps,rng)
38            elif method=='momentum':
39                v=gamma*v-dt*grad+math.sqrt(2*eps*dt)*rng.normal(); x+=v
40            else: x-=dt*grad-math.sqrt(2*eps*dt)*rng.normal()
41            if x>0: hit=(j+1)*dt; break
42        if hit is not None: hits.append(hit)
43    return {'hit_fraction':len(hits)/trials,'median_hit':float(np.median(hits)) if hits else None}
44
45def main():
46    # Prediction 1: asymptotic saddle growth is the positive cubic root.
47    gamma,kappa,dt=1.3,2.7,.001
48    pred=cubic_rate(gamma,kappa); obs=measure_growth(kappa,gamma,dt,T=30.)
49    p1={'gamma':gamma,'kappa':kappa,'dt':dt,'predicted_rate':pred,'observed_rate':obs,'relative_error':abs(obs-pred)/pred}
50
51    # Prediction 2: weak damping gives r proportional to kappa^(1/3).
52    gamma=.03; kappas=np.array([.125,.25,.5,1.,2.,4.])
53    pred=np.array([cubic_rate(gamma,k) for k in kappas])
54    obs=np.array([measure_growth(k,gamma,.002) for k in kappas])
55    p2={'gamma':gamma,'kappas':kappas.tolist(),'predicted_rates':pred.tolist(),'observed_rates':obs.tolist(),
56        'predicted_loglog_slope':float(np.polyfit(np.log(kappas),np.log(pred),1)[0]),
57        'observed_loglog_slope':float(np.polyfit(np.log(kappas),np.log(obs),1)[0]),
58        'relative_rate_errors':(abs(obs-pred)/pred).tolist()}
59
60    # Prediction 3: r is zero at zero negative curvature and increases monotonically with kappa.
61    ks=np.array([0.,1e-5,1e-3,.01,.1,1.])
62    rows=[]
63    for k in ks:
64        r=cubic_rate(1.,k); measured=0. if k==0 else measure_growth(k,1.,.01)
65        rows.append({'kappa':float(k),'predicted_rate':r,'observed_rate':measured})
66    p3={'rows':rows,'predicted_monotone':bool(np.all(np.diff([x['predicted_rate'] for x in rows])>=0)),
67        'observed_monotone':bool(np.all(np.diff([x['observed_rate'] for x in rows])>=0))}
68
69    # Prediction 4 (integration heuristic): dt*r=.2 should retain small rate distortion.
70    gamma,kappa=1.,4.; r=cubic_rate(gamma,kappa); rows=[]
71    for dt in [.01,.05,.1,.15,.2,.3,.5]:
72        o=measure_growth(kappa,gamma,dt,T=20.)
73        rows.append({'dt':dt,'dt_times_rate':dt*r,'observed_rate':o,'relative_error':abs(o-r)/r})
74    p4={'rate':r,'recommended_dt_bound':.2/r,'rows':rows}
75
76    bench={m:double_well_escape(m,seed=10+i) for i,m in enumerate(['third','momentum','overdamped'])}
77    result={'prediction_checks':{'rate_at_fixed_params':p1,'curvature_scaling':p2,'curvature_transition':p3,'integration_heuristic':p4},'double_well_escape':bench}
78    Path('results.json').write_text(json.dumps(result,indent=2)); print(json.dumps(result,indent=2))
79if __name__=='__main__': main()