Initial-Only Weight Decay with Tail Averaging / experiment.py
Audited (legacy)
1import argparse
2import json
3import numpy as np
4
5
6def math_check(d=20, gamma=0.5, lam=0.5, trials=200000, seed=0):
7 # Bounded coordinate samples: x = +/- e_i, with diagonal population covariance.
8 rng = np.random.default_rng(seed)
9 eig = np.linspace(0.05, 1.0, d) / d
10 # x_i = +/-sqrt(d*eig_i)e_i gives Sigma=diag(eig), while remaining bounded.
11 idx = rng.integers(0, d, size=trials)
12 signs = rng.choice([-1.0, 1.0], size=trials)
13 xscale = np.sqrt(d * eig[idx])
14 # P is diagonal for each sample, so estimate E[P^2] directly.
15 pdiag = np.full((trials, d), 1.0 - gamma * lam)
16 pdiag[np.arange(trials), idx] -= gamma * xscale**2
17 ep2 = (pdiag * pdiag).mean(axis=0)
18 A = 1.0 - gamma * (eig + lam)
19 rhs1 = (1.0 - gamma * lam) * A
20 rhs2 = (1.0 - gamma * lam) ** 2
21 # Since all matrices are diagonal, these are exact eigenvalue gaps.
22 gap_lemma = float(np.max(ep2 - rhs1))
23 gap_contraction = float(np.max(rhs1 - rhs2))
24 return {
25 "max_E_P2_minus_(1-gamma-lambda)A": gap_lemma,
26 "max_(1-gamma-lambda)A_minus_scalar_bound": gap_contraction,
27 "gamma_lambda": gamma * lam,
28 "A_min": float(A.min()),
29 "A_max": float(A.max()),
30 "passed": bool(gap_lemma <= 0.01 and gap_contraction <= 1e-12 and 0 <= gamma * lam <= 1),
31 }
32
33
34def make_data(seed, n_train=6000, n_test=3000, d=30):
35 rng = np.random.default_rng(seed)
36 eig = np.geomspace(1.0, 0.03, d)
37 # Gaussian inputs with controlled covariance and a mildly noisy target.
38 xtr = rng.normal(size=(n_train, d)) * np.sqrt(eig)
39 xte = rng.normal(size=(n_test, d)) * np.sqrt(eig)
40 teacher = rng.normal(size=d) / np.sqrt(np.arange(1, d + 1))
41 ytr = xtr @ teacher + 0.25 * rng.normal(size=n_train)
42 yte = xte @ teacher + 0.25 * rng.normal(size=n_test)
43 return xtr, ytr, xte, yte
44
45
46def train(x, y, xt, yt, seed, mode, gamma=0.08, lam=0.35, m=300, T=900, batch=32):
47 rng = np.random.default_rng(seed + 10000)
48 d = x.shape[1]
49 theta = np.zeros(d)
50 avg = np.zeros(d)
51 avg_count = 0
52 losses = []
53 update_ratios = []
54 for t in range(T):
55 ix = rng.integers(0, len(x), size=batch)
56 xb, yb = x[ix], y[ix]
57 residual = xb @ theta - yb
58 grad = xb.T @ residual / batch
59 if mode == "constant":
60 lt = lam
61 elif mode == "initial":
62 lt = lam if t < m else 0.0
63 elif mode == "none":
64 lt = 0.0
65 else:
66 raise ValueError(mode)
67 old = theta.copy()
68 theta = (1.0 - gamma * lt) * theta - gamma * grad
69 update_ratios.append(float(np.linalg.norm(theta - old) / max(1.0, np.linalg.norm(old))))
70 if 2 * m <= t < 3 * m:
71 avg += theta
72 avg_count += 1
73 if (t + 1) % 100 == 0:
74 losses.append(float(np.mean((x @ theta - y) ** 2) / 2))
75 tail = avg / max(avg_count, 1)
76 final_mse = float(np.mean((xt @ theta - yt) ** 2))
77 tail_mse = float(np.mean((xt @ tail - yt) ** 2))
78 return {
79 "final_test_mse": final_mse,
80 "tail_test_mse": tail_mse,
81 "final_train_loss": losses[-1],
82 "max_update_ratio": max(update_ratios),
83 "loss_trace": losses,
84 }
85
86
87def main():
88 ap = argparse.ArgumentParser()
89 ap.add_argument("--runs", type=int, default=8)
90 ap.add_argument("--out", default="results.json")
91 args = ap.parse_args()
92 check = math_check()
93 methods = ["constant", "initial", "none"]
94 all_results = {k: [] for k in methods}
95 for seed in range(args.runs):
96 data = make_data(seed)
97 for method in methods:
98 all_results[method].append(train(*data, seed, method))
99 summary = {}
100 for method, rows in all_results.items():
101 summary[method] = {
102 "final_test_mse_mean": float(np.mean([r["final_test_mse"] for r in rows])),
103 "final_test_mse_std": float(np.std([r["final_test_mse"] for r in rows], ddof=1)),
104 "tail_test_mse_mean": float(np.mean([r["tail_test_mse"] for r in rows])),
105 "tail_test_mse_std": float(np.std([r["tail_test_mse"] for r in rows], ddof=1)),
106 "final_train_loss_mean": float(np.mean([r["final_train_loss"] for r in rows])),
107 "max_update_ratio_mean": float(np.mean([r["max_update_ratio"] for r in rows])),
108 }
109 output = {"config": {"runs": args.runs, "gamma": 0.08, "lambda": 0.35, "m": 300, "T": 900, "batch": 32}, "math_check": check, "summary": summary, "raw": all_results}
110 with open(args.out, "w") as f:
111 json.dump(output, f, indent=2)
112 print(json.dumps({"math_check": check, "summary": summary}, indent=2))
113
114
115if __name__ == "__main__":
116 main()