Download code/run_quant_seeds.py from xedro98/quantization-as-a-transfer-constraint: direct link, hf CLI and curl.
- Browser
- Download file 2.76 kB
-
https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/run_quant_seeds.py
- Command line
-
hf download hf://xedro98/quantization-as-a-transfer-constraint/code/run_quant_seeds.py
-
curl -L -o run_quant_seeds.py https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/run_quant_seeds.py
2.76 kB
| """Multi-seed extension: complete seeds 1 and 2 for the full 128-cell protocol | |
| (seeds 0 already in data/grid_quant_full.json). Reads the seed-0 anchor into a | |
| new data/grid_quant_seeds.json, then fills seeds 1,2 resumably in short sandbox | |
| batches (budget_seconds per call). Rows carry a 'seed' field. | |
| """ | |
| import json | |
| import os | |
| import sys | |
| import time | |
| import numpy as np | |
| import torch | |
| sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) | |
| import experiment_quant as eq | |
| PROJ = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) | |
| ANCHOR = os.path.join(PROJ, "data", "grid_quant_full.json") # seed-0 rows | |
| OUT = os.path.join(PROJ, "data", "grid_quant_seeds.json") # seeds 0,1,2 | |
| CONFIGS = ["none", "wa8", "wa4", "w4"] | |
| WIDTHS = [64, 128, 512] | |
| LRS = [round(float(v), 7) for v in np.logspace(-5, -0.7, 8)] | |
| STEPS = 1000 | |
| SEEDS = [0, 1, 2] | |
| def _init(): | |
| if os.path.exists(OUT): | |
| with open(OUT) as f: | |
| return json.load(f) | |
| rows = json.load(open(ANCHOR)) # 128 seed-0 rows | |
| for r in rows: | |
| r["seed"] = 0 | |
| with open(OUT, "w") as f: | |
| json.dump(rows, f, indent=1) | |
| return rows | |
| def save(rows): | |
| with open(OUT, "w") as f: | |
| json.dump(rows, f, indent=1) | |
| def main(budget_seconds=245.0): | |
| loaders = eq.get_loaders() | |
| rows = _init() | |
| done = {(r["qspec"], r["param"], r["width"], r["lr"], r["seed"]) for r in rows} | |
| total = 0 | |
| t0 = time.time() | |
| nd_start = len(done) | |
| for seed in SEEDS: | |
| for qspec in CONFIGS: | |
| for mup_ in (False, True): | |
| pn = "mup" if mup_ else "sp" | |
| for w in WIDTHS: | |
| for lr in LRS: | |
| total += 1 | |
| key = (qspec, pn, w, lr, seed) | |
| if key in done: | |
| continue | |
| ns = len(done) | |
| if ns > nd_start and time.time() - t0 > budget_seconds: | |
| return rows, "budget-remaining", ns, total, round(time.time() - t0, 1) | |
| r = eq.run_cell(w, lr, STEPS, seed, mup_, loaders, qspec) | |
| r["seed"] = seed | |
| rows.append(r) | |
| save(rows) | |
| done.add(key) | |
| print("S%d Q%-4s %s w=%d lr=%.0e loss=%s (t=%.0fs)" % ( | |
| seed, qspec, pn, w, lr, r["final_test_loss"], time.time() - t0), | |
| flush=True) | |
| return rows, "complete", len(done), total, round(time.time() - t0, 1) | |
| if __name__ == "__main__": | |
| b = int(sys.argv[1]) if len(sys.argv) > 1 else 245 | |
| rows, status, nd, total, el = main(b) | |
| print("status=%s done=%d/%d elapsed=%ss" % (status, nd, total, el)) | |