Moment-preserving HT compression / moment_ht_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json
  2import numpy as np
  3
  4
  5def make_basis(n=32):
  6    v = np.linspace(-1.0, 1.0, n)
  7    V1, V2 = np.meshgrid(v, v, indexing="ij")
  8    return np.stack([np.ones((n, n)), V1, V2, 0.5 * (V1**2 + V2**2)])
  9
 10
 11def moments(x, phi):
 12    return np.einsum("kij,ij->k", phi, x)
 13
 14
 15def gram(phi):
 16    return np.einsum("kij,lij->kl", phi, phi)
 17
 18
 19def conserve(target, compressed, phi, ridge=1e-12):
 20    """Smallest Frobenius-norm correction of compressed to target moments."""
 21    G = gram(phi)
 22    delta = moments(target, phi) - moments(compressed, phi)
 23    coeff = np.linalg.solve(G + ridge * np.eye(len(phi)), delta)
 24    return compressed + np.einsum("k,kij->ij", coeff, phi)
 25
 26
 27def rank_svd(x, rank):
 28    u, s, vt = np.linalg.svd(x, full_matrices=False)
 29    return (u[:, :rank] * s[:rank]) @ vt[:rank]
 30
 31
 32def zero_moment_part(x, phi):
 33    """Orthogonal projection of x onto the nullspace of all moments."""
 34    return conserve(np.zeros_like(x), x, phi)
 35
 36
 37def main():
 38    rng = np.random.default_rng(329)
 39    n, rank, steps = 32, 2, 200
 40    phi = make_basis(n)
 41
 42    # Core verification: exact moment preservation and minimum-norm property.
 43    x, y = rng.normal(size=(2, n, n))
 44    corrected = conserve(x, y, phi)
 45    residual = np.max(np.abs(moments(corrected, phi) - moments(x, phi)))
 46    optimal_distance = np.linalg.norm(corrected - y)
 47    random_distances = []
 48    for _ in range(100):
 49        q0 = zero_moment_part(rng.normal(size=(n, n)), phi)
 50        random_distances.append(np.linalg.norm(corrected + q0 - y))
 51
 52    v = np.linspace(-1, 1, n)
 53    V1, V2 = np.meshgrid(v, v, indexing="ij")
 54    # Deliberately not low rank: several separated smooth components.
 55    exact = (np.exp(-18 * ((V1 + .55)**2 + (V2 - .35)**2))
 56             + .8 * np.exp(-14 * ((V1 - .35)**2 + (V2 + .45)**2))
 57             + .15 * np.sin(7 * V1 + 2 * V2) * np.cos(5 * V2))
 58    target0 = moments(exact, phi)
 59    baseline, idea = exact.copy(), exact.copy()
 60    base_drift, idea_drift, base_err, idea_err = [], [], [], []
 61    for _ in range(steps):
 62        # Exact update conserves the four selected moments.
 63        increment = zero_moment_part(rng.normal(size=(n, n)), phi)
 64        increment *= 0.004 * np.linalg.norm(exact) / np.linalg.norm(increment)
 65        exact = exact + increment
 66        baseline = rank_svd(baseline + increment, rank)
 67        candidate = idea + increment
 68        idea = conserve(candidate, rank_svd(candidate, rank), phi)
 69        base_drift.append(np.linalg.norm(moments(baseline, phi) - target0))
 70        idea_drift.append(np.linalg.norm(moments(idea, phi) - target0))
 71        base_err.append(np.linalg.norm(baseline - exact) / np.linalg.norm(exact))
 72        idea_err.append(np.linalg.norm(idea - exact) / np.linalg.norm(exact))
 73
 74    dense = n * n
 75    factor_units = rank * (2 * n + 1)
 76    corrected_units = (rank + len(phi)) * (2 * n + 1)
 77    result = {
 78        "math_check": {
 79            "max_moment_residual": float(residual),
 80            "min_random_feasible_distance_minus_optimal": float(min(random_distances) - optimal_distance),
 81            "passed": bool(residual < 1e-9 and min(random_distances) > optimal_distance),
 82        },
 83        "experiment": {
 84            "n": n, "rank": rank, "steps": steps,
 85            "baseline_final_moment_drift_l2": float(base_drift[-1]),
 86            "idea_final_moment_drift_l2": float(idea_drift[-1]),
 87            "baseline_max_moment_drift_l2": float(max(base_drift)),
 88            "idea_max_moment_drift_l2": float(max(idea_drift)),
 89            "baseline_mean_relative_state_error": float(np.mean(base_err)),
 90            "idea_mean_relative_state_error": float(np.mean(idea_err)),
 91            "dense_storage_units": dense,
 92            "baseline_factor_storage_units": factor_units,
 93            "idea_factor_storage_upper_bound": corrected_units,
 94            "baseline_storage_ratio": factor_units / dense,
 95            "idea_storage_ratio_upper_bound": corrected_units / dense,
 96        },
 97    }
 98    print(json.dumps(result, indent=2))
 99
100
101if __name__ == "__main__":
102    main()