Download code/experiment_quant.py from xedro98/quantization-as-a-transfer-constraint: direct link, hf CLI and curl.
- Browser
- Download file 6.63 kB
-
https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/experiment_quant.py
- Command line
-
hf download hf://xedro98/quantization-as-a-transfer-constraint/code/experiment_quant.py
-
curl -L -o experiment_quant.py https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/experiment_quant.py
6.63 kB
| """ORIGINAL experiment: does muP's zero-shot learning-rate transfer survive | |
| low-precision (quantized) training, and where is the precision floor? | |
| Reuses the SP-vs-muP LR-transfer harness on FashionMNIST MLPs, adding a | |
| straight-through per-tensor symmetric quantizer applied to weights and/or | |
| activations. For each (param, quant-config, width) we sweep LR and record the | |
| final test loss, so we can measure opt-LR-vs-width stability and zero-shot | |
| transfer regret as a function of bit-precision. | |
| quant-config: 'none' (fp32 baseline), 'wa8'/'wa4' (weights+activations), 'w4' | |
| (weights-only 4-bit, to attribute the effect to weights vs activations). | |
| """ | |
| import json | |
| import math | |
| import os | |
| import time | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from experiment import (MLP, build_model, make_optim, get_loaders, DATA_DIR, DEVICE) | |
| def straight_through_quant(x, bits): | |
| """Per-tensor symmetric fake-quant with a straight-through gradient.""" | |
| qmax = float(2 ** (bits - 1) - 1) | |
| scale = x.abs().amax().clamp_min(1e-12) / qmax | |
| q = (x / scale).round().clamp(-qmax, qmax) | |
| out = q * scale | |
| return x + (out - x).detach() | |
| def quant_per_lastdim(x, bits): | |
| """Per-output-channel symmetric fake-quant (axis = last dim), straight-through.""" | |
| qmax = float(2 ** (bits - 1) - 1) | |
| d = x.dim() | |
| scale = x.abs().amax(dim=d - 1, keepdim=True).clamp_min(1e-12) / qmax | |
| q = (x / scale).round().clamp(-qmax, qmax) | |
| out = q * scale | |
| return x + (out - x).detach() | |
| def quant_per_dim0(x, bits): | |
| """Per-output-row symmetric fake-quant (axis = dim 0), straight-through.""" | |
| qmax = float(2 ** (bits - 1) - 1) | |
| scale = x.abs().amax(dim=0, keepdim=True).clamp_min(1e-12) / qmax | |
| q = (x / scale).round().clamp(-qmax, qmax) | |
| out = q * scale | |
| return x + (out - x).detach() | |
| def quant_blocklast(x, bits, block=16): | |
| """Micro-scaling style block quantizer: a scale per contiguous block over the | |
| last axis (the shared-exponent rule of MXFP/FP8-style formats), straight-through. | |
| Requires the last axis length to be a multiple of block.""" | |
| qmax = float(2 ** (bits - 1) - 1) | |
| shape = x.shape | |
| flat = x.reshape(-1, block) | |
| scale = flat.abs().amax(dim=-1, keepdim=True).clamp_min(1e-12) / qmax | |
| q = (flat / scale).round().clamp(-qmax, qmax) * scale | |
| out = q.reshape(shape) | |
| return x + (out - x).detach() | |
| def apply_quant(model, bits, quant_weights, quant_acts, per_channel=False, block=False): | |
| """Wrap each Linear so its forward quantizes the weight (and optionally output).""" | |
| for m in model.modules(): | |
| if isinstance(m, nn.Linear): | |
| orig = m.forward | |
| def fwd(x, m=m, orig=orig): | |
| W = m.weight | |
| if quant_weights: | |
| if block: | |
| Wq = quant_blocklast(W, bits) | |
| elif per_channel: | |
| Wq = quant_per_dim0(W, bits) | |
| else: | |
| Wq = straight_through_quant(W, bits) | |
| else: | |
| Wq = W | |
| out = F.linear(x, Wq, m.bias) | |
| if quant_acts: | |
| if block: | |
| out = quant_blocklast(out, bits) | |
| elif per_channel: | |
| out = quant_per_lastdim(out, bits) | |
| else: | |
| out = straight_through_quant(out, bits) | |
| return out | |
| m.forward = fwd | |
| return model | |
| def evaluate(model, loader): | |
| model.eval() | |
| tot, correct, n = 0.0, 0, 0 | |
| for x, y in loader: | |
| out = model(x) | |
| tot += F.cross_entropy(out, y, reduction="sum").item() | |
| correct += (out.argmax(1) == y).sum().item() | |
| n += y.numel() | |
| return tot / n, correct / n | |
| def run_cell(width, lr, steps, seed, mup, loaders, qspec): | |
| torch.manual_seed(seed) | |
| np.random.seed(seed) | |
| model = build_model(mup, width) | |
| bits = {"none": 32, "wa8": 8, "wa4": 4, "w4": 4, "wa4c": 4, "b8": 8, "b4": 4}[qspec] | |
| qw = qspec in ("wa8", "wa4", "w4", "wa4c", "b8", "b4") | |
| qa = qspec in ("wa8", "wa4", "wa4c", "b8", "b4") | |
| if qspec != "none": | |
| apply_quant(model, bits, qw, qa, per_channel=(qspec == "wa4c"), | |
| block=(qspec in ("b8", "b4"))) | |
| opt = make_optim(model, mup, lr) | |
| it = iter(loaders[0]) | |
| t0 = time.time() | |
| for step in range(steps): | |
| try: | |
| x, y = next(it) | |
| except StopIteration: | |
| it = iter(loaders[0]) | |
| x, y = next(it) | |
| opt.zero_grad() | |
| loss = F.cross_entropy(model(x), y) | |
| loss.backward() | |
| opt.step() | |
| loss, acc = evaluate(model, loaders[1]) | |
| return { | |
| "param": "mup" if mup else "sp", | |
| "width": width, | |
| "lr": lr, | |
| "seed": seed, | |
| "steps": steps, | |
| "qspec": qspec, | |
| "bits": bits, | |
| "final_test_loss": round(loss, 5), | |
| "final_test_acc": round(acc, 5), | |
| "diverged": not math.isfinite(loss), | |
| "seconds": round(time.time() - t0, 1), | |
| } | |
| # --- config can be shrunk for a smoke test --- | |
| import sys | |
| SMOKE = "--smoke" in sys.argv | |
| def config(): | |
| if SMOKE: | |
| return {"widths": [64, 1024], "lrs": [3e-4, 3e-3, 1e-2, 5e-2], "steps": 150} | |
| return {"widths": [64, 1024], "lrs": [round(v, 7) for v in np.logspace(-5, -0.7, 8)], | |
| "steps": 1000} | |
| def main(): | |
| QSPECS = ["none", "wa8", "wa4", "w4"] if not SMOKE else ["wa4", "w4"] | |
| cfg = config() | |
| loaders = get_loaders() | |
| out_path = os.path.join(DATA_DIR, "grid_quant_results.json") | |
| results = [] | |
| if os.path.exists(out_path): | |
| with open(out_path) as f: | |
| results = json.load(f) | |
| done = {(r["param"], r["width"], r["lr"], r["seed"], r["qspec"]) for r in results} | |
| for qspec in QSPECS: | |
| for mup in (False, True): | |
| pname = "mup" if mup else "sp" | |
| for width in cfg["widths"]: | |
| for lr in cfg["lrs"]: | |
| key = (pname, width, lr, 0, qspec) | |
| if key in done: | |
| continue | |
| r = run_cell(width, lr, cfg["steps"], 0, mup, loaders, qspec) | |
| results.append(r) | |
| with open(out_path, "w") as f: | |
| json.dump(results, f, indent=1) | |
| print("QCELL %s %-4s w=%d lr=%.3e loss=%s (%ss)" % ( | |
| pname, qspec, width, lr, r["final_test_loss"], r["seconds"]), | |
| flush=True) | |
| print("QUANT_DONE", flush=True) | |
| if __name__ == "__main__": | |
| main() | |