import json, math, random from pathlib import Path import numpy as np import torch SEED = 1333 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) def exact_decay(eta, eta_max, lam, steps, scaled): w = 1.0 for _ in range(steps): lt = lam * min(max(eta / eta_max, 0.0), 1.0) if scaled else lam w *= (1.0 - eta * lt) return w def schedule(step, total, peak, warmup=20): if step < warmup: return peak * (step + 1) / warmup x = (step - warmup) / max(1, total - warmup - 1) return peak * 0.5 * (1.0 + math.cos(math.pi * x)) class AdamWVariant: def __init__(self, model, peak_lr, decay, scaled): self.model = model self.peak_lr = peak_lr self.decay = decay self.scaled = scaled self.m = {id(p): torch.zeros_like(p) for p in model.parameters()} self.v = {id(p): torch.zeros_like(p) for p in model.parameters()} self.t = 0 @torch.no_grad() def step(self, lr): self.t += 1 b1, b2, eps = 0.9, 0.999, 1e-8 for p in self.model.parameters(): if p.grad is None: continue g = p.grad m, v = self.m[id(p)], self.v[id(p)] m.mul_(b1).add_(g, alpha=1-b1) v.mul_(b2).addcmul_(g, g, value=1-b2) u = (m / (1-b1**self.t)) / ((v / (1-b2**self.t)).sqrt() + eps) # Decay is deliberately decoupled from the adaptive update. frac = min(max(lr / self.peak_lr, 0.0), 1.0) lam_t = self.decay * frac if self.scaled else self.decay p.mul_(1.0 - lr * lam_t).add_(u, alpha=-lr) def train(scaled, seed=1333, total=500): torch.manual_seed(seed) n, d = 512, 12 x = torch.randn(n, d) true_w = torch.randn(d, 1) y = x @ true_w + 0.15 * torch.randn(n, 1) model = torch.nn.Linear(d, 1, bias=True) opt = AdamWVariant(model, peak_lr=0.025, decay=0.08, scaled=scaled) losses, norms, lrs = [], [], [] for step in range(total): lr = schedule(step, total, 0.025, 40) pred = model(x) loss = ((pred-y)**2).mean() model.zero_grad(); loss.backward(); opt.step(lr) losses.append(float(loss)); norms.append(float(torch.linalg.vector_norm(model.weight))) lrs.append(lr) return {"final_loss": losses[-1], "best_loss": min(losses), "final_norm": norms[-1], "start_norm": norms[0], "loss_at_100": losses[99], "loss_at_300": losses[299], "norm_at_300": norms[299], "losses": losses, "norms": norms, "lrs": lrs} def main(): # Prediction 1: at constant eta fraction r, scaled/constant log shrinkage ratio is r. eta_max, lam, steps = 0.02, 0.3, 200 fractions = [0.1, 0.25, 0.5, 0.75, 1.0] ratio_rows = [] for r in fractions: eta = eta_max*r a = exact_decay(eta, eta_max, lam, steps, True) b = exact_decay(eta, eta_max, lam, steps, False) observed = math.log(a) / math.log(b) ratio_rows.append({"r": r, "predicted_log_ratio": r, "observed_log_ratio": observed, "relative_error": abs(observed-r)/r}) # Prediction 2: cumulative scaled decay exponent is sum(lambda*eta^2/eta_max). # Use a schedule and compare exact log shrinkage to the analytical sum. T = 300; peak = 0.03; lam2 = 0.2 etas = [schedule(t, T, peak, 30) for t in range(T)] exact = exact_decay(1, 1, 1, 0, True) # only to keep helper semantics explicit w = 1.0 for eta in etas: w *= 1 - lam2*eta*eta/peak predicted_log = sum(math.log(1-lam2*eta*eta/peak) for eta in etas) schedule_row = {"predicted_log_shrinkage": predicted_log, "observed_log_shrinkage": math.log(w), "absolute_error": abs(math.log(w)-predicted_log), "constant_log_shrinkage": sum(math.log(1-lam2*eta) for eta in etas)} # Prediction 3: in the weak-decay regime, scaled log shrinkage is linear in lambda # with slope -sum(eta^2/eta_max); measure that slope over a lambda sweep. lambda_sweep = [] slope_sum = sum(eta*eta/peak for eta in etas) for lv in [0.0, 0.02, 0.05, 0.1, 0.2, 0.4]: ww = 1.0 for eta in etas: ww *= 1 - lv*eta*eta/peak first_order = -lv*slope_sum lambda_sweep.append({"lambda":lv, "observed_log_shrinkage":math.log(ww), "predicted_first_order":first_order, "normalized_observed":(math.log(ww)/lv if lv else 0.0), "predicted_slope":-slope_sum}) # Prediction 4: lambda=0 makes the two recursions exactly identical. zero_rows=[] for r in [0.1, 0.5, 1.0]: a=exact_decay(peak*r, peak, 0.0, 100, True) b=exact_decay(peak*r, peak, 0.0, 100, False) zero_rows.append({"r":r,"absolute_difference":abs(a-b)}) seeds = [1333, 1334, 1335] runs = {"baseline": [train(False, seed=s) for s in seeds], "idea": [train(True, seed=s) for s in seeds]} def aggregate(rs): keys = ["final_loss", "loss_at_100", "loss_at_300", "final_norm", "norm_at_300"] return {k:{"mean":float(np.mean([r[k] for r in rs])), "std":float(np.std([r[k] for r in rs], ddof=1))} for k in keys} baseline=aggregate(runs["baseline"]); idea=aggregate(runs["idea"]) result = {"math_verification": {"log_ratio_sweep": ratio_rows, "schedule_cumulative": schedule_row, "lambda_zero_sweep": zero_rows, "lambda_scaling_sweep": lambda_sweep}, "mini_experiment": {"baseline":baseline, "idea":idea, "seeds":seeds}, "settings":{"steps":500,"peak_lr":0.025,"nominal_decay":0.08,"data":"fixed synthetic linear regression","seed":SEED}} Path("results.json").write_text(json.dumps(result, indent=2)) print(json.dumps(result, indent=2)) if __name__ == "__main__": main()