Residual-to-State Update Throttle / sweeps.py

Failed on benchmark

Raw ⬇ ZIP
 1import json, math
 2import numpy as np
 3from experiment import throttle
 4
 5EPS = 1e-8
 6
 7def transition_sweep():
 8    kappa = 0.4
 9    rows = []
10    for q in [0.25, 1.0, 4.0, 16.0, 64.0]:
11        predicted = kappa * math.sqrt(q)
12        # Dense grid gives an observed bracket for the first a<1 point.
13        norms = np.linspace(0.1 * predicted, 2.0 * predicted, 1001)
14        gains = [throttle(np.array([x]), q, kappa)[0] for x in norms]
15        observed = float(norms[next(i for i, a in enumerate(gains) if a < 1.0)])
16        rows.append({"q": q, "predicted_threshold": predicted,
17                     "observed_threshold": observed,
18                     "relative_error": abs(observed-predicted)/predicted})
19    return rows
20
21def inverse_scaling_sweep():
22    # For delta>kappa, a = kappa*sqrt(q)/||r||, so ||r||a/sqrt(q)=kappa.
23    kappa = 0.3
24    rows = []
25    for q in [0.5, 2.0, 8.0, 32.0]:
26        for residual_norm in [2.0, 5.0, 10.0]:
27            r = np.array([residual_norm, 0.0])
28            a, delta = throttle(r, q, kappa)
29            observed = residual_norm*a/math.sqrt(q)
30            rows.append({"q": q, "residual_norm": residual_norm,
31                         "delta": delta, "gain": a,
32                         "predicted_normalized_update": kappa,
33                         "observed_normalized_update": observed})
34    return rows
35
36def scalar_stability_sweep():
37    # For scalar least squares, r=theta, q=1, and theta <- (1-eta*a)theta.
38    # Plain GD has multiplier |1-eta| and diverges for eta>2. Throttle has
39    # a=min(1,kappa/(|theta|+eps)); for large theta its decrement is eta*kappa.
40    kappa = 0.4
41    rows = []
42    for eta in [0.5, 1.5, 2.0, 2.5, 3.0, 5.0, 10.0]:
43        plain = 1.0; gated = 1.0
44        for _ in range(100):
45            plain *= (1.0-eta)
46            a, _ = throttle(np.array([gated]), 1.0, kappa)
47            gated -= eta*a*gated
48        rows.append({"eta": eta,
49                     "predicted_plain_multiplier": abs(1-eta),
50                     "observed_plain_final_abs": abs(plain),
51                     "observed_throttle_final_abs": abs(gated),
52                     "observed_throttle_max_abs": abs(gated) if eta <= 2 else None})
53    return rows
54
55def main():
56    out = {"transition_sweep": transition_sweep(),
57           "inverse_scaling_sweep": inverse_scaling_sweep(),
58           "scalar_stability_sweep": scalar_stability_sweep()}
59    with open("sweep_results.json", "w") as f: json.dump(out, f, indent=2)
60    print(json.dumps(out, indent=2))
61
62if __name__ == '__main__': main()