Download code/run_quant_refined.py from xedro98/quantization-as-a-transfer-constraint: direct link, hf CLI and curl.
- Browser
- Download file 2.52 kB
-
https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/run_quant_refined.py
- Command line
-
hf download hf://xedro98/quantization-as-a-transfer-constraint/code/run_quant_refined.py
-
curl -L -o run_quant_refined.py https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/run_quant_refined.py
2.52 kB
| """Refined LR grid to localize the optimum better than the coarse factor-~4 grid. | |
| Per-param fine grids (factor ~1.6) around the observed optima, configs none/wa8/wa4, | |
| widths 64/512, 3 seeds. Writes data/grid_refined.json (rows carry 'seed'). | |
| """ | |
| import json | |
| import os | |
| import sys | |
| import time | |
| 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_refined.json") | |
| CONFIGS = ["none", "wa8", "wa4", "w4", "wa4c"] | |
| WIDTHS = [64, 128, 512] | |
| PER_PARAM_LR = { | |
| "sp": [round(float(v), 7) for v in (3e-4, 5e-4, 8e-4, 1.3e-3, 2e-3, 3e-3)], | |
| "mup": [round(float(v), 7) for v in (3e-3, 5e-3, 8e-3, 1.3e-2, 2e-2, 3e-2)], | |
| } | |
| STEPS = 1000 | |
| SEEDS = [0, 1, 2] | |
| 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=245.0): | |
| loaders = eq.get_loaders() | |
| rows = load() | |
| 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 PER_PARAM_LR[pn]: | |
| 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("RF %-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), 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)) | |