Adaptive Proximal Quasi-Newton Training / experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4
  5SEED = 2394
  6np.random.seed(SEED); random.seed(SEED)
  7
  8
  9def prox_l1(z, eta, lam, h):
 10    # Diagonal-metric prox for R(w)=lambda*||w||_1.
 11    return np.sign(z) * np.maximum(np.abs(z) - eta * lam * h, 0.0)
 12
 13
 14def adaptive_quadratic(A, h, w0, eta0=0.05, gamma_up=1.5, gamma_down=.5,
 15                       q=3, lam=0.0, steps=120):
 16    w = w0.copy(); eta = eta0; success = 0; accepts = 0; rejects = 0
 17    eta_hist=[]; obj_hist=[]; resid_hist=[]
 18    def f(x): return .5 * x @ A @ x
 19    for _ in range(steps):
 20        g = A @ w
 21        trial = prox_l1(w - eta * h * g, eta, lam, h)
 22        old, new = f(w), f(trial); s = trial-w
 23        # Objective Armijo safeguard. For R=0 this is the stated smooth merit.
 24        rhs = old + 1e-4 * (g @ s)
 25        if np.isfinite(new) and new <= rhs:
 26            w = trial; accepts += 1; success += 1
 27            if success >= q:
 28                eta *= gamma_up; success = 0
 29        else:
 30            rejects += 1; success = 0; eta *= gamma_down
 31        eta_hist.append(eta); obj_hist.append(f(w)); resid_hist.append(np.linalg.norm(A@w))
 32    return dict(w=w, eta=np.array(eta_hist), obj=np.array(obj_hist), residual=np.array(resid_hist), accepts=accepts, rejects=rejects)
 33
 34
 35def fixed_quadratic(A, h, w0, eta, steps=100):
 36    w=w0.copy(); vals=[]
 37    for _ in range(steps):
 38        vals.append(.5*w@A@w); w=w-eta*h*(A@w)
 39    vals.append(.5*w@A@w)
 40    return np.array(vals)
 41
 42
 43def run_stability_sweeps():
 44    # Diagonal H and A make the exact transformed eigenvalues transparent.
 45    dim=8; h=np.array([.5, .8, 1.0, 1.2, 1.5, .7, 1.1, .9])
 46    lambdas=np.array([1.,2.,4.,7.,10.,14.,18.,25.])
 47    A=np.diag(lambdas); w0=np.ones(dim)
 48    L=float(np.max(h*lambdas)); eta_c=2/L
 49    # Prediction 1: a fixed linear iteration is stable iff eta*L < 2.
 50    etas=np.linspace(.1*eta_c, 2.2*eta_c, 22)
 51    stable=[]
 52    for eta in etas:
 53        vals=fixed_quadratic(A,h,w0,eta,steps=80)
 54        stable.append(bool(np.all(np.isfinite(vals)) and vals[-1] < vals[0] and np.max(vals)<1e12))
 55    boundary=next((float(etas[i]) for i in range(len(etas)) if not stable[i]), float('nan'))
 56    # Refine empirical boundary by binary search on final norm.
 57    lo,hi=0.,3*eta_c
 58    for _ in range(45):
 59        mid=(lo+hi)/2
 60        vals=fixed_quadratic(A,h,w0,mid,steps=100)
 61        ok=np.all(np.isfinite(vals)) and vals[-1] < vals[0] and np.max(vals)<1e10
 62        if ok: lo=mid
 63        else: hi=mid
 64    empirical=float(lo)
 65    # Prediction 2: asymptotic contraction factor is max_i |1-eta*h_i*lambda_i|.
 66    eta_test=.55*eta_c
 67    pred_factor=float(np.max(np.abs(1-eta_test*h*lambdas)))
 68    vals=fixed_quadratic(A,h,w0,eta_test,steps=40)
 69    # Objective ratio tends to squared factor; estimate from late ratios.
 70    observed=float(np.median(np.sqrt(np.maximum(vals[-10:-1],1e-300)/np.maximum(vals[-11:-2],1e-300))))
 71    # Prediction 3: adaptive safeguard settles below the boundary and rejects unsafe trials.
 72    ad=adaptive_quadratic(A,h,w0,eta0=eta_c*.18,steps=180)
 73    late_eta=float(np.median(ad['eta'][-30:]))
 74    late_max=float(np.max(ad['eta'][-30:]))
 75    return {
 76      'L_exact':L, 'eta_critical_predicted':eta_c,
 77      'boundary_coarse_first_unstable':boundary, 'boundary_binary_observed':empirical,
 78      'boundary_relative_error':abs(empirical-eta_c)/eta_c,
 79      'contraction_eta':eta_test, 'contraction_factor_predicted':pred_factor,
 80      'contraction_factor_observed':observed, 'contraction_abs_error':abs(observed-pred_factor),
 81      'adaptive_late_median_eta':late_eta, 'adaptive_late_max_eta':late_max,
 82      'adaptive_eta_ratio_to_boundary':late_eta/eta_c,
 83      'adaptive_accepts':ad['accepts'], 'adaptive_rejects':ad['rejects'],
 84      'adaptive_final_objective':float(ad['obj'][-1]),
 85      'stable_sweep': [{'eta':float(e),'etaL':float(e*L),'stable':s} for e,s in zip(etas,stable)]
 86    }
 87
 88
 89def sparse_regression():
 90    rng=np.random.default_rng(SEED+7); n,d=256,24
 91    X=rng.normal(size=(n,d)); true=np.zeros(d); true[[1,5,9,15,20]]=rng.normal(size=5)
 92    y=X@true+.08*rng.normal(size=n); A=X.T@X/n; b=X.T@y/n; w0=np.zeros(d)
 93    lam=.025; L=np.linalg.eigvalsh(A).max(); h=np.ones(d); eta=1/L
 94    def obj(w): return .5*np.mean((X@w-y)**2)+lam*np.abs(w).sum()
 95    def run_sgd():
 96        w=w0.copy(); hist=[]
 97        for _ in range(250):
 98            ix=rng.choice(n,64,replace=False); g=X[ix].T@(X[ix]@w-y[ix])/64
 99            w-=.35*g; hist.append(obj(w))
100        return w,hist
101    def run_prox():
102        w=w0.copy(); hist=[]
103        for _ in range(250):
104            w=prox_l1(w-eta*(A@w-b),eta,lam,h); hist.append(obj(w))
105        return w,hist
106    def run_ad():
107        # Use deterministic full-batch gradients and the same safeguarded update.
108        w=w0.copy(); H=np.ones(d); et=eta*.5; succ=0; hist=[]; acc=rej=0
109        for _ in range(250):
110            g=A@w-b; trial=prox_l1(w-et*H*g,et,lam,H)
111            old=obj(w); new=obj(trial); s=trial-w
112            if np.isfinite(new) and new <= old+1e-4*(g@s):
113                oldw=w; oldg=g; w=trial; acc+=1; succ+=1
114                # Diagonal inverse-BFGS-like secant update, clipped for robustness.
115                yk=(A@w-b)-oldg; sk=w-oldw; den=yk*sk
116                mask=den>1e-10
117                H[mask]=np.clip(sk[mask]/den[mask],1e-4,100.)
118                if succ>=3: et*=1.5; succ=0
119            else: rej+=1; succ=0; et*=.5
120            hist.append(obj(w))
121        return w,hist,acc,rej
122    ws,hs=run_sgd(); wp,hp=run_prox(); wa,ha,ac,re=run_ad()
123    return {'sgd_final':float(hs[-1]),'prox_final':float(hp[-1]),'adaptive_final':float(ha[-1]),
124            'sgd_nonzeros':int(np.sum(np.abs(ws)>1e-4)), 'prox_nonzeros':int(np.sum(np.abs(wp)>1e-4)),
125            'adaptive_nonzeros':int(np.sum(np.abs(wa)>1e-4)), 'adaptive_accepts':ac,'adaptive_rejects':re,
126            'iterations':250,'lambda':lam,'L':float(L)}
127
128
129def main():
130    out={'seed':SEED,'stability':run_stability_sweeps(),'sparse_regression':sparse_regression()}
131    Path('results.json').write_text(json.dumps(out,indent=2))
132    s=out['stability']; r=out['sparse_regression']
133    print(json.dumps({'stability':s,'sparse_regression':r},indent=2))
134
135if __name__=='__main__': main()