Amortized low-rank Laplace hyperparameter marginalization / experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json
  2import time
  3import numpy as np
  4
  5
  6def logdet_cholesky(H):
  7    L = np.linalg.cholesky(H)
  8    return 2.0 * np.log(np.diag(L)).sum()
  9
 10
 11def low_rank_eval(U, eigvals, Pdiag, sigma2, b):
 12    # H ~= P + U diag(eigvals / sigma2) U.T
 13    pinv = 1.0 / Pdiag
 14    lam = eigvals / sigma2
 15    # Woodbury and determinant lemma, with diagonal P.
 16    K = np.diag(1.0 / lam) + (U.T * pinv) @ U
 17    rhs = U.T @ (pinv * b)
 18    x = pinv * b - pinv * (U @ np.linalg.solve(K, rhs))
 19    middle = np.eye(len(lam)) + (U.T * pinv) @ U @ np.diag(lam)
 20    logdet = np.log(Pdiag).sum() + np.linalg.slogdet(middle)[1]
 21    return logdet, x
 22
 23
 24def exact_eval(J, Pdiag, sigma2, b):
 25    H = np.diag(Pdiag) + (J.T @ J) / sigma2
 26    return logdet_cholesky(H), np.linalg.solve(H, b)
 27
 28
 29def main():
 30    rng = np.random.default_rng(394)
 31    d, q = 420, 90
 32    # A deliberately concentrated spectrum: a shared data subspace is meaningful.
 33    left = rng.normal(size=(q, 16))
 34    right = rng.normal(size=(16, d))
 35    J = left @ right / np.sqrt(d * 16)
 36    J += 0.025 * rng.normal(size=(q, d))
 37    G = J.T @ J
 38    evals, evecs = np.linalg.eigh(G)
 39    order = np.argsort(evals)[::-1]
 40    evals, evecs = evals[order], evecs[:, order]
 41    b = rng.normal(size=d)
 42
 43    # Stage 1: direct numerical verification of determinant lemma and Woodbury.
 44    P = 0.7 + 0.8 * rng.random(d)
 45    sigma2 = 0.8
 46    H = np.diag(P) + G / sigma2
 47    Ufull = evecs[:, :q]
 48    lf, xf = low_rank_eval(Ufull, evals[:q], P, sigma2, b)
 49    le = logdet_cholesky(H)
 50    xe = np.linalg.solve(H, b)
 51    verification = {
 52        "full_rank_logdet_abs_error": float(abs(lf - le)),
 53        "full_rank_solve_relative_error": float(np.linalg.norm(xf - xe) / np.linalg.norm(xe)),
 54        "woodbury_identity_pass": bool(abs(lf - le) < 1e-8 and np.linalg.norm(xf-xe)/np.linalg.norm(xe) < 1e-8),
 55    }
 56
 57    # 64 candidates emulate repeated prior/noise hyperparameter evaluations.
 58    n_candidates = 64
 59    prior_scales = np.exp(np.linspace(np.log(0.35), np.log(2.5), n_candidates))
 60    noises = np.exp(np.linspace(np.log(0.35), np.log(2.0), n_candidates))
 61    exact_logs, exact_q = [], []
 62    t0 = time.perf_counter()
 63    for a, noise in zip(prior_scales, noises):
 64        ld, x = exact_eval(J, np.full(d, a), noise, b)
 65        exact_logs.append(ld)
 66        exact_q.append(float(b @ x))
 67    exact_time = time.perf_counter() - t0
 68
 69    results = {"verification": verification, "candidates": n_candidates, "d": d, "q": q}
 70    for r in [8, 16, 32, 64]:
 71        U = evecs[:, :r]
 72        vals = evals[:r]
 73        low_logs, low_q = [], []
 74        t0 = time.perf_counter()
 75        for a, noise in zip(prior_scales, noises):
 76            ld, x = low_rank_eval(U, vals, np.full(d, a), noise, b)
 77            low_logs.append(ld)
 78            low_q.append(float(b @ x))
 79        low_time = time.perf_counter() - t0
 80        logs = np.asarray(exact_logs)
 81        approx = np.asarray(low_logs)
 82        qexact = np.asarray(exact_q)
 83        qapprox = np.asarray(low_q)
 84        # Ranking agreement is Spearman correlation, implemented from ranks.
 85        rank_e = np.argsort(np.argsort(logs))
 86        rank_a = np.argsort(np.argsort(approx))
 87        corr = np.corrcoef(rank_e, rank_a)[0, 1]
 88        results[f"rank_{r}"] = {
 89            "time_sec": low_time,
 90            "speedup_vs_exact": exact_time / low_time,
 91            "max_abs_logdet_error": float(np.max(np.abs(approx-logs))),
 92            "relative_logdet_rmse": float(np.sqrt(np.mean((approx-logs)**2)) / np.std(logs)),
 93            "max_relative_quadratic_error": float(np.max(np.abs(qapprox-qexact) / np.maximum(np.abs(qexact), 1e-12))),
 94            "spearman_rank_correlation": float(corr),
 95            "captured_trace_fraction": float(vals.sum() / evals.sum()),
 96        }
 97    results["exact_time_sec"] = exact_time
 98    print(json.dumps(results, indent=2))
 99
100
101if __name__ == "__main__":
102    main()