Adaptive Proximal Quasi-Newton Training / bench_stage2.py
Failed on benchmark
1import os, sys, json, random, math
2import numpy as np
3import torch
4import torch.nn as nn
5sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
6from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
7
8SEEDS = tuple(range(8))
9SWEEP_SEEDS = (0,1,2,3)
10EPOCHS = 15
11BATCH = 128
12# Union of every learning rate tried by either method; baseline also sweeps its
13# decisive regularization knob (AdamW weight decay).
14LRS = [1e-3, 3e-3, 1e-2]
15WDS = [0.0, 1e-4, 1e-3]
16IDEA_SETTINGS = [
17 {"lr": 1e-3, "lam": 1e-5, "eta_up": 1.35},
18 {"lr": 3e-3, "lam": 1e-5, "eta_up": 1.5},
19 {"lr": 1e-2, "lam": 1e-5, "eta_up": 1.5},
20]
21
22
23def seed_all(seed):
24 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
25 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
26
27
28def baseline_one(cfg, seed, keep_model=False):
29 seed_all(seed)
30 d = get_dataset("tabular", seed=seed, n_train=1200, n_test=400)
31 net = make_model("mlp_tiny", d["input_shape"], d["out_dim"])
32 net, metric, hist = train_model(net, d, epochs=EPOCHS, lr=cfg["lr"],
33 batch=BATCH, weight_decay=cfg["wd"], log=lambda *_: None)
34 return (metric, net, d) if keep_model else metric
35
36
37def idea_one(cfg, seed, keep_model=False):
38 seed_all(seed)
39 d = get_dataset("tabular", seed=seed, n_train=1200, n_test=400)
40 net = make_model("mlp_tiny", d["input_shape"], d["out_dim"])
41 device = "cuda" if torch.cuda.is_available() else "cpu"
42 try:
43 net = net.to(device); xtr, ytr = d["xtr"].to(device), d["ytr"].to(device)
44 xte, yte = d["xte"].to(device), d["yte"].to(device)
45 lossf = nn.MSELoss(); params = [p for p in net.parameters() if p.requires_grad]
46 H = [torch.ones_like(p) for p in params]
47 eta = float(cfg["lr"]); succ = 0; accepts = rejects = 0; eta_hist=[]; curv=[]
48 for ep in range(EPOCHS):
49 net.train(); perm = torch.randperm(len(xtr), device=device)
50 for start in range(0, len(xtr), BATCH):
51 ix = perm[start:start+BATCH]; out = net(xtr[ix]); loss = lossf(out, ytr[ix])
52 grads = torch.autograd.grad(loss, params, allow_unused=True)
53 old = [p.detach().clone() for p in params]
54 oldg = [torch.zeros_like(p) if g is None else g.detach() for p,g in zip(params,grads)]
55 trial=[]
56 for p,g,h in zip(params,oldg,H):
57 z = p.detach() - eta*h*g
58 # prox in diagonal inverse-Hessian metric for lambda ||w||_1
59 trial.append(torch.sign(z)*torch.clamp(torch.abs(z)-eta*cfg["lam"]*h, min=0.0))
60 with torch.no_grad():
61 for p,v in zip(params,trial): p.copy_(v)
62 new = lossf(net(xtr[ix]), ytr[ix]).detach()
63 pred = loss.detach() + 1e-4*sum((g*(v-o)).sum() for g,v,o in zip(oldg,trial,old))
64 ok = bool(torch.isfinite(new) and new <= pred)
65 if ok:
66 accepts += 1; succ += 1
67 newg = torch.autograd.grad(new, params, allow_unused=True) if new.requires_grad else [None]*len(params)
68 # new is detached above, so estimate curvature from a fresh graph
69 ng = torch.autograd.grad(lossf(net(xtr[ix]), ytr[ix]), params, allow_unused=True)
70 for j,(p,g,h,o,gg) in enumerate(zip(params,oldg,H,old,ng)):
71 yy = torch.zeros_like(p) if gg is None else gg.detach()-g
72 ss = p.detach()-o; den=yy*ss
73 mask=den>1e-10
74 H[j] = torch.where(mask, torch.clamp(ss/(den+1e-12), 1e-4, 100.0), h)
75 if mask.any(): curv.append(float(torch.max(torch.abs(h[mask]*yy[mask]/(ss[mask]+1e-12))).cpu()))
76 if succ >= 3: eta *= cfg["eta_up"]; succ=0
77 else:
78 with torch.no_grad():
79 for p,o in zip(params,old): p.copy_(o)
80 rejects += 1; succ=0; eta *= .5; H=[h.clamp(1e-4,100.) for h in H]
81 eta_hist.append(eta)
82 net.eval()
83 with torch.no_grad(): metric=float(lossf(net(xte),yte).cpu())
84 sig={"accepted":accepts,"rejected":rejects,"final_eta":eta,
85 "median_eta":float(np.median(eta_hist)) if eta_hist else None,
86 "max_secant_curvature":float(max(curv)) if curv else None}
87 return (metric, net, d, sig) if keep_model else metric
88 except RuntimeError:
89 # Explicit CPU fallback for constrained/shared CUDA environments.
90 torch.cuda.empty_cache() if torch.cuda.is_available() else None
91 os.environ["CUDA_VISIBLE_DEVICES"]=""
92 return idea_one(cfg, seed, keep_model)
93
94
95def main():
96 base_grid=[{"lr":lr,"wd":wd} for lr in LRS for wd in WDS]
97 base=sweep_baseline(lambda c: lambda s: baseline_one(c,s), base_grid, seeds=SWEEP_SEEDS)
98 # Required full baseline is already re-evaluated by sweep_baseline on 8 seeds.
99 best_lr=base["best_cfg"]["lr"]
100 idea_grid=[dict(x) for x in IDEA_SETTINGS]
101 idea_results=[]
102 for cfg in idea_grid:
103 r=evaluate(lambda s,cfg=cfg: idea_one(cfg,s), seeds=SEEDS)
104 idea_results.append((r,cfg))
105 idea, idea_cfg=min(idea_results,key=lambda z:z[0]["mean"])
106 # Signature comes from trained models, not an analytic toy identity.
107 sigs=[idea_one(idea_cfg,s,True)[3] for s in SEEDS]
108 med_eta=float(np.median([x["median_eta"] for x in sigs]))
109 maxcur=max(x["max_secant_curvature"] for x in sigs if x["max_secant_curvature"] is not None)
110 pred_ec=2.0/maxcur if maxcur else float("nan")
111 observed=max(x["final_eta"] for x in sigs)
112 signature={"prediction":"adaptive eta should remain below local secant stability boundary 2/kappa",
113 "predicted_eta_c_from_trained_secants":pred_ec,"observed_max_final_eta":observed,
114 "observed_median_eta":med_eta,"accepted_mean":float(np.mean([x["accepted"] for x in sigs])),
115 "rejected_mean":float(np.mean([x["rejected"] for x in sigs])),
116 "confirmed":bool(np.isfinite(pred_ec) and observed < 1.2*pred_ec)}
117 rep=make_report("tabular","mlp_tiny",base,idea,signature)
118 rep["idea"]["chosen_cfg"]=idea_cfg; rep["idea"]["sweep"]=[{"cfg":c,"mean":r["mean"],"per_seed":r["per_seed"]} for r,c in idea_results]
119 rep["notes"]="Matched Friedman#1 regression and shared mlp_tiny; optimizer is the only intervention. Baseline AdamW sweep covers lr and weight decay."
120 with open("bench_report.json","w") as f: json.dump(rep,f,indent=2)
121 print(json.dumps(rep,indent=2))
122
123if __name__=="__main__": main()