Adaptive Proximal Quasi-Newton Training / bench_stage2.py

Failed on benchmark

Raw ⬇ ZIP
  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()