Learning-Rate-Scaled Weight Decay / experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5
  6SEED = 1333
  7random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
  8torch.set_num_threads(4)
  9
 10
 11def exact_decay(eta, eta_max, lam, steps, scaled):
 12    w = 1.0
 13    for _ in range(steps):
 14        lt = lam * min(max(eta / eta_max, 0.0), 1.0) if scaled else lam
 15        w *= (1.0 - eta * lt)
 16    return w
 17
 18
 19def schedule(step, total, peak, warmup=20):
 20    if step < warmup:
 21        return peak * (step + 1) / warmup
 22    x = (step - warmup) / max(1, total - warmup - 1)
 23    return peak * 0.5 * (1.0 + math.cos(math.pi * x))
 24
 25
 26class AdamWVariant:
 27    def __init__(self, model, peak_lr, decay, scaled):
 28        self.model = model
 29        self.peak_lr = peak_lr
 30        self.decay = decay
 31        self.scaled = scaled
 32        self.m = {id(p): torch.zeros_like(p) for p in model.parameters()}
 33        self.v = {id(p): torch.zeros_like(p) for p in model.parameters()}
 34        self.t = 0
 35
 36    @torch.no_grad()
 37    def step(self, lr):
 38        self.t += 1
 39        b1, b2, eps = 0.9, 0.999, 1e-8
 40        for p in self.model.parameters():
 41            if p.grad is None: continue
 42            g = p.grad
 43            m, v = self.m[id(p)], self.v[id(p)]
 44            m.mul_(b1).add_(g, alpha=1-b1)
 45            v.mul_(b2).addcmul_(g, g, value=1-b2)
 46            u = (m / (1-b1**self.t)) / ((v / (1-b2**self.t)).sqrt() + eps)
 47            # Decay is deliberately decoupled from the adaptive update.
 48            frac = min(max(lr / self.peak_lr, 0.0), 1.0)
 49            lam_t = self.decay * frac if self.scaled else self.decay
 50            p.mul_(1.0 - lr * lam_t).add_(u, alpha=-lr)
 51
 52
 53def train(scaled, seed=1333, total=500):
 54    torch.manual_seed(seed)
 55    n, d = 512, 12
 56    x = torch.randn(n, d)
 57    true_w = torch.randn(d, 1)
 58    y = x @ true_w + 0.15 * torch.randn(n, 1)
 59    model = torch.nn.Linear(d, 1, bias=True)
 60    opt = AdamWVariant(model, peak_lr=0.025, decay=0.08, scaled=scaled)
 61    losses, norms, lrs = [], [], []
 62    for step in range(total):
 63        lr = schedule(step, total, 0.025, 40)
 64        pred = model(x)
 65        loss = ((pred-y)**2).mean()
 66        model.zero_grad(); loss.backward(); opt.step(lr)
 67        losses.append(float(loss)); norms.append(float(torch.linalg.vector_norm(model.weight)))
 68        lrs.append(lr)
 69    return {"final_loss": losses[-1], "best_loss": min(losses),
 70            "final_norm": norms[-1], "start_norm": norms[0],
 71            "loss_at_100": losses[99], "loss_at_300": losses[299],
 72            "norm_at_300": norms[299], "losses": losses, "norms": norms, "lrs": lrs}
 73
 74
 75def main():
 76    # Prediction 1: at constant eta fraction r, scaled/constant log shrinkage ratio is r.
 77    eta_max, lam, steps = 0.02, 0.3, 200
 78    fractions = [0.1, 0.25, 0.5, 0.75, 1.0]
 79    ratio_rows = []
 80    for r in fractions:
 81        eta = eta_max*r
 82        a = exact_decay(eta, eta_max, lam, steps, True)
 83        b = exact_decay(eta, eta_max, lam, steps, False)
 84        observed = math.log(a) / math.log(b)
 85        ratio_rows.append({"r": r, "predicted_log_ratio": r, "observed_log_ratio": observed,
 86                           "relative_error": abs(observed-r)/r})
 87    # Prediction 2: cumulative scaled decay exponent is sum(lambda*eta^2/eta_max).
 88    # Use a schedule and compare exact log shrinkage to the analytical sum.
 89    T = 300; peak = 0.03; lam2 = 0.2
 90    etas = [schedule(t, T, peak, 30) for t in range(T)]
 91    exact = exact_decay(1, 1, 1, 0, True)  # only to keep helper semantics explicit
 92    w = 1.0
 93    for eta in etas: w *= 1 - lam2*eta*eta/peak
 94    predicted_log = sum(math.log(1-lam2*eta*eta/peak) for eta in etas)
 95    schedule_row = {"predicted_log_shrinkage": predicted_log, "observed_log_shrinkage": math.log(w),
 96                    "absolute_error": abs(math.log(w)-predicted_log),
 97                    "constant_log_shrinkage": sum(math.log(1-lam2*eta) for eta in etas)}
 98    # Prediction 3: in the weak-decay regime, scaled log shrinkage is linear in lambda
 99    # with slope -sum(eta^2/eta_max); measure that slope over a lambda sweep.
100    lambda_sweep = []
101    slope_sum = sum(eta*eta/peak for eta in etas)
102    for lv in [0.0, 0.02, 0.05, 0.1, 0.2, 0.4]:
103        ww = 1.0
104        for eta in etas: ww *= 1 - lv*eta*eta/peak
105        first_order = -lv*slope_sum
106        lambda_sweep.append({"lambda":lv, "observed_log_shrinkage":math.log(ww),
107                             "predicted_first_order":first_order,
108                             "normalized_observed":(math.log(ww)/lv if lv else 0.0),
109                             "predicted_slope":-slope_sum})
110    # Prediction 4: lambda=0 makes the two recursions exactly identical.
111    zero_rows=[]
112    for r in [0.1, 0.5, 1.0]:
113        a=exact_decay(peak*r, peak, 0.0, 100, True)
114        b=exact_decay(peak*r, peak, 0.0, 100, False)
115        zero_rows.append({"r":r,"absolute_difference":abs(a-b)})
116
117    seeds = [1333, 1334, 1335]
118    runs = {"baseline": [train(False, seed=s) for s in seeds],
119            "idea": [train(True, seed=s) for s in seeds]}
120    def aggregate(rs):
121        keys = ["final_loss", "loss_at_100", "loss_at_300", "final_norm", "norm_at_300"]
122        return {k:{"mean":float(np.mean([r[k] for r in rs])),
123                   "std":float(np.std([r[k] for r in rs], ddof=1))} for k in keys}
124    baseline=aggregate(runs["baseline"]); idea=aggregate(runs["idea"])
125    result = {"math_verification": {"log_ratio_sweep": ratio_rows,
126                                     "schedule_cumulative": schedule_row,
127                                     "lambda_zero_sweep": zero_rows,
128                                     "lambda_scaling_sweep": lambda_sweep},
129              "mini_experiment": {"baseline":baseline, "idea":idea, "seeds":seeds},
130              "settings":{"steps":500,"peak_lr":0.025,"nominal_decay":0.08,"data":"fixed synthetic linear regression","seed":SEED}}
131    Path("results.json").write_text(json.dumps(result, indent=2))
132    print(json.dumps(result, indent=2))
133
134if __name__ == "__main__": main()