Centered Heavy-Tail Clipping Optimizer / bench_experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json
  2import sys
  3from pathlib import Path
  4import numpy as np
  5import torch
  6from torch.func import functional_call, grad, vmap
  7
  8sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
  9from bench import (get_dataset, make_model, sweep_baseline, evaluate,
 10                   make_report, permutation_pvalue)
 11
 12# Optimizer interventions are tested on the designated optimizer track.
 13TRACK = "tabular"
 14MODEL = "mlp_tiny"
 15SEEDS = tuple(range(8))
 16EPOCHS = 8
 17BATCH = 128
 18WEIGHT_DECAY = 0.0
 19LR_GRID = [0.0015, 0.003, 0.006, 0.012, 0.024]
 20TAU_GRID = [0.5, 1.0, 2.0]
 21
 22
 23def seed_all(seed):
 24    np.random.seed(seed)
 25    torch.manual_seed(seed)
 26    if torch.cuda.is_available():
 27        torch.cuda.manual_seed_all(seed)
 28
 29
 30def centered_clip(gs, tau):
 31    # gs: [batch, number_of_parameters], coordinate-wise robust center.
 32    c = torch.median(gs, dim=0).values
 33    r = gs - c
 34    n = torch.linalg.vector_norm(r, dim=1, keepdim=True)
 35    scale = torch.minimum(torch.ones_like(n), torch.tensor(tau, device=gs.device) / (n + 1e-12))
 36    return c + r * scale, c, n
 37
 38
 39def per_example_gradients(net, xb, yb):
 40    """Vectorized exact per-example gradients, equivalent to individual autograd."""
 41    names, params = zip(*[(n, p) for n, p in net.named_parameters() if p.requires_grad])
 42    pmap = {n: p for n, p in zip(names, params)}
 43    buffers = {n: b for n, b in net.named_buffers()}
 44    def one_loss(pm, x, y):
 45        out = functional_call(net, (pm, buffers), (x.unsqueeze(0),))
 46        return ((out - y.unsqueeze(0)) ** 2).mean()
 47    grads = vmap(grad(one_loss), in_dims=(None, 0, 0))(pmap, xb, yb)
 48    return torch.cat([grads[n].reshape(len(xb), -1) for n in names], dim=1), None
 49
 50
 51def train_one(seed, method, lr, tau, keep_model=False):
 52    seed_all(seed)
 53    ds = get_dataset(TRACK, seed=seed, n_train=400, n_test=400)
 54    net = make_model(MODEL, ds["input_shape"], ds["out_dim"])
 55    # CPU is deliberately the safe fallback; CUDA errors are caught and retried.
 56    device = "cuda" if torch.cuda.is_available() else "cpu"
 57    try:
 58        net = net.to(device)
 59        xtr, ytr = ds["xtr"].to(device), ds["ytr"].to(device)
 60        xte, yte = ds["xte"].to(device), ds["yte"].to(device)
 61        params = [p for p in net.parameters() if p.requires_grad]
 62        m = [torch.zeros_like(p) for p in params]
 63        v = [torch.zeros_like(p) for p in params]
 64        b1, b2, eps = 0.9, 0.999, 1e-8
 65        step = 0
 66        for ep in range(EPOCHS):
 67            gen = torch.Generator(device=device).manual_seed(seed * 1000 + ep)
 68            perm = torch.randperm(len(xtr), generator=gen, device=device)
 69            for start in range(0, len(xtr), BATCH):
 70                idx = perm[start:start+BATCH]
 71                gs, _ = per_example_gradients(net, xtr[idx], ytr[idx])
 72                if method == "baseline":
 73                    h = gs.mean(dim=0)
 74                    hn = torch.linalg.vector_norm(h)
 75                    h = h * torch.minimum(torch.tensor(1.0, device=device),
 76                                          torch.tensor(tau, device=device) / (hn + 1e-12))
 77                else:
 78                    h, _, _ = centered_clip(gs, tau)
 79                    h = h.mean(dim=0)
 80                off = 0
 81                step += 1
 82                with torch.no_grad():
 83                    for k, p in enumerate(params):
 84                        z = h[off:off+p.numel()].reshape_as(p); off += p.numel()
 85                        m[k].mul_(b1).add_(z, alpha=1-b1)
 86                        v[k].mul_(b2).addcmul_(z, z, value=1-b2)
 87                        mh = m[k] / (1 - b1 ** step)
 88                        vh = v[k] / (1 - b2 ** step)
 89                        p.addcdiv_(mh, torch.sqrt(vh) + eps, value=-lr)
 90        net.eval()
 91        with torch.no_grad():
 92            metric = float(((net(xte) - yte) ** 2).mean())
 93        return metric, (net, ds) if keep_model else None
 94    except (RuntimeError, torch.cuda.OutOfMemoryError):
 95        if device == "cuda":
 96            torch.cuda.empty_cache()
 97            # Retry the exact same seeded run on CPU.
 98            old = torch.cuda.is_available
 99            # Re-enter after making CUDA invisible is unreliable; use a CPU-only helper.
100            return train_one_cpu(seed, method, lr, tau, keep_model)
101        raise
102
103
104def train_one_cpu(seed, method, lr, tau, keep_model=False):
105    # CPU fallback uses the same algorithm and data/order, with no CUDA assumptions.
106    seed_all(seed)
107    ds = get_dataset(TRACK, seed=seed, n_train=400, n_test=400)
108    net = make_model(MODEL, ds["input_shape"], ds["out_dim"])
109    xtr, ytr, xte, yte = ds["xtr"], ds["ytr"], ds["xte"], ds["yte"]
110    params = list(net.parameters()); m=[torch.zeros_like(p) for p in params]; v=[torch.zeros_like(p) for p in params]
111    step=0
112    for ep in range(EPOCHS):
113        gen=torch.Generator().manual_seed(seed*1000+ep); perm=torch.randperm(len(xtr),generator=gen)
114        for start in range(0,len(xtr),BATCH):
115            gs,_=per_example_gradients(net,xtr[perm[start:start+BATCH]],ytr[perm[start:start+BATCH]])
116            if method=="baseline":
117                h=gs.mean(0); n=torch.linalg.vector_norm(h); h=h*min(1.,tau/(float(n)+1e-12))
118            else: h=centered_clip(gs,tau)[0].mean(0)
119            step+=1; off=0
120            with torch.no_grad():
121                for k,p in enumerate(params):
122                    z=h[off:off+p.numel()].reshape_as(p); off+=p.numel(); m[k]=.9*m[k]+.1*z; v[k]=.999*v[k]+.001*z*z
123                    p.addcdiv_(m[k]/(1-.9**step),torch.sqrt(v[k]/(1-.999**step))+1e-8,value=-lr)
124    with torch.no_grad(): metric=float(((net(xte)-yte)**2).mean())
125    return metric,(net,ds) if keep_model else None
126
127
128def fn(cfg, method):
129    return lambda seed: train_one(seed, method, cfg["lr"], cfg["tau"])[0]
130
131
132def signature(cfg):
133    vals=[]
134    for s in SEEDS:
135        metric, obj=train_one(s,"idea",cfg["lr"],cfg["tau"],True)
136        net,ds=obj; dev=next(net.parameters()).device; x,y=ds["xtr"][:BATCH].to(dev),ds["ytr"][:BATCH].to(dev)
137        gs,_=per_example_gradients(net,x,y); clipped,c,n=centered_clip(gs,cfg["tau"])
138        h=clipped.mean(0); raw=gs.mean(0)-c
139        cos=float(torch.dot(h-c,raw)/(torch.linalg.vector_norm(h-c)*torch.linalg.vector_norm(raw)+1e-12))
140        vals.append({"metric":metric,"observed_max_residual":float(torch.linalg.vector_norm(clipped-c,dim=1).max()),"predicted_tau":cfg["tau"],"direction_cosine":cos})
141    return {"prediction":"residual norm <= tau and common-direction cosine remains near 1","predicted_residual_bound":cfg["tau"],"observed_max_residual_mean":float(np.mean([z["observed_max_residual"] for z in vals])),"observed_direction_cosine_mean":float(np.mean([z["direction_cosine"] for z in vals])),"confirmed":bool(max(z["observed_max_residual"] for z in vals) <= cfg["tau"]*1.00001 and min(z["direction_cosine"] for z in vals)>0.99),"per_seed":vals}
142
143
144def main():
145    grid=[{"lr":lr,"tau":tau} for lr in LR_GRID for tau in TAU_GRID]
146    base=sweep_baseline(lambda cfg: fn(cfg,"baseline"),grid,seeds=SEEDS)
147    best=base["best_cfg"]
148    idea_cfgs=[best,{"lr":0.006,"tau":best["tau"]},{"lr":0.024,"tau":best["tau"]}]
149    idea_runs=[]
150    for cfg in idea_cfgs:
151        r=evaluate(fn(cfg,"idea"),seeds=SEEDS); r["cfg"]=cfg; idea_runs.append(r)
152    ibest=min(idea_runs,key=lambda x:x["mean"])
153    diffs=[a-b for a,b in zip(base["full"]["per_seed"],ibest["per_seed"])]
154    rep=make_report(TRACK,MODEL,base,ibest,{"mechanism_signature":signature(ibest["cfg"]),"idea_grid":idea_runs,"paired_diffs_baseline_minus_idea":diffs,"permutation_pvalue":permutation_pvalue(diffs)})
155    Path("bench_report.json").write_text(json.dumps(rep,indent=2))
156    print(json.dumps(rep,indent=2))
157
158if __name__=="__main__": main()