Flatness-Calibrated Constant-Step SGD / flat_sgd_experiment.py
Mechanism failed
1import json
2import math
3from dataclasses import dataclass
4from pathlib import Path
5import numpy as np
6
7@dataclass
8class SimResult:
9 rms: float
10 tau: float
11 values: np.ndarray
12
13def grad(x, m):
14 return np.sign(x) * abs(x) ** (m - 1)
15
16def simulate(m, alpha, sigma=1.0, n=120000, burn=30000, seed=0):
17 rng = np.random.default_rng(seed)
18 x = 0.0
19 out = np.empty(n - burn)
20 for t in range(n):
21 x -= alpha * (grad(x, m) + sigma * rng.normal())
22 if not np.isfinite(x) or abs(x) > 1e8:
23 return SimResult(float('nan'), float('nan'), np.empty(0))
24 if t >= burn:
25 out[t - burn] = x
26 z = out - out.mean()
27 var = np.mean(z*z)
28 tau = 1.0
29 if var > 0:
30 for lag in range(1, min(4000, len(z)//10)):
31 ac = np.mean(z[:-lag] * z[lag:]) / var
32 if ac <= 0:
33 break
34 tau += 2*ac
35 return SimResult(float(np.sqrt(np.mean(out*out))), float(tau), out)
36
37def slope(x, y):
38 return float(np.polyfit(np.log(x), np.log(y), 1)[0])
39
40def calibrate(target_r, m, sigma=1.0, C=1.0, amin=1e-5, amax=0.25):
41 # stated law in 1D: r = C * alpha^(1/m) * sigma
42 return float(np.clip((target_r/(C*sigma))**m, amin, amax))
43
44def main():
45 # Small alphas keep Euler updates stable and make asymptotic laws visible.
46 ms = [2, 3, 4, 5]
47 alphas = np.array([0.002, 0.004, 0.008, 0.016])
48 rows = []
49 for m in ms:
50 radii, taus = [], []
51 for j, a in enumerate(alphas):
52 r = simulate(m, a, seed=1000 + 10*m + j)
53 radii.append(r.rms); taus.append(r.tau)
54 rows.append({'m':m, 'alpha':float(a), 'radius':r.rms, 'tau':r.tau})
55 rows.append({'m':m, 'radius_slope':slope(alphas, radii),
56 'tau_slope':slope(alphas, taus),
57 'pred_radius_slope':1/m, 'pred_tau_slope':-(m-1)})
58
59 # Empirically calibrate the unknown C at a reference rate, then target a radius.
60 calibrated = []
61 for m in ms:
62 ref_alpha = 0.004
63 ref = simulate(m, ref_alpha, seed=4000+m, n=180000, burn=50000)
64 C_hat = ref.rms / (ref_alpha**(1/m))
65 a = calibrate(0.12, m, C=C_hat)
66 got = simulate(m, a, seed=4500+m, n=180000, burn=50000)
67 calibrated.append({'m':m, 'C_hat':C_hat, 'alpha':a, 'target_radius':0.12,
68 'observed_radius':got.rms,
69 'relative_error':abs(got.rms-0.12)/0.12})
70
71 # Noise prediction: radius is proportional to sigma^(2/m), since v=Sigma=sigma^2.
72 noise_sweep = []
73 for m in [2, 3, 4, 5]:
74 a = 0.004
75 sigmas = np.array([0.5, 1.0, 2.0])
76 radii = []
77 for j, sig in enumerate(sigmas):
78 radii.append(simulate(m, a, sigma=float(sig), seed=6000+10*m+j,
79 n=160000, burn=45000).rms)
80 noise_sweep.append({'m':m, 'observed_slope':slope(sigmas, radii),
81 'predicted_slope':2/m, 'sigmas':sigmas.tolist(),
82 'radii':radii})
83
84 # Adaptation test: infer m from curvature q(rho) for H=|x|^m/m,
85 # where q(rho) proportional to rho^(m-2), then invert the radius law.
86 adapt = []
87 target = 0.12
88 for m in ms:
89 rho = 0.03
90 q1 = rho**(m-2)
91 q2 = (2*rho)**(m-2)
92 mhat = 2 + math.log((q2+1e-12)/(q1+1e-12), 2) if m != 2 else 2.0
93 a = calibrate(target, mhat)
94 rr = simulate(m, a, seed=5000+m, n=140000, burn=40000)
95 adapt.append({'true_m':m, 'mhat':mhat, 'alpha':a, 'target_radius':target,
96 'observed_radius':rr.rms, 'relative_radius_error':abs(rr.rms-target)/target})
97
98 # Same target-radius request: calibrated rate versus quadratic miscalibration.
99 compare = []
100 for m in [3, 4, 5]:
101 ai = calibrate(target, m)
102 aq = target**2
103 ri = simulate(m, ai, seed=8000+m, n=140000, burn=40000).rms
104 rq = simulate(m, aq, seed=9000+m, n=140000, burn=40000).rms
105 compare.append({'m':m, 'idea_alpha':ai, 'idea_radius':ri,
106 'quadratic_alpha':aq, 'quadratic_radius':rq,
107 'idea_error':abs(ri-target), 'quadratic_error':abs(rq-target)})
108
109 result = {'radius_and_mixing_sweep':rows, 'adaptation':adapt,
110 'empirical_C_calibration':calibrated, 'noise_sweep':noise_sweep,
111 'target_radius_comparison':compare,
112 'notes':'RMS stationary radius and integrated autocorrelation time from 1D constant-step SGD.'}
113 Path('results.json').write_text(json.dumps(result, indent=2))
114 print(json.dumps(result, indent=2))
115
116if __name__ == '__main__':
117 main()