Buckets:
| """Claim 6 -- Sec. 3.3 / Fig. 4: unbalanced-OT DRO on MNIST with feature-dependent | |
| label noise (25%). | |
| Objective (20): F(theta) = lam*beta * log( (1/n) sum_i exp( lhat_i(theta)/(lam*beta) ) ), | |
| lhat_i(theta) = sup_z { l(theta; z) - lam*c(z, x_i) }. | |
| Baseline (21) : SGD on (1/n) sum_i exp( lhat_i/(lam*beta) ) (Wang et al., 2024). | |
| Proposed (22) : SGD on (1/n) sum_i [ alpha + (lam*beta/rho) log(1 + rho e^{(lhat_i-alpha)/(lam*beta)}) ]. | |
| ERM : plain cross-entropy SGD, for the noisy-fit reference curve. | |
| Paper settings kept: batch size 1, lam = beta = 1, rho = 0.1, 5 Nesterov steps for the | |
| inner maximisation, same CNN (32/64 conv + 128 hidden), official noisy-label file. | |
| REDUCED SCALE (CPU-only budget): training subset instead of the full 50k split, few | |
| epochs, fewer seeds -- see --n-train / --epochs / --seeds in the saved config. | |
| """ | |
| import argparse | |
| import itertools | |
| import json | |
| import math | |
| import os | |
| import time | |
| from multiprocessing import Pool | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as Fn | |
| ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) | |
| OUT = os.path.join(ROOT, "outputs") | |
| DATA = os.path.join(ROOT, "data") | |
| os.makedirs(OUT, exist_ok=True) | |
| MEAN, STD = 0.1307, 0.3081 | |
| class SimpleCNN(nn.Module): | |
| def __init__(self, with_alpha=False): | |
| super().__init__() | |
| self.features = nn.Sequential( | |
| nn.Conv2d(1, 32, 3, padding=1), | |
| nn.ReLU(), | |
| nn.MaxPool2d(2), | |
| nn.Conv2d(32, 64, 3, padding=1), | |
| nn.ReLU(), | |
| nn.MaxPool2d(2), | |
| ) | |
| self.classifier = nn.Sequential( | |
| nn.Flatten(), nn.Linear(64 * 7 * 7, 128), nn.ReLU(), nn.Linear(128, 10) | |
| ) | |
| if with_alpha: | |
| self.alpha = nn.Parameter(torch.randn(1)) | |
| def forward(self, x): | |
| return self.classifier(self.features(x)) | |
| def load_data(n_train, n_eval, n_val, n_test, seed=42): | |
| X = np.load(os.path.join(DATA, "mnist_X.npy")).astype(np.float32) / 255.0 | |
| y = np.load(os.path.join(DATA, "mnist_y.npy")).astype(np.int64) | |
| y_noisy = np.load(os.path.join(DATA, "feature-dependent_25_ytrain.npy")).astype( | |
| np.int64 | |
| ) | |
| X = (X - MEAN) / STD | |
| X = X.reshape(-1, 1, 28, 28) | |
| Xtr_all, ytr_clean = X[:60000], y[:60000] | |
| rng = np.random.default_rng(seed) | |
| perm = rng.permutation(60000) | |
| tr_idx, val_idx = perm[:50000], perm[50000:] | |
| noise_rate = float(np.mean(y_noisy != ytr_clean)) | |
| tr_sub = tr_idx[:n_train] | |
| ev_sub = tr_idx[:n_eval] | |
| val_sub = val_idx[:n_val] | |
| d = { | |
| "Xtr": torch.from_numpy(Xtr_all[tr_sub]), | |
| "ytr": torch.from_numpy(y_noisy[tr_sub]), | |
| "Xev": torch.from_numpy(Xtr_all[ev_sub]), | |
| "yev": torch.from_numpy(y_noisy[ev_sub]), | |
| "Xval": torch.from_numpy(Xtr_all[val_sub]), | |
| "yval": torch.from_numpy(y_noisy[val_sub]), | |
| "Xte": torch.from_numpy(X[60000 : 60000 + n_test]), | |
| "yte": torch.from_numpy(y[60000 : 60000 + n_test]), | |
| "noise_rate": noise_rate, | |
| } | |
| return d | |
| def inner_max(model, X, y, lam, n_steps=5, momentum=0.4): | |
| """sup_z l(theta;z) - lam ||z - x||^2 by 5 Nesterov steps (lr = 0.1/lam).""" | |
| lr = 0.1 / lam | |
| U = X.clone().requires_grad_(True) | |
| v = torch.zeros_like(U) | |
| for _ in range(n_steps): | |
| U_ahead = (U + momentum * v).requires_grad_(True) | |
| preds = model(U_ahead) | |
| if not torch.isfinite(preds).all(): | |
| return None | |
| loss = ( | |
| Fn.cross_entropy(preds, y, reduction="sum") | |
| - lam * (X - U_ahead).pow(2).sum() | |
| ) | |
| (g,) = torch.autograd.grad(loss, U_ahead) | |
| v = momentum * v + lr * g | |
| U = (U + v).detach() | |
| return U.detach() | |
| def exponents(model, X, y, lam): | |
| U = inner_max(model, X, y, lam) | |
| if U is None: | |
| return None | |
| preds = model(U) | |
| if not torch.isfinite(preds).all(): | |
| return None | |
| losses = Fn.cross_entropy(preds, y, reduction="none") | |
| costs = (X - U).pow(2).sum(dim=(1, 2, 3)) | |
| return losses - lam * costs | |
| def accuracy(model, X, y, bs=500): | |
| model.eval() | |
| correct = 0 | |
| for s in range(0, len(X), bs): | |
| out = model(X[s : s + bs]) | |
| correct += int((out.argmax(1) == y[s : s + bs]).sum()) | |
| model.train() | |
| return correct / len(X) | |
| def eval_objective(model, X, y, lam, lam_beta, bs=100): | |
| """F(theta) of Eq. (20) on the eval subset.""" | |
| for p in model.parameters(): | |
| p.requires_grad_(False) | |
| model.eval() | |
| ex = [] | |
| for s in range(0, len(X), bs): | |
| e = exponents(model, X[s : s + bs], y[s : s + bs], lam) | |
| if e is None: | |
| for p in model.parameters(): | |
| p.requires_grad_(True) | |
| model.train() | |
| return None | |
| ex.append(e.detach()) | |
| ex = torch.cat(ex) | |
| val = lam_beta * (torch.logsumexp(ex / lam_beta, dim=0) - math.log(len(ex))) | |
| for p in model.parameters(): | |
| p.requires_grad_(True) | |
| model.train() | |
| return float(val) | |
| def train(method, lr, seed, rho, n_train, epochs, cfg): | |
| torch.set_num_threads(1) | |
| torch.manual_seed(seed) | |
| lam = lam_beta = 1.0 | |
| d = cfg["data"] | |
| model = SimpleCNN(with_alpha=(method == "proposed")) | |
| opt = torch.optim.SGD(model.parameters(), lr=lr) | |
| n = len(d["Xtr"]) | |
| hist = [] | |
| fpe_iter = None | |
| t0 = time.time() | |
| def snapshot(step): | |
| obj = eval_objective(model, d["Xev"], d["yev"], lam, lam_beta) | |
| hist.append( | |
| { | |
| "step": step, | |
| "epoch": step / n, | |
| "F": obj, | |
| "val_acc": accuracy(model, d["Xval"], d["yval"]), | |
| "test_acc": accuracy(model, d["Xte"], d["yte"]), | |
| } | |
| ) | |
| snapshot(0) | |
| rng = np.random.default_rng(seed) | |
| step = 0 | |
| stop = False | |
| for ep in range(epochs): | |
| order = rng.permutation(n) | |
| for i in order: | |
| X = d["Xtr"][i : i + 1] if isinstance(i, (int, np.integer)) else None | |
| X = d["Xtr"][int(i)].unsqueeze(0) | |
| y = d["ytr"][int(i)].unsqueeze(0) | |
| if method == "erm": | |
| loss = Fn.cross_entropy(model(X), y) | |
| else: | |
| for p in model.parameters(): | |
| p.requires_grad_(False) | |
| U = inner_max(model, X, y, lam) | |
| for p in model.parameters(): | |
| p.requires_grad_(True) | |
| if U is None: | |
| fpe_iter = step | |
| stop = True | |
| break | |
| preds = model(U) | |
| losses = Fn.cross_entropy(preds, y, reduction="none") | |
| costs = (X - U).pow(2).sum(dim=(1, 2, 3)) | |
| ex = losses - lam * costs | |
| if method == "baseline": # Eq. (21) | |
| loss = torch.exp(ex / lam_beta).mean() | |
| else: # Eq. (22) | |
| adj = (ex - model.alpha) / lam_beta + math.log(rho) | |
| loss = (lam_beta / rho) * Fn.softplus( | |
| adj | |
| ).mean() + model.alpha.squeeze() | |
| if not torch.isfinite(loss): | |
| fpe_iter = step | |
| stop = True | |
| break | |
| opt.zero_grad() | |
| loss.backward() | |
| if not all( | |
| torch.isfinite(p.grad).all() | |
| for p in model.parameters() | |
| if p.grad is not None | |
| ): | |
| fpe_iter = step | |
| stop = True | |
| break | |
| opt.step() | |
| if not all(torch.isfinite(p).all() for p in model.parameters()): | |
| fpe_iter = step | |
| stop = True | |
| break | |
| step += 1 | |
| if step % max(1, n // 2) == 0: | |
| snapshot(step) | |
| if stop: | |
| break | |
| if not stop and (not hist or hist[-1]["step"] != step): | |
| snapshot(step) | |
| rec = { | |
| "method": method, | |
| "lr": lr, | |
| "seed": seed, | |
| "rho": rho, | |
| "n_train": n_train, | |
| "epochs": epochs, | |
| "hist": hist, | |
| "fpe_iter": fpe_iter, | |
| "diverged": fpe_iter is not None, | |
| "final_F": hist[-1]["F"] if hist else None, | |
| "seconds": time.time() - t0, | |
| } | |
| tag = f"{method}_lr{lr}_rho{rho}_seed{seed}" | |
| with open(os.path.join(OUT, "uot_traj" + cfg.get("tag", ""), f"{tag}.json"), "w") as fh: | |
| json.dump(rec, fh) | |
| print(tag, "done", rec["final_F"], "fpe", fpe_iter, flush=True) | |
| return rec | |
| _CFG = {} | |
| def _init(n_train, n_eval, n_val, n_test, tag=""): | |
| _CFG["data"] = load_data(n_train, n_eval, n_val, n_test) | |
| _CFG["tag"] = tag | |
| def _worker(t): | |
| return train(*t, cfg=_CFG) | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--procs", type=int, default=30) | |
| ap.add_argument("--seeds", type=int, default=3) | |
| ap.add_argument("--n-train", type=int, default=3000) | |
| ap.add_argument("--n-eval", type=int, default=1000) | |
| ap.add_argument("--n-val", type=int, default=2000) | |
| ap.add_argument("--n-test", type=int, default=2000) | |
| ap.add_argument("--epochs", type=int, default=2) | |
| ap.add_argument("--tag", type=str, default="") | |
| args = ap.parse_args() | |
| os.makedirs(os.path.join(OUT, "uot_traj" + args.tag), exist_ok=True) | |
| tasks = [] | |
| for lr in [1e-3, 1e-4, 1e-5, 1e-6]: | |
| for s in range(args.seeds): | |
| tasks.append(("baseline", lr, s, None, args.n_train, args.epochs)) | |
| for lr in [1e-3, 1e-4, 1e-5]: | |
| for s in range(args.seeds): | |
| tasks.append(("proposed", lr, s, 0.1, args.n_train, args.epochs)) | |
| for s in range(args.seeds): | |
| tasks.append(("erm", 1e-3, s, None, args.n_train, args.epochs)) | |
| print(f"{len(tasks)} runs", flush=True) | |
| t0 = time.time() | |
| with Pool( | |
| args.procs, | |
| initializer=_init, | |
| initargs=(args.n_train, args.n_eval, args.n_val, args.n_test, args.tag), | |
| ) as pool: | |
| recs = pool.map(_worker, tasks, chunksize=1) | |
| d = load_data(args.n_train, args.n_eval, args.n_val, args.n_test) | |
| summary = { | |
| "config": vars(args), | |
| "noise_rate": d["noise_rate"], | |
| "elapsed": time.time() - t0, | |
| "runs": recs, | |
| } | |
| with open(os.path.join(OUT, f"claim6_uot_dro{args.tag}.json"), "w") as fh: | |
| json.dump(summary, fh) | |
| for meth in ["baseline", "proposed", "erm"]: | |
| for lr in [1e-3, 1e-4, 1e-5, 1e-6]: | |
| rs = [r for r in recs if r["method"] == meth and r["lr"] == lr] | |
| if not rs: | |
| continue | |
| div = sum(r["diverged"] for r in rs) | |
| fs = [ | |
| r["final_F"] | |
| for r in rs | |
| if r["final_F"] is not None and not r["diverged"] | |
| ] | |
| print( | |
| meth, | |
| lr, | |
| "diverged", | |
| f"{div}/{len(rs)}", | |
| "mean final F", | |
| None if not fs else round(float(np.mean(fs)), 4), | |
| flush=True, | |
| ) | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 11.1 kB
- Xet hash:
- 1092d15a2465a7daf48b5a32aee12ee1bc8c6febe27bf796748ecfd0ee9bb1f8
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.