import json import time import numpy as np def logdet_cholesky(H): L = np.linalg.cholesky(H) return 2.0 * np.log(np.diag(L)).sum() def low_rank_eval(U, eigvals, Pdiag, sigma2, b): # H ~= P + U diag(eigvals / sigma2) U.T pinv = 1.0 / Pdiag lam = eigvals / sigma2 # Woodbury and determinant lemma, with diagonal P. K = np.diag(1.0 / lam) + (U.T * pinv) @ U rhs = U.T @ (pinv * b) x = pinv * b - pinv * (U @ np.linalg.solve(K, rhs)) middle = np.eye(len(lam)) + (U.T * pinv) @ U @ np.diag(lam) logdet = np.log(Pdiag).sum() + np.linalg.slogdet(middle)[1] return logdet, x def exact_eval(J, Pdiag, sigma2, b): H = np.diag(Pdiag) + (J.T @ J) / sigma2 return logdet_cholesky(H), np.linalg.solve(H, b) def main(): rng = np.random.default_rng(394) d, q = 420, 90 # A deliberately concentrated spectrum: a shared data subspace is meaningful. left = rng.normal(size=(q, 16)) right = rng.normal(size=(16, d)) J = left @ right / np.sqrt(d * 16) J += 0.025 * rng.normal(size=(q, d)) G = J.T @ J evals, evecs = np.linalg.eigh(G) order = np.argsort(evals)[::-1] evals, evecs = evals[order], evecs[:, order] b = rng.normal(size=d) # Stage 1: direct numerical verification of determinant lemma and Woodbury. P = 0.7 + 0.8 * rng.random(d) sigma2 = 0.8 H = np.diag(P) + G / sigma2 Ufull = evecs[:, :q] lf, xf = low_rank_eval(Ufull, evals[:q], P, sigma2, b) le = logdet_cholesky(H) xe = np.linalg.solve(H, b) verification = { "full_rank_logdet_abs_error": float(abs(lf - le)), "full_rank_solve_relative_error": float(np.linalg.norm(xf - xe) / np.linalg.norm(xe)), "woodbury_identity_pass": bool(abs(lf - le) < 1e-8 and np.linalg.norm(xf-xe)/np.linalg.norm(xe) < 1e-8), } # 64 candidates emulate repeated prior/noise hyperparameter evaluations. n_candidates = 64 prior_scales = np.exp(np.linspace(np.log(0.35), np.log(2.5), n_candidates)) noises = np.exp(np.linspace(np.log(0.35), np.log(2.0), n_candidates)) exact_logs, exact_q = [], [] t0 = time.perf_counter() for a, noise in zip(prior_scales, noises): ld, x = exact_eval(J, np.full(d, a), noise, b) exact_logs.append(ld) exact_q.append(float(b @ x)) exact_time = time.perf_counter() - t0 results = {"verification": verification, "candidates": n_candidates, "d": d, "q": q} for r in [8, 16, 32, 64]: U = evecs[:, :r] vals = evals[:r] low_logs, low_q = [], [] t0 = time.perf_counter() for a, noise in zip(prior_scales, noises): ld, x = low_rank_eval(U, vals, np.full(d, a), noise, b) low_logs.append(ld) low_q.append(float(b @ x)) low_time = time.perf_counter() - t0 logs = np.asarray(exact_logs) approx = np.asarray(low_logs) qexact = np.asarray(exact_q) qapprox = np.asarray(low_q) # Ranking agreement is Spearman correlation, implemented from ranks. rank_e = np.argsort(np.argsort(logs)) rank_a = np.argsort(np.argsort(approx)) corr = np.corrcoef(rank_e, rank_a)[0, 1] results[f"rank_{r}"] = { "time_sec": low_time, "speedup_vs_exact": exact_time / low_time, "max_abs_logdet_error": float(np.max(np.abs(approx-logs))), "relative_logdet_rmse": float(np.sqrt(np.mean((approx-logs)**2)) / np.std(logs)), "max_relative_quadratic_error": float(np.max(np.abs(qapprox-qexact) / np.maximum(np.abs(qexact), 1e-12))), "spearman_rank_correlation": float(corr), "captured_trace_fraction": float(vals.sum() / evals.sum()), } results["exact_time_sec"] = exact_time print(json.dumps(results, indent=2)) if __name__ == "__main__": main()