Flatness-Calibrated Constant-Step SGD / flat_sgd_experiment.py

Mechanism failed

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