Initial-Only Weight Decay with Tail Averaging / experiment.py

Audited (legacy)

Raw ⬇ ZIP
  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()