mtRgrp0u7Z-repro / code /blr_driver.py
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
2.88 kB
#!/usr/bin/env python3
"""Driver: full Bayesian-linear-regression sweep, parallel over CPU workers."""
import json, sys, time
from concurrent.futures import ProcessPoolExecutor
import numpy as np
import blr
def job(j):
t0 = time.time()
r = blr.run(**{k: v for k, v in j.items() if k != "tag"})
r["seconds"] = round(time.time() - t0, 2)
r["tag"] = j.get("tag", "")
return r
def main():
jobs = []
samplers = ["sgld", "sglrw", "clipped_sgld"]
# 1. anchored cell B=8, delta0=1e-3, 8 seeds
for s in samplers:
for seed in range(8):
jobs.append({"sampler": s, "B": 8, "delta0": 1e-3, "seed": seed, "tag": "anchor"})
# unclipped SGLRW control at the anchored cell
for seed in range(4):
jobs.append({"sampler": "sglrw_unclipped", "B": 8, "delta0": 1e-3, "seed": seed, "tag": "anchor_unclipped"})
# 2. batch/step-size grid
for B in [8, 16, 32, 64]:
for d0 in [1e-3, 1e-4]:
for s in samplers:
for seed in range(3):
if B == 8 and d0 == 1e-3 and seed < 8:
continue
jobs.append({"sampler": s, "B": B, "delta0": d0, "seed": seed, "tag": "grid"})
# 3. design-matrix robustness at the anchored cell
for design in ["uniform", "correlated", "illcond"]:
for s in samplers:
for seed in range(2):
jobs.append({"sampler": s, "B": 8, "delta0": 1e-3, "seed": seed,
"design": design, "tag": "design"})
# 4. initialisation robustness
for init in ["zero", "warm"]:
for s in samplers:
for seed in range(2):
jobs.append({"sampler": s, "B": 8, "delta0": 1e-3, "seed": seed,
"init": init, "tag": "init"})
# 5. constant-step schedule
for d0 in [1e-3, 1e-4, 1e-5]:
for s in samplers:
for seed in range(2):
jobs.append({"sampler": s, "B": 8, "delta0": d0, "seed": seed,
"schedule": "const", "tag": "const"})
print(f"{len(jobs)} jobs", flush=True)
out = []
with ProcessPoolExecutor(max_workers=6) as ex:
for i, r in enumerate(ex.map(job, jobs)):
out.append(r)
if i % 10 == 0:
print(i, r["tag"], r["sampler"], r["B"], r["delta0"], round(r["kl"], 4), flush=True)
with open("/Users/sshpro/icml-queue/fix8-work/blr_results.json", "w") as f:
json.dump(out, f, indent=1)
# MC reference
refs = [{"seed": s, "design": "gaussian", "mc_reference_kl": blr.mc_reference(s)} for s in range(8)]
with open("/Users/sshpro/icml-queue/fix8-work/blr_results.json", "w") as f:
json.dump({"runs": out, "mc_reference": refs}, f, indent=1)
print("DONE", len(out), flush=True)
if __name__ == "__main__":
main()