import json import time import numpy as np from scipy.optimize import minimize def group_prox(w, d, lam, groups): """Exact prox of lam*sum_g ||x_g|| under diagonal metric. For a nonzero group x_i = w_i/(1+c/d_i), c=lam/||x_g||. """ x = w.copy() for sl in groups: wg, dg = w[sl], d[sl] if np.linalg.norm(dg * wg) <= lam: x[sl] = 0.0 continue def phi(c): xx = wg / (1.0 + c / dg) return c * np.linalg.norm(xx) - lam lo, hi = 0.0, max(lam / max(np.linalg.norm(wg), 1e-30), 1e-12) while phi(hi) < 0.0: hi *= 2.0 for _ in range(24): mid = 0.5 * (lo + hi) if phi(mid) < 0.0: lo = mid else: hi = mid c = 0.5 * (lo + hi) x[sl] = wg / (1.0 + c / dg) return x def objective(theta, X, y, lam, groups): r = X @ theta - y return 0.5 * np.mean(r * r) + lam * sum(np.linalg.norm(theta[s]) for s in groups) def grad(theta, X, y): return X.T @ (X @ theta - y) / X.shape[0] def rank_one_step(theta, g, d, a, rho, lam, groups, root_steps=12): dinv = 1.0 / d v = dinv * a den = 1.0 + rho * np.dot(a, v) Minvg = dinv * g - (rho * v * np.dot(a, dinv * g)) / den z = theta - Minvg def xr(q): x = group_prox(z - v * q, d, lam, groups) return x, q - rho * np.dot(a, x - z) q0 = 0.0 _, r0 = xr(q0) if abs(r0) < 1e-14: return _, q0, r0, 1 width = max(1.0, abs(r0)) lo, hi = q0 - width, q0 + width _, rlo = xr(lo) _, rhi = xr(hi) for _ in range(60): if rlo <= 0.0 <= rhi: break width *= 2.0 lo, hi = q0 - width, q0 + width _, rlo = xr(lo) _, rhi = xr(hi) else: raise RuntimeError("failed to bracket monotone residual") for _ in range(root_steps): mid = 0.5 * (lo + hi) _, rm = xr(mid) if rm <= 0.0: lo = mid else: hi = mid q = 0.5 * (lo + hi) x, r = xr(q) return x, q, r, root_steps def direct_metric_solution(z, d, a, rho, lam, groups): M = np.diag(d) + rho * np.outer(a, a) fun = lambda x: lam * sum(np.linalg.norm(x[s]) for s in groups) + 0.5 * (x-z) @ M @ (x-z) # Smooth approximation is not used: SLSQP handles the small nonsmooth check. res = minimize(fun, z.copy(), method="SLSQP", options={"ftol": 1e-12, "maxiter": 1000}) return res.x, res.fun, res.success def math_check(rng): n, bs = 24, 3 groups = [slice(i, i + bs) for i in range(0, n, bs)] d = rng.uniform(0.5, 2.0, n) a = rng.normal(size=n) rho = 1.7 z = rng.normal(size=n) lam = 0.08 # Solve the stated residual equation and compare to direct metric minimization. def solve_from_z(): v = a / d def xr(q): x = group_prox(z - v*q, d, lam, groups) return x, q-rho*np.dot(a, x-z) lo, hi = -1.0, 1.0 while xr(lo)[1] > 0: lo *= 2 while xr(hi)[1] < 0: hi *= 2 for _ in range(60): m = (lo+hi)/2 if xr(m)[1] <= 0: lo = m else: hi = m q=(lo+hi)/2 return xr(q)[0], q, xr(q)[1] x, q, residual = solve_from_z() xd, _, ok = direct_metric_solution(z, d, a, rho, lam, groups) M = np.diag(d) + rho*np.outer(a,a) Minv_sm = np.diag(1/d) - rho*np.outer(a/d, a/d)/(1+rho*np.dot(a,a/d)) sm_err = np.max(np.abs(Minv_sm @ M - np.eye(n))) 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)} def run_optimizer(X, y, groups, lam, rank_one, seed=0, steps=70): rng = np.random.default_rng(seed) p = X.shape[1] theta = np.zeros(p) d = np.ones(p) * 0.25 prev_theta = theta.copy() prev_g = grad(theta, X, y) residuals, vals = [], [] t0 = time.perf_counter() for it in range(steps): g = grad(theta, X, y) # Adam/RMS-style diagonal curvature estimate, clipped for stability. d = np.clip(0.92*d + 0.08*(g*g + 1e-3), 0.03, 3.0) if rank_one and it >= 1: s = theta - prev_theta yy = g - prev_g yn = np.linalg.norm(yy) a = yy / (yn + 1e-12) rho = max(0.0, float(np.dot(s, yy)/(np.dot(s,s)+1e-12) - np.mean(d))) rho = min(rho, 2.0) theta_new, q, rr, _ = rank_one_step(theta, g, d, a, rho, lam, groups) residuals.append(abs(rr)) else: # Matching diagonal proximal-gradient baseline. z = theta - g/d theta_new = group_prox(z, d, lam, groups) residuals.append(0.0) prev_theta, prev_g, theta = theta, g, theta_new vals.append(objective(theta, X, y, lam, groups)) elapsed = time.perf_counter()-t0 active = sum(np.linalg.norm(theta[s]) > 1e-7 for s in groups) 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))} def main(): rng = np.random.default_rng(123) n, p, bs = 300, 40, 4 X = rng.normal(size=(n,p)); X /= np.sqrt(np.mean(X*X, axis=0, keepdims=True)) groups = [slice(i, i+bs) for i in range(0,p,bs)] true = np.zeros(p) chosen = [1, 3, 6, 8] for k in chosen: true[groups[k]] = rng.normal(size=bs) y = X @ true + 0.08*rng.normal(size=n) lam = 0.025 check = math_check(np.random.default_rng(9)) baseline = run_optimizer(X,y,groups,lam,False,seed=4) idea = run_optimizer(X,y,groups,lam,True,seed=4) result = {"math_check":check, "baseline":baseline, "idea":idea, "settings":{"n":n,"p":p,"steps":180,"lambda":lam,"seed":123}} with open("results.json","w") as f: json.dump(result,f,indent=2) print(json.dumps(result, indent=2)) if __name__ == "__main__": main()