Condition-number-aware restarted PAGE / page_experiment.py
Mechanism failed
1import json, math
2from dataclasses import dataclass
3import numpy as np
4
5SEED = 7
6
7@dataclass
8class QuadProblem:
9 H: np.ndarray # [n,d,d]
10 c: np.ndarray # [n,d]
11 xstar: np.ndarray
12 fstar: float
13
14 @property
15 def n(self): return self.H.shape[0]
16 def component_grads(self, x):
17 return np.einsum('nij,j->ni', self.H, x) + self.c
18 def grad(self, x): return self.component_grads(x).mean(0)
19 def loss(self, x):
20 # f_i(x)=.5*x'H_i x+c_i'x; constants chosen so F(xstar)=0
21 d=x-self.xstar
22 Hm=self.H.mean(0)
23 return .5*d@Hm@d
24
25
26def make_problem(kappa, n=64, d=2, seed=0):
27 rng=np.random.default_rng(seed)
28 # Common eigenvectors, mean Hessian has condition number kappa.
29 lo, hi = 1.0, float(kappa)
30 vals=np.array([lo, hi]) if d==2 else np.linspace(lo,hi,d)
31 H=np.tile(np.diag(vals), (n,1,1)).copy()
32 # Component-wise curvature and linear noise: mean linear term is zero,
33 # while individual gradients differ, making PAGE variance visible.
34 H += rng.normal(0, .04, size=H.shape)
35 H=(H+H.transpose(0,2,1))/2
36 # Keep positive definite.
37 for i in range(n): H[i] += max(0, .2-np.linalg.eigvalsh(H[i]).min())*np.eye(d)
38 xstar=np.array([1.0,-1.0])[:d]
39 noise=rng.normal(0, .35, size=(n,d)); noise-=noise.mean(0)
40 c=-np.einsum('nij,j->ni',H,xstar)+noise
41 return QuadProblem(H,c,xstar,0.)
42
43
44def constants(prob):
45 Hm=prob.H.mean(0)
46 mu=np.linalg.eigvalsh(Hm).min()
47 Lms=np.sqrt(np.mean(np.sum(prob.H**2,axis=(1,2))))
48 # For this quadratic, mean Hessian smoothness is also the exact PL curvature.
49 return mu, Lms, Lms/mu
50
51
52def run_page(prob, phase, steps=1200, batch=1, eta_scale=.25, seed=0):
53 rng=np.random.default_rng(seed); n=prob.n
54 # Conservative step based on mean smoothness; fixed across policies.
55 L=np.linalg.eigvalsh(prob.H.mean(0)).max()
56 eta=eta_scale/L
57 x=np.array([4.,-4.]); xp=x.copy(); v=None
58 calls=0; rows=[]; refreshes=0
59 for t in range(steps):
60 if t % phase == 0 or v is None:
61 v=prob.grad(x); calls += n; refreshes += 1
62 xp=x.copy()
63 else:
64 ids=rng.integers(0,n,size=batch)
65 gx=prob.component_grads(x)[ids]
66 gp=prob.component_grads(xp)[ids]
67 v=v+(gx-gp).mean(0); calls += batch
68 xp=x.copy()
69 x=x-eta*v
70 if t in (0,1,2,4,8,16,32,64,128,256,512,799,1199):
71 rows.append((calls,prob.loss(x),refreshes))
72 return rows
73
74
75def interpolate(rows, target):
76 for calls, loss, _ in rows:
77 if loss <= target: return calls
78 return math.inf
79
80
81def variance_check(prob, seed=3):
82 rng=np.random.default_rng(seed); x=np.array([2.,-2.]); xp=x+np.array([.05,-.03])
83 true=prob.grad(x); vals=[]
84 gs=prob.component_grads(x); gps=prob.component_grads(xp)
85 for _ in range(4000):
86 i=rng.integers(prob.n)
87 vals.append(true + (gs[i]-gps[i]) - (prob.grad(x)-prob.grad(xp)))
88 return float(np.mean(np.sum(np.asarray(vals)**2,axis=1)))
89
90
91def main():
92 # Core math sanity: PL ratio on the mean quadratic equals its smallest eigenvalue;
93 # PAGE difference estimator is unbiased conditional on the previous reference.
94 report={'math_check':{},'experiments':{}}
95 for k in (2, 16):
96 p=make_problem(k,seed=100+k)
97 mu,lms,kap=constants(p)
98 Hm=p.H.mean(0); eig=np.linalg.eigvalsh(Hm)
99 x=np.array([1.7,-2.2]); gap=p.loss(x); ratio=np.linalg.norm(p.grad(x))**2/(2*gap)
100 # fixed high-condition schedule and regime-aware short schedule
101 fixed=run_page(p, phase=32, seed=19+k)
102 adaptive_phase=4 if kap <= math.sqrt(p.n) else min(32,max(4,round(kap*math.sqrt(p.n))))
103 adaptive=run_page(p, phase=adaptive_phase, seed=19+k)
104 target=1e-4
105 report['experiments'][str(k)]={
106 'mu':mu,'L_ms':lms,'kappa':kap,'sqrt_n':math.sqrt(p.n),
107 'fixed_phase':32,'adaptive_phase':adaptive_phase,
108 'fixed_calls_to_1e-4':interpolate(fixed,target),
109 'adaptive_calls_to_1e-4':interpolate(adaptive,target),
110 'fixed_trace':fixed,'adaptive_trace':adaptive,
111 'difference_estimator_second_moment':variance_check(p)
112 }
113 report['math_check'][str(k)]={
114 'eigenvalues_mean_H':eig.tolist(),
115 'PL_ratio_at_test_x':ratio,
116 'PL_ratio_expected_range':[float(eig.min()),float(eig.max())],
117 'ratio_within_eigenvalue_range':bool(eig.min()-1e-8 <= ratio <= eig.max()+1e-8)
118 }
119 with open('results.json','w') as f: json.dump(report,f,indent=2)
120 print(json.dumps(report,indent=2))
121
122if __name__=='__main__': main()