Condition-number-aware restarted PAGE / page_experiment.py

Mechanism failed

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