Inexact High-Order Moreau DC Optimizer / experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, time
2import numpy as np
3from scipy.optimize import minimize
4
5SEED = 7
6rng = np.random.default_rng(SEED)
7D = 24
8xtrue = np.zeros(D)
9xtrue[[1, 6, 14, 20]] = [2.2, -1.7, 1.3, -2.5]
10y = xtrue + 0.35 * rng.standard_normal(D)
11lam, eps, c, tau = 0.34, 1e-4, 0.72, 0.08
12
13def softplus(z):
14 return np.logaddexp(0.0, z)
15
16def sigmoid(z):
17 return 1.0 / (1.0 + np.exp(-np.clip(z, -50, 50)))
18
19def g(x):
20 return 0.5 * np.sum((x-y)**2) + lam * np.sum(np.sqrt(x*x + eps))
21
22def h(x):
23 return lam * tau * np.sum(softplus((np.abs(x)-c)/tau))
24
25def phi(x):
26 return g(x) - h(x)
27
28def gg(x):
29 return x-y + lam*x/np.sqrt(x*x + eps)
30
31def hh(x):
32 return lam * sigmoid((np.abs(x)-c)/tau) * np.sign(x)
33
34def rawgrad(x):
35 return gg(x) - hh(x)
36
37def penalty_grad(z, s, p, gamma):
38 d = z-s
39 n = np.linalg.norm(d)
40 if n < 1e-14:
41 return np.zeros_like(z)
42 return d * n**(p-2) / gamma
43
44def prox_obj(z, s, which, p, gamma):
45 base = g(z) if which == 'g' else h(z)
46 return base + np.linalg.norm(z-s)**p / (p*gamma)
47
48def exact_prox(s, which, p, gamma):
49 basegrad = gg if which == 'g' else hh
50 r = minimize(lambda z: prox_obj(z, s, which, p, gamma), s.copy(),
51 jac=lambda z: basegrad(z) + penalty_grad(z, s, p, gamma),
52 method='L-BFGS-B', options={'maxiter': 1000, 'ftol': 1e-13, 'gtol': 1e-10})
53 return r.x, r
54
55def inexact_prox(s, which, p, gamma, K=8):
56 z = s.copy()
57 basegrad = gg if which == 'g' else hh
58 # Fixed conservative steps make the inexactness explicit and reproducible.
59 step = 0.08 if p == 2 else 0.035
60 for _ in range(K):
61 z -= step * (basegrad(z) + penalty_grad(z, s, p, gamma))
62 return z
63
64def moreau_grad(s, p, gamma, exact=False, K=8):
65 if exact:
66 u, _ = exact_prox(s, 'g', p, gamma)
67 v, _ = exact_prox(s, 'h', p, gamma)
68 else:
69 u = inexact_prox(s, 'g', p, gamma, K)
70 v = inexact_prox(s, 'h', p, gamma, K)
71 ag = penalty_grad(s, u, p, gamma) # sign is s-u
72 ah = penalty_grad(s, v, p, gamma)
73 return ag-ah, u, v
74
75def run(method, steps=180, eta=0.18, gamma0=0.7, p=2, K=8):
76 s = np.zeros(D)
77 vals, norms, spikes = [], [], 0
78 for t in range(steps):
79 gamma = gamma0 * (0.985 ** t)
80 if method == 'raw':
81 grad = rawgrad(s)
82 else:
83 grad, _, _ = moreau_grad(s, p, gamma, exact=False, K=K)
84 n = np.linalg.norm(grad)
85 if n > 10: spikes += 1
86 s -= eta * grad
87 vals.append(float(phi(s))); norms.append(float(n))
88 if not np.all(np.isfinite(s)): break
89 return {'final_phi': vals[-1], 'best_phi': min(vals),
90 'grad_rms': float(np.sqrt(np.mean(np.array(norms)**2))),
91 'spikes_gt10': spikes, 'trajectory': vals, 'norms': norms}
92
93def math_check():
94 s = rng.normal(size=D)
95 out = {}
96 for p in (2, 4):
97 gamma = .43
98 u, ru = exact_prox(s, 'g', p, gamma)
99 v, rv = exact_prox(s, 'h', p, gamma)
100 ag = penalty_grad(s, u, p, gamma)
101 ah = penalty_grad(s, v, p, gamma)
102 rg = np.linalg.norm(gg(u) + penalty_grad(u, s, p, gamma))
103 rh = np.linalg.norm(hh(v) + penalty_grad(v, s, p, gamma))
104 # Envelope finite differences in a random direction.
105 d = rng.normal(size=D); d /= np.linalg.norm(d); delta = 1e-5
106 ep = prox_obj(u, s+delta*d, 'g', p, gamma) # objective at old prox: upper check only
107 um, _ = exact_prox(s-delta*d, 'g', p, gamma)
108 up, _ = exact_prox(s+delta*d, 'g', p, gamma)
109 em = prox_obj(um, s-delta*d, 'g', p, gamma)
110 eplus = prox_obj(up, s+delta*d, 'g', p, gamma)
111 fd = (eplus-em)/(2*delta)
112 out[str(p)] = {'g_residual': float(rg), 'h_residual': float(rh),
113 'gradient_fd_abs_error': float(abs(fd-np.dot(ag,d))),
114 'prox_success': bool(ru.success and rv.success)}
115 # Claimed lifted critical interval for g=|x|, h=2|x|: report endpoint scaling.
116 out['critical_interval'] = {str(p): float(.43**(p/(p-1)-1)) for p in (2,4)}
117 return out
118
119def main():
120 check = math_check()
121 results = {'math_check': check, 'seed': SEED,
122 'settings': {'dimension': D, 'steps': 180, 'eta': .18, 'gamma0': .7, 'K': 8}}
123 for name, p in [('raw', 2), ('moreau_p2', 2), ('moreau_p4', 4)]:
124 results[name] = run(name, p=p)
125 with open('results.json', 'w') as f:
126 json.dump(results, f, indent=2)
127 summary = {'math_check': check}
128 for name in ('raw', 'moreau_p2', 'moreau_p4'):
129 summary[name] = {m: results[name][m] for m in ('final_phi','best_phi','grad_rms','spikes_gt10')}
130 print(json.dumps(summary, indent=2))
131
132if __name__ == '__main__':
133 main()