Rank-One Proximal Quasi-Newton Optimizer / run_experiment.py
Mechanism failed
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()