Download code/run_quant_transfer.py from xedro98/quantization-as-a-transfer-constraint: direct link, hf CLI and curl.
- Browser
- Download file 7.8 kB
-
https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/run_quant_transfer.py
- Command line
-
hf download hf://xedro98/quantization-as-a-transfer-constraint/code/run_quant_transfer.py
-
curl -L -o run_quant_transfer.py https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/run_quant_transfer.py
7.8 kB
| """CRITICAL fix for Concern 1/7: a true zero-shot transfer experiment. | |
| Protocol (per config x parametrization x seed): | |
| Split FashionMNIST train (60k) into train-T (54k) and val-V (6k); test (10k) | |
| stays held out. Models always train on train-T and are evaluated on test. | |
| SOURCE : at width 64, sweep the 8-point LR grid, select LR* = argmin on val-V. | |
| TRANSFER: freeze LR*, train at widths {128,512} on train-T, report TEST loss. | |
| ORACLE : at each target width, sweep the same grid on val-V, select its argmin | |
| and report that LR's TEST loss (the "independently tuned" baseline). | |
| The paper's tabulated arc-min analyses used test-set selection; this section | |
| deliberately does NOT, and is the direct demonstration the review requests. | |
| Writes data/grid_transfer.json. | |
| """ | |
| import json | |
| import os | |
| import sys | |
| import time | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| from torch.utils.data import DataLoader | |
| from torchvision import datasets, transforms | |
| sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) | |
| import experiment as ex | |
| import experiment_quant as eq | |
| PROJ = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) | |
| DATA_DIR = os.path.join(PROJ, "data") | |
| OUT = os.path.join(PROJ, "data", "grid_transfer.json") | |
| LRS = [round(float(v), 7) for v in np.logspace(-5, -0.7, 8)] | |
| STEPS = 1000 | |
| CONFIGS = ["none", "wa4"] # transfer-claim cores | |
| TARGETS = [128, 512] | |
| SEEDS = [0, 1, 2] | |
| def split_loaders(batch_size=128): | |
| tf = transforms.Compose([transforms.ToTensor(), | |
| transforms.Normalize((0.2860,), (0.3530,))]) | |
| train = datasets.FashionMNIST(DATA_DIR, train=True, download=True, transform=tf) | |
| test = datasets.FashionMNIST(DATA_DIR, train=False, download=True, transform=tf) | |
| tv = torch.utils.data.random_split(train, [54000, 6000], | |
| generator=torch.Generator().manual_seed(0)) | |
| tr, va = tv | |
| g = torch.Generator().manual_seed(1) | |
| tl = DataLoader(tr, batch_size=batch_size, shuffle=True, generator=g, | |
| num_workers=0, drop_last=True) | |
| vl = DataLoader(va, batch_size=batch_size, shuffle=False, num_workers=0) | |
| el = DataLoader(test, batch_size=512, shuffle=False, num_workers=0) | |
| return tl, vl, el | |
| def train_eval(model, mup, lr, tr, ev, steps=STEPS): | |
| opt = ex.make_optim(model, mup, lr) | |
| it = iter(tr) | |
| for _ in range(steps): | |
| try: | |
| x, y = next(it) | |
| except StopIteration: | |
| it = iter(tr); x, y = next(it) | |
| opt.zero_grad() | |
| loss = F.cross_entropy(model(x), y) | |
| loss.backward() | |
| opt.step() | |
| return eq.evaluate(model, ev)[0] | |
| def run_one(width, lr, seed, mup, qspec, tr, el, steps=STEPS): | |
| torch.manual_seed(seed); np.random.seed(seed) | |
| model = ex.build_model(mup, width) | |
| if qspec != "none": | |
| eq.apply_quant(model, 4, True, True, per_channel=(qspec == "wa4c"), | |
| block=False) | |
| return train_eval(model, mup, lr, tr, el, steps) | |
| def main(budget_seconds=235.0): | |
| tl, vl, el = split_loaders() | |
| rows = [] | |
| if os.path.exists(OUT): | |
| with open(OUT) as f: | |
| rows = json.load(f) | |
| done = {(r["kind"], r["qspec"], r["param"], r["width"], r["lr"], r["seed"]) | |
| for r in rows} | |
| done_ot = {(r["kind"], r["qspec"], r["param"], r["width"], r["seed"]) | |
| for r in rows if r["kind"] == "oraT"} | |
| t0 = time.time(); nd0 = len(done) | |
| for qspec in CONFIGS: | |
| for mup_ in (False, True): | |
| pn = "mup" if mup_ else "sp" | |
| for seed in SEEDS: | |
| # ---- SOURCE: select LR* on val at width 64 ---- | |
| for lr in LRS: | |
| key = ("src", qspec, pn, 64, lr, seed) | |
| if key in done: | |
| continue | |
| if len(done) > nd0 and time.time() - t0 > budget_seconds: | |
| _flush(rows); return rows, "budget-remaining", len(done) | |
| v = run_one(64, lr, seed, mup_, qspec, tl, vl) # val-loss | |
| r = {"kind": "src", "qspec": qspec, "param": pn, "width": 64, | |
| "lr": lr, "seed": seed, "val_loss": round(v, 5)} | |
| rows.append(r); done.add(key) | |
| print("SRC %-4s %s w64 lr=%.0e val=%.3f (t=%.0fs)" % ( | |
| qspec, pn, lr, v, time.time() - t0), flush=True) | |
| _flush(rows) | |
| cands = [r for r in rows if (r["kind"], r["qspec"], r["param"], r["width"], r["seed"]) | |
| == ("src", qspec, pn, 64, seed)] | |
| lrstar = min(cands, key=lambda r: r["val_loss"])["lr"] | |
| # ---- TRANSFER + ORACLE at target widths ---- | |
| for w in TARGETS: | |
| tkey = ("trf", qspec, pn, w, lrstar, seed) | |
| if tkey not in done: | |
| if len(done) > nd0 and time.time() - t0 > budget_seconds: | |
| _flush(rows); return rows, "budget-remaining", len(done) | |
| tl_ = run_one(w, lrstar, seed, mup_, qspec, tl, el) # TEST loss | |
| rows.append({"kind": "trf", "qspec": qspec, "param": pn, "width": w, | |
| "lr": lrstar, "seed": seed, "test_loss": round(tl_, 5)}) | |
| done.add(tkey) | |
| print("TRF %-4s %s w%s lr=%.0e test=%.3f (t=%.0fs)" % ( | |
| qspec, pn, w, lrstar, tl_, time.time() - t0), flush=True) | |
| for lr in LRS: | |
| ok = ("ora", qspec, pn, w, lr, seed) | |
| if ok in done: | |
| continue | |
| if len(done) > nd0 and time.time() - t0 > budget_seconds: | |
| _flush(rows); return rows, "budget-remaining", len(done) | |
| o = run_one(w, lr, seed, mup_, qspec, tl, vl) # val-loss for oracle | |
| rows.append({"kind": "ora", "qspec": qspec, "param": pn, "width": w, | |
| "lr": lr, "seed": seed, "val_loss": round(o, 5)}) | |
| done.add(ok) | |
| print("ORA %-4s %s w%s lr=%.0e val=%.3f (t=%.0fs)" % ( | |
| qspec, pn, w, lr, o, time.time() - t0), flush=True) | |
| # ---- ORACLE-TEST: eval the target's val-argmin LR on TEST ---- | |
| ot = ("oraT", qspec, pn, w, seed) | |
| if ot not in done_ot: | |
| osc = [r for r in rows if (r["kind"], r["qspec"], r["param"], r["width"], r["seed"]) | |
| == ("ora", qspec, pn, w, seed)] | |
| lr_oracle = min(osc, key=lambda r: r["val_loss"])["lr"] | |
| if len(done) > nd0 and time.time() - t0 > budget_seconds: | |
| _flush(rows); return rows, "budget-remaining", len(done) | |
| te = run_one(w, lr_oracle, seed, mup_, qspec, tl, el) # TEST loss | |
| rows.append({"kind": "oraT", "qspec": qspec, "param": pn, "width": w, | |
| "lr": lr_oracle, "seed": seed, "test_loss": round(te, 5)}) | |
| done.add(("oraT", qspec, pn, w, lr_oracle, seed)) | |
| done_ot.add(ot) | |
| print("ORT %-4s %s w%s lr=%.0e test=%.3f (t=%.0fs)" % ( | |
| qspec, pn, w, lr_oracle, te, time.time() - t0), flush=True) | |
| _flush(rows) | |
| return rows, "complete", len(done) | |
| def _flush(rows): | |
| with open(OUT, "w") as f: | |
| json.dump(rows, f, indent=1) | |
| if __name__ == "__main__": | |
| b = int(sys.argv[1]) if len(sys.argv) > 1 else 235 | |
| rows, status, nd = main(b) | |
| print("status=%s done=%d" % (status, nd)) | |