SabaPivot's picture
download
raw
11.1 kB
"""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
@torch.no_grad()
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.