Buckets:
| """Meta-train L-SR1 on analytic problems (Sec. 5.1). | |
| Two experiment presets: | |
| * ``quad2`` : quadratics at N=2 (train) for the Newton-alignment study | |
| (L=8, K=16, gamma1=0.4, lambda_sec=100). | |
| * ``mixed`` : quadratics + Rosenbrock + Rastrigin at N=100 for the 30-problem | |
| performance-profile study (gamma1=0.1). | |
| The same script runs at tiny scale locally (smoke test) and at paper scale on a | |
| GPU Job -- only ``--iters``/``--batch``/``--device`` change. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import os | |
| import time | |
| import torch | |
| from lsr1 import (LSR1Config, LSR1Optimizer, Quadratic, Rastrigin, Rosenbrock, | |
| make_quadratics, init_points, rollout) | |
| def sample_batch(family, B, N, gen, device, dtype): | |
| if family == "quadratic": | |
| cond = float(torch.empty(1).uniform_(1.0, 1000.0).item()) | |
| prob = make_quadratics(B, N, cond=cond, generator=gen, device=device, | |
| dtype=dtype) | |
| x0 = init_points(B, N, "quadratic", generator=gen, device=device, | |
| dtype=dtype) | |
| elif family == "rosenbrock": | |
| prob = Rosenbrock() | |
| x0 = init_points(B, N, "rosenbrock", generator=gen, device=device, | |
| dtype=dtype) | |
| elif family == "rastrigin": | |
| prob = Rastrigin() | |
| x0 = init_points(B, N, "rastrigin", generator=gen, device=device, | |
| dtype=dtype) | |
| else: | |
| raise ValueError(family) | |
| return prob, x0 | |
| # Per-experiment presets. Quadratic settings follow Sec 5.1.1 / Table 6; the | |
| # Rosenbrock & Rastrigin settings follow Table 6 (L=32 train, K=64, lambda_sec=1, | |
| # gamma1=0.1) with a scale-invariant (relative) secant residual + trust region, | |
| # which are reproduction stabilizers for their large/ill-scaled gradients. | |
| PRESETS = { | |
| # Sec 5.1.1 Newton-alignment / projection study. | |
| "quad2": dict(N=2, families=["quadratic"], buffer_L=8, K=16, gamma1=0.4, | |
| gamma2=0.001, lambda_sec=100.0, relative_secant=False, | |
| trust_region=None), | |
| # Sec 5.1.2 performance-profile, per family: | |
| "quad100": dict(N=100, families=["quadratic"], buffer_L=16, K=32, | |
| gamma1=0.1, gamma2=0.001, lambda_sec=10.0, | |
| relative_secant=False, trust_region=10.0), | |
| "rosen100": dict(N=100, families=["rosenbrock"], buffer_L=32, K=64, | |
| gamma1=0.1, gamma2=0.001, lambda_sec=1.0, | |
| relative_secant=True, trust_region=1.0), | |
| "rastr100": dict(N=100, families=["rastrigin"], buffer_L=32, K=64, | |
| gamma1=0.1, gamma2=0.001, lambda_sec=1.0, | |
| relative_secant=True, trust_region=1.0), | |
| } | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--preset", choices=list(PRESETS), default="mixed") | |
| ap.add_argument("--N", type=int, default=0, help="override problem dim") | |
| ap.add_argument("--K", type=int, default=0, help="override unroll length") | |
| ap.add_argument("--iters", type=int, default=10000) | |
| ap.add_argument("--batch", type=int, default=128) | |
| ap.add_argument("--lr", type=float, default=1e-4) | |
| ap.add_argument("--projection", type=int, default=1, | |
| help="1 = keep secant penalty (learned PSD projection); " | |
| "0 = ablate (lambda_sec forced to 0)") | |
| ap.add_argument("--trust-region", type=float, default=-1.0, | |
| help="override preset trust region (>=0 sets it; <0 keeps " | |
| "preset; 0 disables)") | |
| ap.add_argument("--step-identity", type=int, default=1, | |
| help="1 = step preconditioner is B0=I + Sum vv^T (SR1 starts " | |
| "from identity); 0 = Sum vv^T only") | |
| ap.add_argument("--device", default="cpu") | |
| ap.add_argument("--seed", type=int, default=0) | |
| ap.add_argument("--out", default="outputs/model.pt") | |
| ap.add_argument("--log-every", type=int, default=100) | |
| ap.add_argument("--trackio", default="") | |
| args = ap.parse_args() | |
| preset = PRESETS[args.preset] | |
| device = args.device | |
| dtype = torch.float32 | |
| torch.manual_seed(args.seed) | |
| gen = torch.Generator(device=device).manual_seed(args.seed) | |
| lam = preset["lambda_sec"] if args.projection else 0.0 | |
| tr = preset["trust_region"] | |
| if args.trust_region >= 0: | |
| tr = None if args.trust_region == 0 else args.trust_region | |
| preset = {**preset, "trust_region": tr} | |
| K = args.K or preset["K"] | |
| cfg = LSR1Config(buffer_L=preset["buffer_L"], K=K, | |
| gamma1=preset["gamma1"], gamma2=preset["gamma2"], | |
| lambda_sec=lam, relative_secant=preset["relative_secant"], | |
| trust_region=preset["trust_region"], | |
| step_includes_identity=bool(args.step_identity)) | |
| opt = LSR1Optimizer(cfg).to(device) | |
| meta_opt = torch.optim.AdamW(opt.parameters(), lr=args.lr, betas=(0.9, 0.999), | |
| weight_decay=0.01) | |
| run = None | |
| if args.trackio: | |
| import trackio | |
| run = trackio.init(project=args.trackio, | |
| name=f"{args.preset}-proj{args.projection}-s{args.seed}", | |
| config={**vars(args), **preset}) | |
| N = args.N or preset["N"] | |
| families = preset["families"] | |
| hist = [] | |
| t0 = time.time() | |
| opt.train() | |
| for it in range(1, args.iters + 1): | |
| family = families[it % len(families)] | |
| prob, x0 = sample_batch(family, args.batch, N, gen, device, dtype) | |
| out = rollout(opt, prob, x0, create_graph=True) | |
| loss = out["meta_loss"] | |
| meta_opt.zero_grad() | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(opt.parameters(), 1.0) | |
| meta_opt.step() | |
| if it % args.log_every == 0 or it == 1: | |
| rec = dict(iter=it, meta_loss=float(loss.detach()), | |
| f_final=float(out["f_final"].mean()), | |
| sec_final=float(out["sec_traj"][-1].mean()), | |
| family=family, elapsed=time.time() - t0) | |
| hist.append(rec) | |
| print(f"[{it:6d}/{args.iters}] loss={rec['meta_loss']:.4e} " | |
| f"f_final={rec['f_final']:.4e} sec={rec['sec_final']:.3e} " | |
| f"({family}) {rec['elapsed']:.1f}s", flush=True) | |
| if run is not None: | |
| trackio.log({k: rec[k] for k in | |
| ("meta_loss", "f_final", "sec_final")}, step=it) | |
| os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True) | |
| torch.save({"state_dict": opt.state_dict(), "cfg": cfg.__dict__, | |
| "preset": args.preset, "args": vars(args)}, args.out) | |
| with open(args.out.replace(".pt", "_hist.json"), "w") as f: | |
| json.dump({"history": hist, "preset": preset, "args": vars(args), | |
| "wall_time_s": time.time() - t0}, f, indent=2) | |
| print(f"saved {args.out} wall={time.time()-t0:.1f}s") | |
| if run is not None: | |
| run.finish() | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 7.03 kB
- Xet hash:
- 10f8f02ef1c0e71ede6a2f3fda1b5179f05ddb43df1be2dc377363db7734c1d8
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.