Download code/experiment.py from xedro98/quantization-as-a-transfer-constraint: direct link, hf CLI and curl.
- Browser
- Download file 6.82 kB
-
https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/experiment.py
- Command line
-
hf download hf://xedro98/quantization-as-a-transfer-constraint/code/experiment.py
-
curl -L -o experiment.py https://huggingface.co/xedro98/quantization-as-a-transfer-constraint/resolve/main/code/experiment.py
6.82 kB
| """Original experiment: learning-rate transfer under SP vs muP on FashionMNIST MLPs. | |
| Produces: | |
| data/grid_results.json final test loss/acc + train-loss curves over an LR grid | |
| for widths x {SP, muP} | |
| data/coord_check.json activation scale (l1) vs width at t=0 and t=1 | |
| """ | |
| import argparse | |
| 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 torch.utils.data import DataLoader | |
| from torchvision import datasets, transforms | |
| from mup import MuReadout, MuAdam, set_base_shapes | |
| ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) | |
| DATA_DIR = os.path.join(ROOT, "data") | |
| DEVICE = torch.device("cpu") | |
| class MLP(nn.Module): | |
| def __init__(self, width, n_hidden=5, in_dim=784, 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(n_hidden - 1)]) | |
| 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 get_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) | |
| 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 evaluate(model, loader): | |
| model.eval() | |
| tot_loss, correct, n = 0.0, 0, 0 | |
| for x, y in loader: | |
| logits = model(x) | |
| tot_loss += F.cross_entropy(logits, y, reduction="sum").item() | |
| correct += (logits.argmax(1) == y).sum().item() | |
| n += y.numel() | |
| return tot_loss / n, correct / n | |
| def build_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 make_optim(model, mup, lr): | |
| if mup: | |
| return MuAdam(model.parameters(), lr=lr) | |
| return torch.optim.Adam(model.parameters(), lr=lr) | |
| def run_cell(width, lr, steps, seed, mup, loaders, record_curve=True): | |
| torch.manual_seed(seed) | |
| np.random.seed(seed) | |
| model = build_model(mup, width) | |
| opt = make_optim(model, mup, lr) | |
| curve = [] | |
| 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() | |
| if record_curve and (step % 250 == 0 or step == steps - 1): | |
| curve.append({"step": step, "train_loss": round(loss.item(), 5)}) | |
| test_loss, test_acc = evaluate(model, loaders[1]) | |
| dt = time.time() - t0 | |
| return { | |
| "param": "mup" if mup else "sp", | |
| "width": width, | |
| "lr": lr, | |
| "seed": seed, | |
| "steps": steps, | |
| "final_test_loss": round(test_loss, 5), | |
| "final_test_acc": round(test_acc, 5), | |
| "diverged": not math.isfinite(test_loss), | |
| "train_curve": curve, | |
| "seconds": round(dt, 1), | |
| } | |
| def coord_check(loaders, widths=(32, 128, 512, 2048), lr=1e-3, seed=0): | |
| """Average absolute activation coordinate per layer, t=0 vs t=1.""" | |
| out = [] | |
| for mup in (False, True): | |
| for width in widths: | |
| torch.manual_seed(seed) | |
| model = build_model(mup, width) | |
| acts = [] | |
| def hook(module, inp, output): | |
| acts.append(output.detach().abs().mean().item()) | |
| handles = [ | |
| model.fc_in.register_forward_hook(hook), | |
| *[h.register_forward_hook(hook) for h in model.hidden], | |
| model.readout.register_forward_hook(hook), | |
| ] | |
| x, y = next(iter(loaders[0])) | |
| model.eval() | |
| with torch.no_grad(): | |
| model(x) | |
| t0 = list(acts) | |
| acts.clear() | |
| opt = make_optim(model, mup, lr) | |
| loss = F.cross_entropy(model(x), y) | |
| opt.zero_grad() | |
| loss.backward() | |
| opt.step() | |
| with torch.no_grad(): | |
| model(x) | |
| t1 = list(acts) | |
| for h in handles: | |
| h.remove() | |
| row = {"param": "mup" if mup else "sp", "width": width, | |
| "t0": [round(v, 5) for v in t0], | |
| "t1": [round(v, 5) for v in t1]} | |
| out.append(row) | |
| print("COORD", row["param"], width, "t1_max=%.4f" % max(t1), flush=True) | |
| return out | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--quick", action="store_true") | |
| args = ap.parse_args() | |
| if args.quick: | |
| widths = [64, 256] | |
| lrs = [3e-4, 1e-3, 3e-3, 1e-2] | |
| steps = 300 | |
| else: | |
| widths = [64, 256, 1024] | |
| lrs = [round(v, 7) for v in np.logspace(-5, -1.3, 12)] | |
| steps = 1500 | |
| os.makedirs(DATA_DIR, exist_ok=True) | |
| loaders = get_loaders() | |
| print("data ready", flush=True) | |
| cc_path = os.path.join(DATA_DIR, "coord_check.json") | |
| if not os.path.exists(cc_path): | |
| rows = coord_check(loaders, | |
| widths=(32, 128, 512, 2048) if not args.quick else (32, 128)) | |
| with open(cc_path, "w") as f: | |
| json.dump(rows, f, indent=1) | |
| out_path = os.path.join(DATA_DIR, "grid_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"]) for r in results} | |
| for mup in (False, True): | |
| pname = "mup" if mup else "sp" | |
| for width in widths: | |
| for lr in lrs: | |
| key = (pname, width, lr, 0) | |
| if key in done: | |
| continue | |
| r = run_cell(width, lr, steps, 0, mup, loaders) | |
| results.append(r) | |
| with open(out_path, "w") as f: | |
| json.dump(results, f, indent=1) | |
| print("CELL {} w={} lr={:.5f} loss={} ({:.0f}s)".format( | |
| pname, width, lr, r["final_test_loss"], r["seconds"]), flush=True) | |
| print("ALL_DONE", flush=True) | |
| if __name__ == "__main__": | |
| main() | |