import json, math import numpy as np from experiment import throttle EPS = 1e-8 def transition_sweep(): kappa = 0.4 rows = [] for q in [0.25, 1.0, 4.0, 16.0, 64.0]: predicted = kappa * math.sqrt(q) # Dense grid gives an observed bracket for the first a<1 point. norms = np.linspace(0.1 * predicted, 2.0 * predicted, 1001) gains = [throttle(np.array([x]), q, kappa)[0] for x in norms] observed = float(norms[next(i for i, a in enumerate(gains) if a < 1.0)]) rows.append({"q": q, "predicted_threshold": predicted, "observed_threshold": observed, "relative_error": abs(observed-predicted)/predicted}) return rows def inverse_scaling_sweep(): # For delta>kappa, a = kappa*sqrt(q)/||r||, so ||r||a/sqrt(q)=kappa. kappa = 0.3 rows = [] for q in [0.5, 2.0, 8.0, 32.0]: for residual_norm in [2.0, 5.0, 10.0]: r = np.array([residual_norm, 0.0]) a, delta = throttle(r, q, kappa) observed = residual_norm*a/math.sqrt(q) rows.append({"q": q, "residual_norm": residual_norm, "delta": delta, "gain": a, "predicted_normalized_update": kappa, "observed_normalized_update": observed}) return rows def scalar_stability_sweep(): # For scalar least squares, r=theta, q=1, and theta <- (1-eta*a)theta. # Plain GD has multiplier |1-eta| and diverges for eta>2. Throttle has # a=min(1,kappa/(|theta|+eps)); for large theta its decrement is eta*kappa. kappa = 0.4 rows = [] for eta in [0.5, 1.5, 2.0, 2.5, 3.0, 5.0, 10.0]: plain = 1.0; gated = 1.0 for _ in range(100): plain *= (1.0-eta) a, _ = throttle(np.array([gated]), 1.0, kappa) gated -= eta*a*gated rows.append({"eta": eta, "predicted_plain_multiplier": abs(1-eta), "observed_plain_final_abs": abs(plain), "observed_throttle_final_abs": abs(gated), "observed_throttle_max_abs": abs(gated) if eta <= 2 else None}) return rows def main(): out = {"transition_sweep": transition_sweep(), "inverse_scaling_sweep": inverse_scaling_sweep(), "scalar_stability_sweep": scalar_stability_sweep()} with open("sweep_results.json", "w") as f: json.dump(out, f, indent=2) print(json.dumps(out, indent=2)) if __name__ == '__main__': main()