Amortized low-rank Laplace hyperparameter marginalization / experiment.py
Failed on benchmark
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()