Overshoot Budget Controller / overshoot_controller.py
Mechanism confirmed, baseline not beaten
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))