diffusion-bridge-flow-matching-repro / reproduce_mfvi_variance.py
ProCreations's picture
Map exact current anchored claims to existing independent evidence
3fefb06 verified
Raw
History Blame
10.6 kB
#!/usr/bin/env python3
"""Exact identities and randomized BLR checks for MFVI predictive-variance overestimation."""
from __future__ import annotations
import hashlib
import json
import warnings
from pathlib import Path
import matplotlib.pyplot as plt
import numpy as np
from scipy.optimize import minimize
ROOT = Path(__file__).resolve().parent
OUT = ROOT / "outputs"
OUT.mkdir(exist_ok=True)
warnings.filterwarnings("ignore", message=r".*encountered in matmul", category=RuntimeWarning)
def random_rotation(rng, dimension):
q, r = np.linalg.qr(rng.normal(size=(dimension, dimension)))
q *= np.sign(np.diag(r))[None, :]
return q
def one_problem(dimension, seed):
rng = np.random.default_rng(2500000 + 1000 * dimension + seed)
samples = max(8 * dimension, 160)
condition = 10.0 ** rng.uniform(0.5, 4.0)
rotation = random_rotation(rng, dimension)
data_covariance = rotation @ np.diag(np.geomspace(condition, 1.0, dimension)) @ rotation.T
x = rng.multivariate_normal(np.zeros(dimension), data_covariance, size=samples)
x -= x.mean(axis=0)
alpha = 10.0 ** rng.uniform(-3.0, 3.0)
noise_variance = 10.0 ** rng.uniform(-3.0, 3.0)
gram = x.T @ x
precision = gram / noise_variance + alpha * np.eye(dimension)
posterior = np.linalg.inv(precision)
mfvi = np.diag(1.0 / np.diag(precision))
eigenvalues, eigenvectors = np.linalg.eigh(posterior)
v_min, v_max = eigenvectors[:, 0], eigenvectors[:, -1]
trace_margin = float(np.trace(posterior) - np.trace(mfvi))
min_direction_margin = float(v_min @ (mfvi - posterior) @ v_min)
max_direction_control = float(v_max @ (mfvi - posterior) @ v_max)
empirical_data_covariance = gram / samples
training_distribution_margin = float(np.trace(empirical_data_covariance @ (posterior - mfvi)))
identity_rhs = float(-noise_variance * alpha / samples * trace_margin)
gram_eigenvalues, gram_eigenvectors = np.linalg.eigh(gram)
first_pc = gram_eigenvectors[:, -1]
last_pc = gram_eigenvectors[:, 0]
first_pc_margin = float(first_pc @ (mfvi - posterior) @ first_pc)
last_pc_control = float(last_pc @ (mfvi - posterior) @ last_pc)
alignment = float(abs(first_pc @ v_min))
competing_scale = float(np.linalg.eigvalsh(gram / noise_variance).max() + alpha)
nonspherical_prior = np.diag(np.geomspace(competing_scale * 1e3, competing_scale * 1e6, dimension))
posterior_ns = np.linalg.inv(gram / noise_variance + nonspherical_prior)
_, eigenvectors_ns = np.linalg.eigh(posterior_ns)
nonspherical_alignment = float(abs(first_pc @ eigenvectors_ns[:, 0]))
# Reverse-KL numerical recovery of the diagonal optimum for a subset of problems.
optimizer_relative_error = None
if dimension <= 8 and seed < 6:
def objective(log_variance):
variance = np.exp(log_variance)
return 0.5 * (float(np.diag(precision) @ variance) - float(np.sum(log_variance)))
result = minimize(objective, np.zeros(dimension), method="BFGS", options={"gtol": 1e-11, "maxiter": 1000})
recovered = np.exp(result.x)
optimizer_relative_error = float(np.linalg.norm(recovered - np.diag(mfvi)) / np.linalg.norm(np.diag(mfvi)))
return {
"dimension": dimension,
"seed": seed,
"samples": samples,
"data_covariance_condition": condition,
"alpha": alpha,
"noise_variance": noise_variance,
"posterior_condition": float(np.linalg.cond(posterior)),
"trace_exact_posterior": float(np.trace(posterior)),
"trace_mfvi": float(np.trace(mfvi)),
"trace_margin": trace_margin,
"minimum_posterior_direction_margin": min_direction_margin,
"maximum_posterior_direction_control": max_direction_control,
"training_distribution_expected_exact_minus_mfvi": training_distribution_margin,
"training_distribution_identity_rhs": identity_rhs,
"identity_absolute_error": abs(training_distribution_margin - identity_rhs),
"first_pc_min_posterior_alignment": alignment,
"first_pc_mfvi_minus_exact_margin": first_pc_margin,
"last_pc_control_margin": last_pc_control,
"nonspherical_prior_first_pc_alignment_control": nonspherical_alignment,
"kl_optimizer_relative_error": optimizer_relative_error,
"gates": {
"posterior_trace_exceeds_mfvi": trace_margin >= -1e-10,
"mfvi_overestimates_min_direction": min_direction_margin >= -1e-10,
"max_direction_mutation_reverses": max_direction_control < 0,
"training_distribution_mfvi_exceeds_exact": training_distribution_margin <= 1e-10,
"training_identity": abs(training_distribution_margin - identity_rhs) <= 1e-8 * max(1.0, abs(identity_rhs)),
"first_pc_is_min_posterior_direction": alignment >= 1 - 1e-8,
"first_pc_overestimate": first_pc_margin >= -1e-10,
"last_pc_mutation_reverses": last_pc_control < 0,
"nonspherical_prior_breaks_alignment": nonspherical_alignment < 0.999,
},
}
def make_figures(records):
fig, axes = plt.subplots(1, 3, figsize=(14, 4.3))
dims = np.array([r["dimension"] for r in records])
axes[0].scatter(dims - 0.08, [r["minimum_posterior_direction_margin"] for r in records], s=14, alpha=0.6, label="min-eigenvector")
axes[0].scatter(dims + 0.08, [r["maximum_posterior_direction_control"] for r in records], s=14, alpha=0.6, label="max control")
axes[0].axhline(0, color="black", linestyle="--")
axes[0].set(xlabel="dimension", ylabel="MFVI minus exact variance", title="Directional variance inversion")
axes[0].set_yscale("symlog", linthresh=1e-10)
axes[0].legend()
x = np.array([r["training_distribution_identity_rhs"] for r in records])
y = np.array([r["training_distribution_expected_exact_minus_mfvi"] for r in records])
axes[1].scatter(x, y, s=16, alpha=0.65)
lo, hi = min(x.min(), y.min()), max(x.max(), y.max())
axes[1].plot([lo, hi], [lo, hi], "k--")
axes[1].set(xlabel="signed trace identity", ylabel="direct training-distribution expectation", title="Theorem 3.7 identity")
axes[1].set_xscale("symlog", linthresh=1e-12)
axes[1].set_yscale("symlog", linthresh=1e-12)
axes[2].scatter(dims - 0.08, [r["first_pc_mfvi_minus_exact_margin"] for r in records], s=14, alpha=0.6, label="first PC")
axes[2].scatter(dims + 0.08, [r["last_pc_control_margin"] for r in records], s=14, alpha=0.6, label="last PC control")
axes[2].axhline(0, color="black", linestyle="--")
axes[2].set(xlabel="dimension", ylabel="MFVI minus exact variance", title="Overestimation follows data concentration")
axes[2].set_yscale("symlog", linthresh=1e-10)
axes[2].legend()
fig.tight_layout()
fig.savefig(OUT / "mfvi_variance_evidence.png", dpi=180)
plt.close(fig)
def main():
records = [one_problem(d, seed) for d in (2, 3, 5, 8, 16, 32, 64) for seed in range(20)]
gate_names = list(records[0]["gates"])
gate_counts = {name: sum(r["gates"][name] for r in records) for name in gate_names}
optimizer_errors = [r["kl_optimizer_relative_error"] for r in records if r["kl_optimizer_relative_error"] is not None]
summary = {
"problems": len(records),
"dimensions": [2, 3, 5, 8, 16, 32, 64],
"alpha_range": [min(r["alpha"] for r in records), max(r["alpha"] for r in records)],
"noise_variance_range": [min(r["noise_variance"] for r in records), max(r["noise_variance"] for r in records)],
"maximum_posterior_condition": max(r["posterior_condition"] for r in records),
"gate_counts": gate_counts,
"claim_1": {
"minimum_trace_margin": min(r["trace_margin"] for r in records),
"minimum_overestimate_margin": min(r["minimum_posterior_direction_margin"] for r in records),
"maximum_direction_mutation_failures": sum(r["maximum_posterior_direction_control"] < 0 for r in records),
"kl_optimizer_checks": len(optimizer_errors),
"kl_optimizer_worst_relative_error": max(optimizer_errors),
},
"claim_2": {
"largest_expected_exact_minus_mfvi": max(r["training_distribution_expected_exact_minus_mfvi"] for r in records),
"worst_identity_absolute_error": max(r["identity_absolute_error"] for r in records),
"all_training_distribution_expectations_nonpositive": all(r["training_distribution_expected_exact_minus_mfvi"] <= 1e-10 for r in records),
},
"claim_3": {
"worst_first_pc_alignment_error": max(1 - r["first_pc_min_posterior_alignment"] for r in records),
"minimum_first_pc_overestimate": min(r["first_pc_mfvi_minus_exact_margin"] for r in records),
"last_pc_mutation_failures": sum(r["last_pc_control_margin"] < 0 for r in records),
"nonspherical_prior_alignment_breaks": sum(r["nonspherical_prior_first_pc_alignment_control"] < 0.999 for r in records),
},
"general_derivation": {
"mfvi_optimum": "S*=diag(1/diag(Sigma^-1))",
"trace_proof": "For A=Sigma^-1, A_ii <= lambda_max(A), and Schur complements give Sigma_ii >= 1/A_ii; summing yields Tr(Sigma)>=Tr(S*).",
"direction_proof": "min_i 1/A_ii >= 1/lambda_max(A)=lambda_min(Sigma), hence S* dominates Sigma along Sigma's minimum-eigenvalue direction.",
"training_distribution_identity": "tr[(X'X/N)(Sigma-S*)]=-(sigma^2 alpha/N)[Tr(Sigma)-Tr(S*)] <= 0.",
"first_pc_proof": "Sigma^-1=X'X/sigma^2+alpha I commutes with X'X; its largest data eigenvalue maps to Sigma's smallest eigenvalue.",
},
}
results = {
"paper": {
"title": "Gaussian Mean Field Variational Inference can Overestimate Predictive Variance",
"openreview": "RG7maF4bGu",
},
"protocol": {"precision": "float64", "device": "CPU", "external_compute_cost": 0},
"summary": summary,
"records": records,
}
path = OUT / "results.json"
path.write_text(json.dumps(results, indent=2))
make_figures(records)
lines = []
for item in sorted(OUT.glob("*")):
if item.is_file() and item.name != "SHA256SUMS.txt":
lines.append(f"{hashlib.sha256(item.read_bytes()).hexdigest()} {item.name}")
(OUT / "SHA256SUMS.txt").write_text("\n".join(lines) + "\n")
print(json.dumps(summary, indent=2))
print(f"RESULTS_SHA256={hashlib.sha256(path.read_bytes()).hexdigest()}")
if __name__ == "__main__":
main()