Centered Heavy-Tail Clipping Optimizer / experiment.py
Failed on benchmark
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))