Download code/run_quant_extra.py from xedro98/quantization-as-a-transfer-constraint: direct link, hf CLI and curl.
- Browser
- Download file 2.62 kB
-
https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/run_quant_extra.py
- Command line
-
hf download hf://xedro98/quantization-as-a-transfer-constraint/code/run_quant_extra.py
-
curl -L -o run_quant_extra.py https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/run_quant_extra.py
2.62 kB
| """Consolidated extension sweep addressing scale, seeds, and real-ish formats. | |
| Jobs (resumable, budgeted for short sandbox batches; rows carry 'seed'): | |
| A) b4 : block-scaled (MX-style, per-block shared exponent) 4-bit, widths 64/512, seeds 0-2 | |
| B) none & wa4 at width 1024, seeds 0-2 -> extends width range to x16 | |
| C) wa4 at widths 64/512, extra seeds 3-5 -> tightens the high-variance SP estimate | |
| Writes data/grid_extra.json. | |
| """ | |
| 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__))) | |
| OUT = os.path.join(PROJ, "data", "grid_extra.json") | |
| LRS = [round(float(v), 7) for v in np.logspace(-5, -0.7, 8)] | |
| STEPS = 1000 | |
| JOBS = [ | |
| ("b4", [64, 512], [0, 1, 2]), | |
| ("none", [1024], [0, 1, 2]), | |
| ("wa4", [1024], [0, 1, 2]), | |
| ("wa4", [64, 512], [3, 4, 5]), | |
| ] | |
| def load(): | |
| if os.path.exists(OUT): | |
| with open(OUT) as f: | |
| return json.load(f) | |
| return [] | |
| def save(rows): | |
| with open(OUT, "w") as f: | |
| json.dump(rows, f, indent=1) | |
| def main(budget_seconds=235.0): | |
| loaders = eq.get_loaders() | |
| rows = load() | |
| done = {(r["qspec"], r["param"], r["width"], r["lr"], r["seed"]) for r in rows} | |
| t0 = time.time() | |
| nd_start = len(done) | |
| for qspec, widths, seeds in JOBS: | |
| for seed in seeds: | |
| for mup_ in (False, True): | |
| pn = "mup" if mup_ else "sp" | |
| for w in widths: | |
| for lr in LRS: | |
| 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, len(done), 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("EX %-4s %s w=%d lr=%.0e loss=%s (t=%.0fs)" % ( | |
| qspec, pn, w, lr, r["final_test_loss"], time.time() - t0), flush=True) | |
| return rows, "complete", len(done), len(done), round(time.time() - t0, 1) | |
| if __name__ == "__main__": | |
| b = int(sys.argv[1]) if len(sys.argv) > 1 else 235 | |
| rows, status, nd, total, el = main(b) | |
| print("status=%s done=%d elapsed=%ss" % (status, nd, el)) | |