Overshoot Budget Controller / overshoot_controller.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import math, json
 2import numpy as np
 3
 4EPS = 1e-12
 5
 6def huber(x, eps, delta):
 7    if x >= delta: return delta*abs(x) - delta*delta/2
 8    if x <= -eps: return eps*abs(x) - eps*eps/2
 9    return .5*x*x
10
11def huber_grad(x, eps, delta):
12    if x >= delta: return delta
13    if x <= -eps: return -eps
14    return x
15
16def trajectory(etas, m, criterion):
17    """The paper's asymmetric-Huber construction, with 1-indexed m."""
18    before = float(sum(etas[:m-1])); after = float(sum(etas[m:]))
19    delta = 1.0/(1.0 + before)
20    overshoot = (float(etas[m-1])-1.0)*delta
21    eps = overshoot/(1.0 + (2.0*after if criterion == 'R' else after))
22    x = 1.0; grads=[]
23    for eta in etas:
24        g = huber_grad(x, eps, delta)
25        grads.append(g); x -= float(eta)*g
26    initial = huber(1.0, eps, delta)
27    R = huber(x, eps, delta) / .5
28    G = (.5 * huber_grad(x, eps, delta)**2) / initial
29    rhsR = (float(etas[m-1])-1)**2 / ((1+before)**2 * (1+2*after))
30    rhsG = (float(etas[m-1])-1)**2 / ((1+2*before) * (1+after)**2)
31    return R, G, rhsR, rhsG
32
33def verify_math(seed=7, trials=2000):
34    rng=np.random.default_rng(seed); worstR=1e9; worstG=1e9; failures=0
35    for _ in range(trials):
36        n=int(rng.integers(2,12)); etas=rng.uniform(.03,2.5,n)
37        m=int(rng.integers(1,n+1))
38        if etas[m-1] <= 1: etas[m-1]=1.01+rng.random()*1.5
39        RR,_,rR,_=trajectory(etas,m,'R'); _,GG,_,rG=trajectory(etas,m,'G')
40        worstR=min(worstR,RR/rR); worstG=min(worstG,GG/rG)
41        failures += int(RR+1e-9 < rR or GG+1e-9 < rG)
42    a=np.array([.15,.25,1.8,.4,.1]); b=np.array([1.8,.15,.25,.4,.1])
43    return {'trials':int(trials),'failures':int(failures),
44      'min_R_ratio':float(worstR),'min_G_ratio':float(worstG),
45      'ordering_R_rhs':[float(trajectory(a,3,'R')[2]),float(trajectory(b,1,'R')[2])],
46      'ordering_G_rhs':[float(trajectory(a,3,'G')[3]),float(trajectory(b,1,'G')[3])]}
47
48def run_controller(proposals, rho=.02, beta=.9, mode='controller', grad_clip=1.0):
49    x=1.0; S=0.; prevq=None; v=1.0; records=[]
50    for proposal in proposals:
51        raw_g=x; q=raw_g*raw_g
52        ratio=1.0 if prevq is None else (q+EPS)/(prevq+EPS)
53        v=beta*v+(1-beta)*ratio
54        cap=1+math.sqrt(1+2*S)*math.sqrt(rho*v+EPS)
55        eta=min(float(proposal),cap) if mode=='controller' else float(proposal)
56        g=max(-grad_clip,min(grad_clip,raw_g)) if mode=='grad_clip' else raw_g
57        x=x-eta*g; S+=eta; prevq=q
58        records.append((x,eta,cap,ratio,g))
59    return x, records
60
61def experiment():
62    # Late, isolated proposals are deliberately beyond the unit-quadratic stability limit.
63    n=64; proposals=np.full(n,.2); proposals[[7,23,39,55]]=20.0
64    out={}
65    for mode in ['sgd','grad_clip','controller']:
66        x, rec=run_controller(proposals,mode=mode)
67        losses=[.5*r[0]*r[0] for r in rec]
68        out[mode]={'final_loss':float(losses[-1]) if np.isfinite(losses[-1]) else 'nonfinite',
69          'max_abs_x':float(max(abs(r[0]) for r in rec)) if all(np.isfinite(r[0]) for r in rec) else 'nonfinite',
70          'max_loss':float(max(losses)) if all(np.isfinite(z) for z in losses) else 'nonfinite',
71          'clipped_steps':int(sum(r[1]<p-1e-10 for r,p in zip(rec,proposals))),
72          'spike_etas':[float(rec[i][1]) for i in [7,23,39,55]],
73          'loss_after_spikes':[float(losses[i+1]) for i in [7,23,39,55]]}
74    _,r=run_controller(np.where(np.arange(n)==55,20.0,.2),mode='controller')
75    out['single_late_spike']={'proposal':20.0,'accepted':float(r[55][1]),'cap':float(r[55][2])}
76    return out
77
78if __name__=='__main__':
79    print(json.dumps({'math_check':verify_math(),'experiment':experiment()},indent=2))