Parity-block curvature preconditioner / parity_experiment.py
Mechanism confirmed, baseline not beaten
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))