Rank-One Proximal Quasi-Newton Optimizer / run_experiment.py

Mechanism failed

Raw ⬇ ZIP
  1import json
  2import time
  3import numpy as np
  4from scipy.optimize import minimize
  5
  6
  7def group_prox(w, d, lam, groups):
  8    """Exact prox of lam*sum_g ||x_g|| under diagonal metric.
  9
 10    For a nonzero group x_i = w_i/(1+c/d_i), c=lam/||x_g||.
 11    """
 12    x = w.copy()
 13    for sl in groups:
 14        wg, dg = w[sl], d[sl]
 15        if np.linalg.norm(dg * wg) <= lam:
 16            x[sl] = 0.0
 17            continue
 18        def phi(c):
 19            xx = wg / (1.0 + c / dg)
 20            return c * np.linalg.norm(xx) - lam
 21        lo, hi = 0.0, max(lam / max(np.linalg.norm(wg), 1e-30), 1e-12)
 22        while phi(hi) < 0.0:
 23            hi *= 2.0
 24        for _ in range(24):
 25            mid = 0.5 * (lo + hi)
 26            if phi(mid) < 0.0:
 27                lo = mid
 28            else:
 29                hi = mid
 30        c = 0.5 * (lo + hi)
 31        x[sl] = wg / (1.0 + c / dg)
 32    return x
 33
 34
 35def objective(theta, X, y, lam, groups):
 36    r = X @ theta - y
 37    return 0.5 * np.mean(r * r) + lam * sum(np.linalg.norm(theta[s]) for s in groups)
 38
 39
 40def grad(theta, X, y):
 41    return X.T @ (X @ theta - y) / X.shape[0]
 42
 43
 44def rank_one_step(theta, g, d, a, rho, lam, groups, root_steps=12):
 45    dinv = 1.0 / d
 46    v = dinv * a
 47    den = 1.0 + rho * np.dot(a, v)
 48    Minvg = dinv * g - (rho * v * np.dot(a, dinv * g)) / den
 49    z = theta - Minvg
 50
 51    def xr(q):
 52        x = group_prox(z - v * q, d, lam, groups)
 53        return x, q - rho * np.dot(a, x - z)
 54
 55    q0 = 0.0
 56    _, r0 = xr(q0)
 57    if abs(r0) < 1e-14:
 58        return _, q0, r0, 1
 59    width = max(1.0, abs(r0))
 60    lo, hi = q0 - width, q0 + width
 61    _, rlo = xr(lo)
 62    _, rhi = xr(hi)
 63    for _ in range(60):
 64        if rlo <= 0.0 <= rhi:
 65            break
 66        width *= 2.0
 67        lo, hi = q0 - width, q0 + width
 68        _, rlo = xr(lo)
 69        _, rhi = xr(hi)
 70    else:
 71        raise RuntimeError("failed to bracket monotone residual")
 72    for _ in range(root_steps):
 73        mid = 0.5 * (lo + hi)
 74        _, rm = xr(mid)
 75        if rm <= 0.0:
 76            lo = mid
 77        else:
 78            hi = mid
 79    q = 0.5 * (lo + hi)
 80    x, r = xr(q)
 81    return x, q, r, root_steps
 82
 83
 84def direct_metric_solution(z, d, a, rho, lam, groups):
 85    M = np.diag(d) + rho * np.outer(a, a)
 86    fun = lambda x: lam * sum(np.linalg.norm(x[s]) for s in groups) + 0.5 * (x-z) @ M @ (x-z)
 87    # Smooth approximation is not used: SLSQP handles the small nonsmooth check.
 88    res = minimize(fun, z.copy(), method="SLSQP", options={"ftol": 1e-12, "maxiter": 1000})
 89    return res.x, res.fun, res.success
 90
 91
 92def math_check(rng):
 93    n, bs = 24, 3
 94    groups = [slice(i, i + bs) for i in range(0, n, bs)]
 95    d = rng.uniform(0.5, 2.0, n)
 96    a = rng.normal(size=n)
 97    rho = 1.7
 98    z = rng.normal(size=n)
 99    lam = 0.08
100    # Solve the stated residual equation and compare to direct metric minimization.
101    def solve_from_z():
102        v = a / d
103        def xr(q):
104            x = group_prox(z - v*q, d, lam, groups)
105            return x, q-rho*np.dot(a, x-z)
106        lo, hi = -1.0, 1.0
107        while xr(lo)[1] > 0: lo *= 2
108        while xr(hi)[1] < 0: hi *= 2
109        for _ in range(60):
110            m = (lo+hi)/2
111            if xr(m)[1] <= 0: lo = m
112            else: hi = m
113        q=(lo+hi)/2
114        return xr(q)[0], q, xr(q)[1]
115    x, q, residual = solve_from_z()
116    xd, _, ok = direct_metric_solution(z, d, a, rho, lam, groups)
117    M = np.diag(d) + rho*np.outer(a,a)
118    Minv_sm = np.diag(1/d) - rho*np.outer(a/d, a/d)/(1+rho*np.dot(a,a/d))
119    sm_err = np.max(np.abs(Minv_sm @ M - np.eye(n)))
120    return {"root_residual": float(abs(residual)), "direct_solution_max_error": float(np.max(abs(x-xd))), "direct_solver_success": bool(ok), "sherman_morrison_identity_error": float(sm_err), "q": float(q)}
121
122
123def run_optimizer(X, y, groups, lam, rank_one, seed=0, steps=70):
124    rng = np.random.default_rng(seed)
125    p = X.shape[1]
126    theta = np.zeros(p)
127    d = np.ones(p) * 0.25
128    prev_theta = theta.copy()
129    prev_g = grad(theta, X, y)
130    residuals, vals = [], []
131    t0 = time.perf_counter()
132    for it in range(steps):
133        g = grad(theta, X, y)
134        # Adam/RMS-style diagonal curvature estimate, clipped for stability.
135        d = np.clip(0.92*d + 0.08*(g*g + 1e-3), 0.03, 3.0)
136        if rank_one and it >= 1:
137            s = theta - prev_theta
138            yy = g - prev_g
139            yn = np.linalg.norm(yy)
140            a = yy / (yn + 1e-12)
141            rho = max(0.0, float(np.dot(s, yy)/(np.dot(s,s)+1e-12) - np.mean(d)))
142            rho = min(rho, 2.0)
143            theta_new, q, rr, _ = rank_one_step(theta, g, d, a, rho, lam, groups)
144            residuals.append(abs(rr))
145        else:
146            # Matching diagonal proximal-gradient baseline.
147            z = theta - g/d
148            theta_new = group_prox(z, d, lam, groups)
149            residuals.append(0.0)
150        prev_theta, prev_g, theta = theta, g, theta_new
151        vals.append(objective(theta, X, y, lam, groups))
152    elapsed = time.perf_counter()-t0
153    active = sum(np.linalg.norm(theta[s]) > 1e-7 for s in groups)
154    return {"final_objective": float(vals[-1]), "objective_at_30": float(vals[min(29, len(vals)-1)]), "objective_at_90": float(vals[min(89, len(vals)-1)]), "seconds": elapsed, "active_groups": int(active), "mean_abs_residual": float(np.mean(residuals)), "max_abs_residual": float(np.max(residuals))}
155
156
157def main():
158    rng = np.random.default_rng(123)
159    n, p, bs = 300, 40, 4
160    X = rng.normal(size=(n,p)); X /= np.sqrt(np.mean(X*X, axis=0, keepdims=True))
161    groups = [slice(i, i+bs) for i in range(0,p,bs)]
162    true = np.zeros(p)
163    chosen = [1, 3, 6, 8]
164    for k in chosen: true[groups[k]] = rng.normal(size=bs)
165    y = X @ true + 0.08*rng.normal(size=n)
166    lam = 0.025
167    check = math_check(np.random.default_rng(9))
168    baseline = run_optimizer(X,y,groups,lam,False,seed=4)
169    idea = run_optimizer(X,y,groups,lam,True,seed=4)
170    result = {"math_check":check, "baseline":baseline, "idea":idea, "settings":{"n":n,"p":p,"steps":180,"lambda":lam,"seed":123}}
171    with open("results.json","w") as f: json.dump(result,f,indent=2)
172    print(json.dumps(result, indent=2))
173
174if __name__ == "__main__": main()