Download code/run_quant_cifar.py from xedro98/quantization-as-a-transfer-constraint: direct link, hf CLI and curl.
- Browser
- Download file 4.66 kB
-
https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/run_quant_cifar.py
- Command line
-
hf download hf://xedro98/quantization-as-a-transfer-constraint/code/run_quant_cifar.py
-
curl -L -o run_quant_cifar.py https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/run_quant_cifar.py
4.66 kB
| """Second-dataset cross-check on CIFAR-10. | |
| Repeats the transfer-claim configs (fp32 `none` and per-tensor 4-bit `wa4`) on | |
| the SAME four-hidden-layer MLP topology applied to CIFAR-10 (3x32x32), for widths | |
| 64 and 512, SP and muP. Corroborates that the FashionMNIST conclusions are not a | |
| dataset artifact. Writes data/grid_cifar.json. Lighter: 2 seeds, base 8-point grid. | |
| """ | |
| import json | |
| import os | |
| import sys | |
| import time | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from torch.utils.data import DataLoader | |
| from torchvision import datasets, transforms | |
| from mup import MuReadout, MuAdam, set_base_shapes | |
| 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__))) | |
| DATA_DIR = os.path.join(PROJ, "data") | |
| OUT = os.path.join(PROJ, "data", "grid_cifar.json") | |
| LRS = [round(float(v), 7) for v in np.logspace(-5, -0.7, 8)] | |
| class MLP(nn.Module): | |
| def __init__(self, width, in_dim=3072, out_dim=10, mup=False): | |
| super().__init__() | |
| Readout = MuReadout if mup else nn.Linear | |
| self.fc_in = nn.Linear(in_dim, width) | |
| self.hidden = nn.ModuleList([nn.Linear(width, width) for _ in range(4)]) | |
| self.readout = Readout(width, out_dim) | |
| def forward(self, x): | |
| x = x.flatten(1) | |
| x = F.relu(self.fc_in(x)) | |
| for h in self.hidden: | |
| x = F.relu(h(x)) | |
| return self.readout(x) | |
| def build_cifar_model(mup, width): | |
| model = MLP(width, mup=mup) | |
| if mup: | |
| base = MLP(1, mup=True) | |
| delta = MLP(2, mup=True) | |
| set_base_shapes(model, base, delta=delta) | |
| return model | |
| def get_cifar_loaders(batch_size=128): | |
| tf = transforms.Compose([ | |
| transforms.ToTensor(), | |
| transforms.Normalize((0.4914, 0.4822, 0.4465), | |
| (0.2470, 0.2435, 0.2616)), | |
| ]) | |
| train = datasets.CIFAR10(DATA_DIR, train=True, download=True, transform=tf) | |
| test = datasets.CIFAR10(DATA_DIR, train=False, download=True, transform=tf) | |
| g = torch.Generator().manual_seed(0) | |
| tl = DataLoader(train, batch_size=batch_size, shuffle=True, generator=g, | |
| num_workers=0, drop_last=True) | |
| el = DataLoader(test, batch_size=512, shuffle=False, num_workers=0) | |
| return tl, el | |
| def run_cell(width, lr, steps, seed, mup, loaders, qspec): | |
| torch.manual_seed(seed) | |
| np.random.seed(seed) | |
| model = build_cifar_model(mup, width) | |
| if qspec != "none": | |
| eq.apply_quant(model, 4 if "4" in qspec else 8, True, True, | |
| per_channel=False, block=False) | |
| opt = MuAdam(model.parameters(), lr=lr) if mup else torch.optim.Adam(model.parameters(), lr=lr) | |
| it = iter(loaders[0]) | |
| for _ 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() | |
| tl, acc = eq.evaluate(model, loaders[1]) | |
| return { | |
| "param": "mup" if mup else "sp", "width": width, "lr": lr, "seed": seed, | |
| "qspec": qspec, "final_test_loss": round(tl, 5), "final_test_acc": round(acc, 5), | |
| } | |
| def main(budget_seconds=235.0): | |
| loaders = get_cifar_loaders() | |
| rows = [] | |
| if os.path.exists(OUT): | |
| with open(OUT) as f: | |
| rows = json.load(f) | |
| done = {(r["qspec"], r["param"], r["width"], r["lr"], r["seed"]) for r in rows} | |
| t0 = time.time(); nd0 = len(done) | |
| for qspec in ("none", "wa4"): | |
| for seed in (0, 1): | |
| for mup_ in (False, True): | |
| pn = "mup" if mup_ else "sp" | |
| for w in (64, 512): | |
| for lr in LRS: | |
| key = (qspec, pn, w, lr, seed) | |
| if key in done: | |
| continue | |
| if len(done) > nd0 and time.time() - t0 > budget_seconds: | |
| return rows, "budget-remaining", len(done) | |
| r = run_cell(w, lr, 1000, seed, mup_, loaders, qspec) | |
| rows.append(r); done.add(key) | |
| with open(OUT, "w") as f: | |
| json.dump(rows, f, indent=1) | |
| print("CIF %-4s %s w=%s lr=%.0e loss=%s (t=%.0fs)" % ( | |
| qspec, pn, w, lr, r["final_test_loss"], time.time() - t0), flush=True) | |
| return rows, "complete", len(done) | |
| 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)) | |