Learning-Rate-Scaled Weight Decay / experiment.py
Beats tuned baseline
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()