Parity-block curvature preconditioner / parity_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json
  2import numpy as np
  3
  4
  5def quadratic_loss(H, x):
  6    return float(0.5 * x @ H @ x)
  7
  8
  9def orthogonal_case(rng):
 10    """Reflection-invariant SPD quadratic with deliberately mismatched sectors."""
 11    signs = np.array([1, 1, 1, -1, -1, -1.0])
 12    S0 = np.diag(signs)
 13    H0 = np.zeros((6, 6))
 14    H0[:3, :3] = np.diag([1.0, 2.0, 3.0])
 15    H0[3:, 3:] = np.diag([18.0, 27.0, 45.0])
 16    U1, _ = np.linalg.qr(rng.normal(size=(3, 3)))
 17    U2, _ = np.linalg.qr(rng.normal(size=(3, 3)))
 18    Q = np.zeros((6, 6)); Q[:3, :3], Q[3:, 3:] = U1, U2
 19    H, S = Q @ H0 @ Q.T, Q @ S0 @ Q.T
 20    I = np.eye(6)
 21    Pp, Pm = (I + S) / 2, (I - S) / 2
 22    cross = np.linalg.norm(Pp @ H @ Pm)
 23    lp = np.linalg.eigvalsh(Pp @ H @ Pp)[-1]
 24    lm = np.linalg.eigvalsh(Pm @ H @ Pm)[-1]
 25
 26    # Exact one-step spectral test, including the stiffest mode in each sector.
 27    # Restrict to a basis of the sector; do not include the zero modes of P.
 28    se, V = np.linalg.eigh(S)
 29    Vp, Vm = V[:, se > 0], V[:, se < 0]
 30    def sector_stability(Vs, eta):
 31        Hs = Vs.T @ H @ Vs
 32        T = np.eye(Hs.shape[0]) - eta * Hs
 33        rho = max(abs(np.linalg.eigvalsh(T)))
 34        return float(rho), bool(rho < 1.0)
 35
 36    stability = {}
 37    for name, Vs, lam in [('plus', Vp, lp), ('minus', Vm, lm)]:
 38        ec = 2.0 / lam
 39        rho_lo, ok_lo = sector_stability(Vs, .99 * ec)
 40        rho_hi, ok_hi = sector_stability(Vs, 1.01 * ec)
 41        stability[name] = {
 42            'predicted_boundary': float(ec),
 43            'rho_at_0.99_boundary': rho_lo, 'stable_at_0.99_boundary': ok_lo,
 44            'rho_at_1.01_boundary': rho_hi, 'stable_at_1.01_boundary': ok_hi,
 45        }
 46
 47    eta_scalar = 1.8 / lm
 48    eta_p, eta_m = 1.8 / lp, 1.8 / lm
 49    x0 = rng.normal(size=6)
 50
 51    def run(sectorwise, steps=30):
 52        x = x0.copy(); losses = []
 53        for _ in range(steps):
 54            losses.append(quadratic_loss(H, x))
 55            g = H @ x
 56            if sectorwise:
 57                x -= eta_p * (Pp @ g) + eta_m * (Pm @ g)
 58            else:
 59                x -= eta_scalar * g
 60        return losses, quadratic_loss(H, x)
 61
 62    scalar, scalar_final = run(False)
 63    sector, sector_final = run(True)
 64    return {
 65        'cross_block_frobenius_norm': float(cross),
 66        'lambda_max_plus': float(lp), 'lambda_max_minus': float(lm),
 67        'curvature_ratio_minus_over_plus': float(lm / lp),
 68        'critical_eta_plus': float(2 / lp), 'critical_eta_minus': float(2 / lm),
 69        'baseline_global_eta': float(eta_scalar),
 70        'idea_eta_plus': float(eta_p), 'idea_eta_minus': float(eta_m),
 71        'baseline_loss_step_10': float(scalar[10]),
 72        'idea_loss_step_10': float(sector[10]),
 73        'baseline_final_loss': float(scalar_final),
 74        'idea_final_loss': float(sector_final),
 75        'stability_scan': stability,
 76    }
 77
 78
 79def nonorthogonal_case():
 80    """Verify that a non-orthogonal parity basis needs its metric G."""
 81    S0 = np.diag([1., 1., -1., -1.])
 82    H0 = np.diag([2., 5., 12., 20.])
 83    B = np.array([[1., .7, 0, 0], [0, 1., 0, 0],
 84                  [0, 0, 1., .5], [0, 0, 0, 1.]])
 85    S = B @ S0 @ np.linalg.inv(B)
 86    H = np.linalg.inv(B).T @ H0 @ np.linalg.inv(B)
 87    plus, minus = np.array([0, 1]), np.array([2, 3])
 88    Gp, Gm = B[:, plus].T @ B[:, plus], B[:, minus].T @ B[:, minus]
 89    Hp, Hm = B[:, plus].T @ H @ B[:, plus], B[:, minus].T @ H @ B[:, minus]
 90    gp = np.linalg.eigvals(np.linalg.solve(Gp, Hp))
 91    gm = np.linalg.eigvals(np.linalg.solve(Gm, Hm))
 92    def residual(vals, Hs, Gs):
 93        return max(abs(np.linalg.det(Hs - v * Gs)) for v in vals)
 94    return {
 95        'involution_error': float(np.linalg.norm(S @ S - np.eye(4))),
 96        'metric_nonorthogonality': float(np.linalg.norm(B.T @ B - np.eye(4))),
 97        'generalized_plus_eigenvalues': np.sort(gp).tolist(),
 98        'generalized_minus_eigenvalues': np.sort(gm).tolist(),
 99        'pencil_residual_plus': float(residual(gp, Hp, Gp)),
100        'pencil_residual_minus': float(residual(gm, Hm, Gm)),
101        'naive_euclidean_plus_eigenvalues': np.sort(np.linalg.eigvalsh(Hp)).tolist(),
102        'naive_euclidean_minus_eigenvalues': np.sort(np.linalg.eigvalsh(Hm)).tolist(),
103    }
104
105
106if __name__ == '__main__':
107    rng = np.random.default_rng(3026)
108    print(json.dumps({'orthogonal': orthogonal_case(rng),
109                      'nonorthogonal': nonorthogonal_case()}, indent=2))