Spaces:
Running
Running
Claims 2 and 3: execute the paper Figure 2 and Table 1 experiments; both VERIFIED with seeds, sweeps and error bars
e812d9a verified | #!/usr/bin/env python3 | |
| """Bayesian linear regression reproduction of Table 1 of the SGLRW paper. | |
| Setup taken verbatim from sections/F_experiments.tex of arXiv 2602.15925v1: | |
| N=1000, d=20, theta* ~ N(0,I), sigma^2=1.5, tau=1e-2, 2000 parallel particles, | |
| 10000 iterations, decaying schedule delta_t = delta_0 (1+t)^-0.55, minibatch B. | |
| Samplers: SGLD, Clipped-SGLD (R = sqrt(2 delta_t), componentwise, drift only), | |
| SGLRW (Definition 4.2 with the implementation clipping of sqrt(delta_t/2)|grad| | |
| to 1 described at the end of Section 4). | |
| Metric: KL( N(mu,Sigma) || N(mu_hat, Sigma_hat) ) between the analytic posterior | |
| and the empirical Gaussian fit of the 2000 final particles (the paper's | |
| "KL divergence between the true posterior and the empirical Gaussian fit"). | |
| """ | |
| from __future__ import annotations | |
| import argparse, json, sys, time | |
| import numpy as np | |
| N, D, SIG2, TAU = 1000, 20, 1.5, 1e-2 | |
| PARTICLES, ITERS = 2000, 10000 | |
| def make_problem(seed, design="gaussian"): | |
| rng = np.random.default_rng(1000 + seed) | |
| if design == "gaussian": | |
| X = rng.standard_normal((N, D)) | |
| elif design == "uniform": | |
| X = rng.uniform(-np.sqrt(3.0), np.sqrt(3.0), size=(N, D)) | |
| elif design == "correlated": | |
| A = rng.standard_normal((D, D)) / np.sqrt(D) | |
| X = rng.standard_normal((N, D)) @ (np.eye(D) * 0.7 + A * 0.7) | |
| elif design == "illcond": | |
| s = np.geomspace(1.0, 0.1, D) | |
| X = rng.standard_normal((N, D)) * s | |
| else: | |
| raise ValueError(design) | |
| theta_star = rng.standard_normal(D) | |
| y = X @ theta_star + np.sqrt(SIG2) * rng.standard_normal(N) | |
| P = X.T @ X / SIG2 + TAU * np.eye(D) | |
| Sig = np.linalg.inv(P) | |
| mu = Sig @ (X.T @ y) / SIG2 | |
| return X, y, mu, Sig | |
| def kl_true_vs_fit(mu, Sig, samples): | |
| """KL( N(mu,Sig) || N(mu_hat,Sig_hat) ).""" | |
| if not np.all(np.isfinite(samples)): | |
| return float("inf") | |
| mh = samples.mean(0) | |
| Sh = np.cov(samples.T) | |
| sign, ldh = np.linalg.slogdet(Sh) | |
| if sign <= 0: | |
| return float("inf") | |
| ish = np.linalg.inv(Sh) | |
| diff = mh - mu | |
| val = 0.5 * (np.trace(ish @ Sig) + diff @ ish @ diff - D + ldh - np.linalg.slogdet(Sig)[1]) | |
| return float(val) | |
| def run(sampler, B, delta0, seed, design="gaussian", iters=ITERS, particles=PARTICLES, | |
| schedule="decay", init="prior"): | |
| X, y, mu, Sig = make_problem(seed, design) | |
| rng = np.random.default_rng(90000 + seed * 17 + hash(sampler) % 1000) | |
| if init == "prior": | |
| theta = rng.standard_normal((particles, D)) | |
| elif init == "zero": | |
| theta = np.zeros((particles, D)) | |
| else: | |
| theta = mu + rng.standard_normal((particles, D)) * 0.01 | |
| scale = (N / B) / SIG2 | |
| for t in range(iters): | |
| dt = delta0 * (1.0 + t) ** (-0.55) if schedule == "decay" else delta0 | |
| idx = rng.integers(0, N, size=(particles, B)) | |
| g = TAU * theta | |
| for b in range(B): | |
| Xb = X[idx[:, b]] | |
| r = np.einsum("pd,pd->p", Xb, theta) - y[idx[:, b]] | |
| g += Xb * (r * scale)[:, None] | |
| if sampler == "sgld": | |
| theta = theta - dt * g + np.sqrt(2 * dt) * rng.standard_normal((particles, D)) | |
| elif sampler == "clipped_sgld": | |
| R = np.sqrt(2 * dt) | |
| upd = np.clip(dt * g, -R, R) | |
| theta = theta - upd + R * rng.standard_normal((particles, D)) | |
| elif sampler == "sglrw": | |
| u = np.clip(np.sqrt(dt / 2.0) * g, -1.0, 1.0) | |
| pplus = 0.5 - 0.5 * u # P[S_i = +sqrt(2 dt)] | |
| s = np.where(rng.random((particles, D)) < pplus, 1.0, -1.0) | |
| theta = theta + np.sqrt(2 * dt) * s | |
| elif sampler == "sglrw_unclipped": | |
| u = np.sqrt(dt / 2.0) * g | |
| pplus = np.clip(0.5 - 0.5 * u, 0.0, 1.0) | |
| s = np.where(rng.random((particles, D)) < pplus, 1.0, -1.0) | |
| theta = theta + np.sqrt(2 * dt) * s | |
| else: | |
| raise ValueError(sampler) | |
| if not np.all(np.isfinite(theta)): | |
| return {"kl": float("inf"), "diverged_at": t, "sampler": sampler, "B": B, | |
| "delta0": delta0, "seed": seed, "design": design} | |
| kl = kl_true_vs_fit(mu, Sig, theta) | |
| mh = theta.mean(0) | |
| Sh = np.cov(theta.T) | |
| return {"kl": kl, "diverged_at": None, "sampler": sampler, "B": B, "delta0": delta0, | |
| "seed": seed, "design": design, "schedule": schedule, "init": init, | |
| "iters": iters, "particles": particles, | |
| "mean_var_ratio": float(np.mean(np.diag(Sh) / np.diag(Sig))), | |
| "mean_shift_norm": float(np.linalg.norm(mh - mu)), | |
| "post_sd_mean": float(np.mean(np.sqrt(np.diag(Sig))))} | |
| def mc_reference(seed, design="gaussian", particles=PARTICLES): | |
| X, y, mu, Sig = make_problem(seed, design) | |
| rng = np.random.default_rng(555 + seed) | |
| L = np.linalg.cholesky(Sig) | |
| smp = mu + rng.standard_normal((particles, D)) @ L.T | |
| return kl_true_vs_fit(mu, Sig, smp) | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--jobs", required=True, help="json list of job dicts") | |
| ap.add_argument("--out", required=True) | |
| a = ap.parse_args() | |
| jobs = json.loads(a.jobs) | |
| out = [] | |
| for j in jobs: | |
| t0 = time.time() | |
| if j.get("mc_reference"): | |
| r = {"mc_reference_kl": mc_reference(j["seed"], j.get("design", "gaussian")), | |
| "seed": j["seed"], "design": j.get("design", "gaussian")} | |
| else: | |
| r = run(**{k: v for k, v in j.items() if k != "tag"}) | |
| r["seconds"] = round(time.time() - t0, 2) | |
| r["tag"] = j.get("tag", "") | |
| out.append(r) | |
| print(json.dumps(r), flush=True) | |
| with open(a.out, "w") as f: | |
| json.dump(out, f, indent=1) | |
| if __name__ == "__main__": | |
| main() | |