Centered Heavy-Tail Clipping Optimizer / experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json
  2import math
  3from pathlib import Path
  4import numpy as np
  5
  6# Centered Heavy-Tail Clipping Optimizer: small reproducible numerical MVP.
  7# The per-example gradients are g_i = grad f(x) + heavy-tailed noise.
  8
  9def clip_centered(g, c, tau):
 10    r = g - c
 11    n = np.linalg.norm(r, axis=-1, keepdims=True)
 12    return c + r * np.minimum(1.0, tau / (n + 1e-12))
 13
 14def coordinate_median(g):
 15    return np.median(g, axis=0)
 16
 17def grad_bias_scaling(seed=0, alpha=1.5):
 18    """Check E||noise 1_{||noise||>tau}|| decreases approximately tau^(1-alpha)."""
 19    rng = np.random.default_rng(seed)
 20    # scalar symmetric Pareto tails make the asserted exponent particularly clear.
 21    n = 2_000_000
 22    signs = rng.choice([-1.0, 1.0], n)
 23    u = rng.random(n)
 24    z = signs * u ** (-1.0 / alpha)  # P(|Z|>t)=t^-alpha, t>=1
 25    taus = np.array([2., 4., 8., 16., 32., 64.])
 26    vals = np.array([np.mean(np.abs(z) * (np.abs(z) > t)) for t in taus])
 27    slope = np.polyfit(np.log(taus), np.log(vals), 1)[0]
 28    # Formula C_tau is checked separately on random vectors.
 29    return {"taus": taus.tolist(), "tail_bias": vals.tolist(),
 30            "loglog_slope": float(slope), "expected_slope": float(1-alpha)}
 31
 32def operator_check(seed=1):
 33    rng = np.random.default_rng(seed)
 34    g = rng.normal(size=(1000, 7)); c = rng.normal(size=7); tau = 0.73
 35    h = clip_centered(g, c, tau)
 36    residual = np.linalg.norm(h-c, axis=1)
 37    raw_residual = np.linalg.norm(g-c, axis=1)
 38    # unchanged inside ball, bounded outside, and direction preserved (up to roundoff)
 39    inside_err = np.max(np.abs(residual[raw_residual <= tau] - raw_residual[raw_residual <= tau]))
 40    max_outside = np.max(residual)
 41    cos = np.sum((h-c)*(g-c), axis=1) / (residual*raw_residual + 1e-30)
 42    return {"inside_max_error": float(inside_err), "max_output_residual": float(max_outside),
 43            "tau": tau, "minimum_direction_cosine": float(np.min(cos))}
 44
 45def draw_gradients(x, batch, d, rng, regime):
 46    if regime == "gaussian":
 47        noise = rng.normal(0, 0.35, size=(batch, d))
 48        # Same occasional contamination rate in both controls, but modest magnitude.
 49        p, mult = 0.002, 8.
 50    else:
 51        df = 1.5
 52        noise = rng.standard_t(df, size=(batch, d)) * 0.20
 53        p, mult = 0.012, 35.
 54    mask = rng.random(batch) < p
 55    if np.any(mask):
 56        noise[mask] += rng.normal(size=(mask.sum(), d)) * mult
 57    return x[None, :] + noise
 58
 59def run_one(seed, method, regime, steps=300):
 60    rng = np.random.default_rng(seed)
 61    d, batch, eta, tau = 20, 64, 0.075, 2.0
 62    x = np.ones(d) * 5.0
 63    losses = []; update_norms = []; divergent = False
 64    for t in range(steps):
 65        gs = draw_gradients(x, batch, d, rng, regime)
 66        if method == "sgd":
 67            h = gs.mean(axis=0)
 68        elif method == "global_clip":
 69            h0 = gs.mean(axis=0)
 70            n = np.linalg.norm(h0)
 71            h = h0 * min(1., tau/(n+1e-12))
 72        elif method == "centered":
 73            c = coordinate_median(gs)
 74            h = clip_centered(gs, c, tau).mean(axis=0)
 75        else:
 76            raise ValueError(method)
 77        x -= eta*h
 78        loss = 0.5*float(np.dot(x,x))
 79        losses.append(loss); update_norms.append(float(eta*np.linalg.norm(h)))
 80        if not np.isfinite(loss) or loss > 1e8:
 81            divergent = True; break
 82    # Catastrophic steps are unusually large parameter changes, measured uniformly.
 83    return {"final_loss": losses[-1], "best_loss": min(losses),
 84            "median_update": float(np.median(update_norms)),
 85            "p99_update": float(np.quantile(update_norms, .99)),
 86            "diverged": divergent, "losses": losses}
 87
 88def aggregate():
 89    out = {"operator_check": operator_check(), "bias_scaling": grad_bias_scaling(), "runs": {}}
 90    for regime in ["heavy_tail", "gaussian"]:
 91        for method in ["sgd", "global_clip", "centered"]:
 92            rs = [run_one(s, method, regime) for s in range(12)]
 93            out["runs"][regime + "/" + method] = {
 94                "final_loss_mean": float(np.mean([r["final_loss"] for r in rs])),
 95                "final_loss_std": float(np.std([r["final_loss"] for r in rs])),
 96                "p99_update_mean": float(np.mean([r["p99_update"] for r in rs])),
 97                "divergences": int(sum(r["diverged"] for r in rs)),
 98                "final_losses": [r["final_loss"] for r in rs]}
 99    return out
100
101if __name__ == "__main__":
102    result = aggregate()
103    Path("results.json").write_text(json.dumps(result, indent=2))
104    print(json.dumps({"operator_check": result["operator_check"],
105                      "bias_scaling": result["bias_scaling"], "runs": result["runs"]}, indent=2))