"""Claims 4-5: temperature-scaled (cold) posteriors correct MFVI's predictive- variance overestimation, with optimal T < 1 in-distribution and T > 1 out-of- distribution. Bayesian basis-function regression, where everything is closed form: exact posterior N(mu, Sigma), Sigma = (Phi^T Phi / s^2 + I/alpha)^-1 MFVI optimum diagonal Gaussian with the same mean and precision diagonal, i.e. var_i = 1 / (Sigma^-1)_ii (<= Sigma_ii, so parameter variance is underestimated) tempered MFVI Sigma_q(T) = T * Sigma_q No sampling: predictive means/variances and Gaussian log-likelihoods are exact. """ import json import numpy as np RESULTS = {} def features(x, d, seed=0): rng = np.random.default_rng(seed) W = rng.normal(size=(d // 2,)) * 2.0 b = rng.uniform(0, 2 * np.pi, size=(d // 2,)) z = np.outer(x, W) + b return np.concatenate([np.cos(z), np.sin(z)], axis=1) / np.sqrt(d // 2) def fit(n=40, d=40, s=0.15, alpha=1.0, seed=0, ood_shift=4.0): rng = np.random.default_rng(seed) xtr = rng.uniform(-1, 1, size=n) Phi = features(xtr, d) wt = rng.normal(size=d) * 0.5 y = Phi @ wt + rng.normal(0, s, size=n) A = Phi.T @ Phi / s ** 2 + np.eye(d) / alpha Sig = np.linalg.inv(A) mu = Sig @ (Phi.T @ y) / s ** 2 var_mf = 1.0 / np.diag(A) # mean-field optimum xin = rng.uniform(-1, 1, size=400) # in-distribution test xoo = ood_shift + rng.uniform(-1, 1, size=400) # out-of-distribution test out = {} for name, xt in (("in", xin), ("ood", xoo)): P = features(xt, d) yt = P @ wt + rng.normal(0, s, size=xt.size) m = P @ mu v_exact = np.einsum("ij,jk,ik->i", P, Sig, P) + s ** 2 v_mf1 = (P ** 2) @ var_mf + s ** 2 out[name] = {"P": P, "yt": yt, "m": m, "v_exact": v_exact, "v_mf1": v_mf1} return out, Sig, var_mf def loglik(y, m, v): return float(np.mean(-0.5 * np.log(2 * np.pi * v) - 0.5 * (y - m) ** 2 / v)) def run(): Ts = np.exp(np.linspace(np.log(0.05), np.log(8.0), 60)) agg = {"in": [], "ood": []} ratios = [] for seed in range(12): out, Sig, var_mf = fit(seed=seed) for name in ("in", "ood"): o = out[name] base = o["v_mf1"] - 0.15 ** 2 lls = [loglik(o["yt"], o["m"], T * base + 0.15 ** 2) for T in Ts] agg[name].append(Ts[int(np.argmax(lls))]) # how much MFVI overestimates predictive variance in-distribution ratios.append(float(np.mean(out["in"]["v_mf1"] / out["in"]["v_exact"]))) tin = float(np.median(agg["in"])); tood = float(np.median(agg["ood"])) RESULTS["claims45_temperature"] = { "seeds": 12, "T_grid": [float(Ts[0]), float(Ts[-1])], "median_optimal_T_in_distribution": round(tin, 4), "median_optimal_T_ood": round(tood, 4), "frac_seeds_T_in_below_1": float(np.mean(np.array(agg["in"]) < 1.0)), "frac_seeds_T_ood_above_1": float(np.mean(np.array(agg["ood"]) > 1.0)), "mean_mfvi_predictive_variance_over_exact": round(float(np.mean(ratios)), 4), "all_optimal_T_in": [round(float(t), 3) for t in agg["in"]], "all_optimal_T_ood": [round(float(t), 3) for t in agg["ood"]]} r = RESULTS["claims45_temperature"] print(" MFVI predictive variance / exact (in-dist): %.4f" % r["mean_mfvi_predictive_variance_over_exact"]) print(" optimal T in-distribution: median %.3f (below 1 in %.0f%% of seeds)" % (tin, 100 * r["frac_seeds_T_in_below_1"])) print(" optimal T out-of-distribution: median %.3f (above 1 in %.0f%% of seeds)" % (tood, 100 * r["frac_seeds_T_ood_above_1"])) if __name__ == "__main__": run() json.dump(RESULTS, open("mfvi_results.json", "w"), indent=1)