Kurtosis-calibrated gradient clipping / experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  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()