txus/lsr1-artifacts / repro_bundle /code /train_analytic.py
txus's picture
download
raw
7.03 kB
"""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.