Download code/run_qdist.py from xedro98/quantization-as-a-transfer-constraint: direct link, hf CLI and curl.
- Browser
- Download file 4.68 kB
-
https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/run_qdist.py
- Command line
-
hf download hf://xedro98/quantization-as-a-transfer-constraint/code/run_qdist.py
-
curl -L -o run_qdist.py https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/run_qdist.py
4.68 kB
| """Mechanism diagnostic for Section 5. | |
| Measures, under per-tensor 4-bit (wa4) for a fixed muP model, as a function of | |
| learning rate and over a short training run, the quantities at the second hidden | |
| pre-activation: | |
| relerr mean |o - q(o)| / mean |o| (relative rounding error) | |
| jump_cells mean |o_t - o_{t-1}| / grid (per-step pre-activation movement in | |
| units of the 4-bit grid step) | |
| changed fraction of coordinates moved by > 0.25 grid step | |
| These instantiate the causal chain LR -> update/grid -> quantization distortion | |
| -> instability. Writes data/qdist.json. Rows: (width,lr,mup,qspec). | |
| """ | |
| import json | |
| import os | |
| import sys | |
| import time | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| 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", "qdist.json") | |
| def measure(width, lr, seed, mup, loaders, steps=180): | |
| torch.manual_seed(seed) | |
| np.random.seed(seed) | |
| model = eq.build_model(mup, width) | |
| # per-tensor 4-bit weight+activation wrappers (exactly the wa4 config) | |
| eq.apply_quant(model, 4, True, True, per_channel=False, block=False) | |
| opt = eq.make_optim(model, mup, lr) | |
| it = iter(loaders[0]) | |
| # instrument the per-tensor quantizer to capture the 4th quant call in each | |
| # forward (that is fc2's output = the second hidden pre-activation). | |
| orig = eq.straight_through_quant | |
| state = {"calls": 0, "prev": None, | |
| "rel": [], "jump": [], "chg": [], "grid": [], "meanabs": []} | |
| o_orig_fwd = model.forward | |
| def counting_fwd(*a, **k): | |
| state["calls"] = 0 | |
| return o_orig_fwd(*a, **k) | |
| model.forward = counting_fwd | |
| def logged(x, bits): | |
| qmax = float(2 ** (bits - 1) - 1) | |
| state["calls"] += 1 | |
| out = orig(x, bits) | |
| if state["calls"] == 4: # fc2 output pre-activation | |
| g = x.abs().amax().clamp_min(1e-12) / qmax | |
| rel = ((x - out).abs().mean() / x.abs().mean().clamp_min(1e-12)).item() | |
| p = state["prev"] | |
| if p is not None: | |
| gv = g.float() | |
| jump = ((x.detach() - p.detach()).abs().mean() / gv.clamp_min(1e-12)).item() | |
| chg = ((x.detach() - p.detach()).abs() > 0.25 * gv).float().mean().item() | |
| else: | |
| jump, chg = float("nan"), float("nan") | |
| state["prev"] = x.detach().clone() | |
| state["rel"].append(rel); state["jump"].append(jump) | |
| state["chg"].append(chg); state["grid"].append(g.item()) | |
| state["meanabs"].append(x.abs().mean().item()) | |
| return out | |
| eq.straight_through_quant = logged | |
| try: | |
| for step in range(steps): | |
| try: | |
| xb, yb = next(it) | |
| except StopIteration: | |
| it = iter(loaders[0]); xb, yb = next(it) | |
| opt.zero_grad() | |
| loss = F.cross_entropy(model(xb), yb) | |
| loss.backward() | |
| opt.step() | |
| finally: | |
| eq.straight_through_quant = orig | |
| model.forward = o_orig_fwd | |
| final = eq.evaluate(model, loaders[1])[0] | |
| # averages over logged training steps (drop the first nan jump) | |
| def avg(a): | |
| a = [v for v in a if v == v] | |
| return float(np.mean(a)) if a else float("nan") | |
| s = state | |
| return { | |
| "width": width, "lr": lr, "seed": seed, "mup": mup, "qspec": "wa4", | |
| "relerr": round(avg(s["rel"]), 5), | |
| "jump_cells": round(avg(s["jump"][1:]), 4), | |
| "changed_frac": round(avg(s["chg"][1:]), 4), | |
| "mean_abs": round(avg(s["meanabs"]), 4), | |
| "grid_step": round(avg(s["grid"]), 6), | |
| "final_loss": round(final, 5), | |
| } | |
| def main(): | |
| loaders = eq.get_loaders() | |
| LRS = [round(float(v), 7) for v in np.logspace(-5, -0.7, 8)] | |
| rows = [] | |
| if os.path.exists(OUT): | |
| with open(OUT) as f: | |
| rows = json.load(f) | |
| done = {(r["lr"], r["seed"]) for r in rows} | |
| t0 = time.time() | |
| for seed in (0, 1): | |
| for lr in LRS: | |
| if (lr, seed) in done: | |
| continue | |
| r = measure(512, lr, seed, True, loaders) # muP, width 512, wa4 | |
| rows.append(r) | |
| with open(OUT, "w") as f: | |
| json.dump(rows, f, indent=1) | |
| print("QD muP w512 lr=%.0e relerr=%.4f jump=%.2f cell chg=%.2f loss=%.3f (t=%.0fs)" | |
| % (lr, r["relerr"], r["jump_cells"], r["changed_frac"], r["final_loss"], | |
| time.time() - t0), flush=True) | |
| return rows | |
| if __name__ == "__main__": | |
| print("rows", len(main())) | |