Kurtosis-calibrated gradient clipping / experiment.py
Beats tuned baseline
1import json, math, random
2import numpy as np
3
4
5def c_transition(k):
6 k = max(float(k), 1.0)
7 return math.sqrt(((k + 1.0) + math.sqrt(max(0.0, (k + 1.0)**2 - 4.0))) / 2.0)
8
9
10def calibrated_t(k, delta):
11 k = max(float(k), 1.0)
12 delta = float(delta)
13 t = math.sqrt(1.0 + math.sqrt(max(0.0, (k - 1.0) * (1.0 / delta - 1.0))))
14 return max(t, c_transition(k))
15
16
17def tail_formula(t, k):
18 return (k - 1.0) / ((t*t - 1.0)**2 + k - 1.0) if k > 1 else 0.0
19
20
21def math_check():
22 rows = []
23 max_inv_err = 0.0
24 min_tail_margin = float('inf')
25 for k in [1.01, 1.1, 2, 3, 10, 100]:
26 c = c_transition(k)
27 poly = c**4 - (k + 1)*c**2 + 1
28 for d in [1e-3, .01, .1, .5]:
29 raw = math.sqrt(1 + math.sqrt((k-1)*(1/d-1)))
30 t = calibrated_t(k, d)
31 # Inversion is exact before regime projection; the clipping projection is deliberate.
32 if raw >= c:
33 inv_err = abs(tail_formula(raw, k) - d)
34 max_inv_err = max(max_inv_err, inv_err)
35 min_tail_margin = min(min_tail_margin, t-c)
36 rows.append({'k': k, 'c': c, 'transition_polynomial': poly})
37 return {'max_inverse_error': max_inv_err, 'minimum_projected_tail_margin': min_tail_margin, 'rows': rows}
38
39
40def sample_kurtosis(x):
41 z = x - np.mean(x)
42 v = np.mean(z*z)
43 return float(np.mean(z**4) / max(v*v, 1e-18))
44
45
46def empirical_tail_check(seed=17, n=2000000):
47 rng = np.random.default_rng(seed)
48 # Standardized Student-t has a known kurtosis 3 + 6/(nu-4), nu>4.
49 out = []
50 for nu in [5, 8, 20]:
51 x = rng.standard_t(nu, n)
52 x = (x - x.mean()) / x.std()
53 k_bound = 3 + 6/(nu-4)
54 for d in [.01, .05]:
55 t = calibrated_t(k_bound*1.05, d) # modest safety factor
56 observed = float(np.mean(x >= t))
57 out.append({'distribution': 'student_t', 'nu': nu, 'kurtosis_bound': k_bound*1.05,
58 'delta': d, 'threshold': t, 'observed_tail': observed,
59 'bound_formula_at_threshold': tail_formula(t, k_bound*1.05)})
60 return out
61
62
63def run_optimizer(seed, mode, steps=300, batch=64, delta=.02):
64 rng = np.random.default_rng(seed)
65 x = 8.0
66 beta = .95
67 mu_ema, v_ema, q_ema = 0., 1., 3.
68 losses, spikes, tails = [], 0, []
69 lr = .08
70 for step in range(steps):
71 # Rare, centered, heavy-tailed gradient contamination.
72 noise = rng.normal(0, .35, batch)
73 rare = rng.random(batch) < .025
74 noise[rare] += rng.choice([-1, 1], rare.sum()) * 12.0
75 g = x + noise
76 if mode == 'calibrated':
77 m = float(g.mean())
78 z = g - m
79 vv = float(np.mean(z*z))
80 qq = float(np.mean(z**4))
81 mu_ema = beta*mu_ema + (1-beta)*m
82 v_ema = beta*v_ema + (1-beta)*vv
83 q_ema = beta*q_ema + (1-beta)*qq
84 k = min(1000., max(1., 1.5*q_ema/(v_ema*v_ema + 1e-12)))
85 t = calibrated_t(k, delta)
86 tau = t*math.sqrt(v_ema + 1e-12)
87 clipped = m + np.clip(g-m, -tau, tau)
88 used = float(clipped.mean())
89 tails.append(float(np.mean(np.abs(g-m) > tau)))
90 spikes += int(np.any(np.abs(g-m) > tau))
91 elif mode == 'fixed':
92 tau = 2.0
93 used = float(np.clip(g, -tau, tau).mean())
94 tails.append(float(np.mean(np.abs(g) > tau)))
95 spikes += int(np.any(np.abs(g) > tau))
96 else:
97 used = float(g.mean())
98 spikes += int(np.any(np.abs(g) > 8.0))
99 x -= lr * used
100 losses.append(.5*x*x)
101 return {'final_loss': losses[-1], 'mean_last50_loss': float(np.mean(losses[-50:])),
102 'max_loss': max(losses), 'spike_steps': spikes,
103 'mean_tail_above_threshold': float(np.mean(tails)) if tails else None,
104 'final_abs_x': abs(x)}
105
106
107def main():
108 results = {'math_check': math_check(), 'empirical_tail_check': empirical_tail_check()}
109 allruns = {}
110 for mode in ['none', 'fixed', 'calibrated']:
111 vals = [run_optimizer(s, mode) for s in range(10, 20)]
112 allruns[mode] = {k: (None if vals[0][k] is None else float(np.mean([v[k] for v in vals]))) for k in vals[0]}
113 allruns[mode]['replicate_std_final_loss'] = float(np.std([v['final_loss'] for v in vals]))
114 results['optimizer'] = allruns
115 with open('results.json', 'w') as f:
116 json.dump(results, f, indent=2)
117 print(json.dumps(results, indent=2))
118
119if __name__ == '__main__':
120 main()