Centered Heavy-Tail Clipping Optimizer / bench_experiment.py
Failed on benchmark
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()