Inexact High-Order Moreau DC Optimizer / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()