import json, math from dataclasses import dataclass import numpy as np SEED = 7 @dataclass class QuadProblem: H: np.ndarray # [n,d,d] c: np.ndarray # [n,d] xstar: np.ndarray fstar: float @property def n(self): return self.H.shape[0] def component_grads(self, x): return np.einsum('nij,j->ni', self.H, x) + self.c def grad(self, x): return self.component_grads(x).mean(0) def loss(self, x): # f_i(x)=.5*x'H_i x+c_i'x; constants chosen so F(xstar)=0 d=x-self.xstar Hm=self.H.mean(0) return .5*d@Hm@d def make_problem(kappa, n=64, d=2, seed=0): rng=np.random.default_rng(seed) # Common eigenvectors, mean Hessian has condition number kappa. lo, hi = 1.0, float(kappa) vals=np.array([lo, hi]) if d==2 else np.linspace(lo,hi,d) H=np.tile(np.diag(vals), (n,1,1)).copy() # Component-wise curvature and linear noise: mean linear term is zero, # while individual gradients differ, making PAGE variance visible. H += rng.normal(0, .04, size=H.shape) H=(H+H.transpose(0,2,1))/2 # Keep positive definite. for i in range(n): H[i] += max(0, .2-np.linalg.eigvalsh(H[i]).min())*np.eye(d) xstar=np.array([1.0,-1.0])[:d] noise=rng.normal(0, .35, size=(n,d)); noise-=noise.mean(0) c=-np.einsum('nij,j->ni',H,xstar)+noise return QuadProblem(H,c,xstar,0.) def constants(prob): Hm=prob.H.mean(0) mu=np.linalg.eigvalsh(Hm).min() Lms=np.sqrt(np.mean(np.sum(prob.H**2,axis=(1,2)))) # For this quadratic, mean Hessian smoothness is also the exact PL curvature. return mu, Lms, Lms/mu def run_page(prob, phase, steps=1200, batch=1, eta_scale=.25, seed=0): rng=np.random.default_rng(seed); n=prob.n # Conservative step based on mean smoothness; fixed across policies. L=np.linalg.eigvalsh(prob.H.mean(0)).max() eta=eta_scale/L x=np.array([4.,-4.]); xp=x.copy(); v=None calls=0; rows=[]; refreshes=0 for t in range(steps): if t % phase == 0 or v is None: v=prob.grad(x); calls += n; refreshes += 1 xp=x.copy() else: ids=rng.integers(0,n,size=batch) gx=prob.component_grads(x)[ids] gp=prob.component_grads(xp)[ids] v=v+(gx-gp).mean(0); calls += batch xp=x.copy() x=x-eta*v if t in (0,1,2,4,8,16,32,64,128,256,512,799,1199): rows.append((calls,prob.loss(x),refreshes)) return rows def interpolate(rows, target): for calls, loss, _ in rows: if loss <= target: return calls return math.inf def variance_check(prob, seed=3): rng=np.random.default_rng(seed); x=np.array([2.,-2.]); xp=x+np.array([.05,-.03]) true=prob.grad(x); vals=[] gs=prob.component_grads(x); gps=prob.component_grads(xp) for _ in range(4000): i=rng.integers(prob.n) vals.append(true + (gs[i]-gps[i]) - (prob.grad(x)-prob.grad(xp))) return float(np.mean(np.sum(np.asarray(vals)**2,axis=1))) def main(): # Core math sanity: PL ratio on the mean quadratic equals its smallest eigenvalue; # PAGE difference estimator is unbiased conditional on the previous reference. report={'math_check':{},'experiments':{}} for k in (2, 16): p=make_problem(k,seed=100+k) mu,lms,kap=constants(p) Hm=p.H.mean(0); eig=np.linalg.eigvalsh(Hm) x=np.array([1.7,-2.2]); gap=p.loss(x); ratio=np.linalg.norm(p.grad(x))**2/(2*gap) # fixed high-condition schedule and regime-aware short schedule fixed=run_page(p, phase=32, seed=19+k) adaptive_phase=4 if kap <= math.sqrt(p.n) else min(32,max(4,round(kap*math.sqrt(p.n)))) adaptive=run_page(p, phase=adaptive_phase, seed=19+k) target=1e-4 report['experiments'][str(k)]={ 'mu':mu,'L_ms':lms,'kappa':kap,'sqrt_n':math.sqrt(p.n), 'fixed_phase':32,'adaptive_phase':adaptive_phase, 'fixed_calls_to_1e-4':interpolate(fixed,target), 'adaptive_calls_to_1e-4':interpolate(adaptive,target), 'fixed_trace':fixed,'adaptive_trace':adaptive, 'difference_estimator_second_moment':variance_check(p) } report['math_check'][str(k)]={ 'eigenvalues_mean_H':eig.tolist(), 'PL_ratio_at_test_x':ratio, 'PL_ratio_expected_range':[float(eig.min()),float(eig.max())], 'ratio_within_eigenvalue_range':bool(eig.min()-1e-8 <= ratio <= eig.max()+1e-8) } with open('results.json','w') as f: json.dump(report,f,indent=2) print(json.dumps(report,indent=2)) if __name__=='__main__': main()