ProCreations's picture
Claims 2 and 3: execute the paper Figure 2 and Table 1 experiments; both VERIFIED with seeds, sweeps and error bars
e812d9a verified
Raw
History Blame Contribute Delete
5.72 kB
#!/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()