Finite-Excitation Orthogonal Gradient Memory / experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json
  2import math
  3from pathlib import Path
  4import numpy as np
  5
  6SEED = 2688
  7rng = np.random.default_rng(SEED)
  8
  9
 10def mgs(X, eps=1e-10):
 11    """Columns of X -> orthonormal columns, rejecting dependent columns."""
 12    qs = []
 13    residuals = []
 14    for i in range(X.shape[1]):
 15        v = X[:, i].copy()
 16        for q in qs:
 17            v -= q * (q @ v)
 18        nv = np.linalg.norm(v)
 19        residuals.append(float(nv))
 20        if nv > eps:
 21            qs.append(v / nv)
 22    Q = np.stack(qs, axis=1) if qs else np.zeros((X.shape[0], 0))
 23    return Q, np.asarray(residuals)
 24
 25
 26def orthogonality_check():
 27    d = 8
 28    # Deliberately badly conditioned but full-rank columns.
 29    U, _ = np.linalg.qr(rng.normal(size=(d, d)))
 30    V, _ = np.linalg.qr(rng.normal(size=(d, d)))
 31    X = U @ np.diag(np.geomspace(1.0, 1e-7, d)) @ V.T
 32    Q, res = mgs(X)
 33    return {
 34        "dimension": d,
 35        "sigma_min_feature_matrix": float(np.linalg.svd(X, compute_uv=False)[-1]),
 36        "mgs_rank": int(Q.shape[1]),
 37        "max_abs_QtQ_minus_I": float(np.max(np.abs(Q.T @ Q - np.eye(d)))),
 38        "min_mgs_residual": float(res.min()),
 39    }
 40
 41
 42def stability_sweep():
 43    # With a complete orthonormal basis, e <- (I-alpha I)e and norm ratio=|1-alpha|.
 44    d = 8
 45    Q = np.eye(d)
 46    e0 = rng.normal(size=d)
 47    e0 /= np.linalg.norm(e0)
 48    alphas = np.array([0.25, 0.75, 1.0, 1.25, 1.75, 1.99, 2.01, 2.5])
 49    rows = []
 50    for a in alphas:
 51        e = e0.copy()
 52        for _ in range(30):
 53            e = e - a * Q @ (Q.T @ e)
 54        ratio = float(np.linalg.norm(e)) # initial norm is 1
 55        predicted = abs(1.0 - a) ** 30
 56        rows.append({"alpha": float(a), "predicted_norm_after_30": predicted,
 57                     "observed_norm_after_30": ratio,
 58                     "observed_stable": bool(ratio < 1.0)})
 59    # empirical boundary is first alpha where amplification per step exceeds 1
 60    boundary = 2.0
 61    return {"predicted_boundary_alpha": boundary, "rows": rows}
 62
 63
 64def excitation_transition():
 65    # Memory of k orthogonal directions. Only those coordinates contract.
 66    d = 8
 67    alpha = 0.8
 68    e0 = np.ones(d) / math.sqrt(d)
 69    rows = []
 70    for k in range(d + 1):
 71        Q = np.eye(d)[:, :k]
 72        e = e0.copy()
 73        norms = []
 74        for t in range(16):
 75            norms.append(float(np.linalg.norm(e)))
 76            e = e - alpha * Q @ (Q.T @ e)
 77        # excited coordinates have exact predicted slope log|1-alpha|; unexcited remain.
 78        expected = math.sqrt((k / d) * abs(1-alpha) ** (2*15) + (d-k)/d)
 79        rows.append({"independent_directions": k,
 80                     "norm_at_start": norms[0], "norm_at_step_15": norms[-1],
 81                     "predicted_norm_at_step_15": expected,
 82                     "all_direction_contraction": bool(k == d)})
 83    return {"predicted_transition_k": d, "rows": rows}
 84
 85
 86def replay_comparison():
 87    # Same realizable linear task, with feature matrices having equal trace but
 88    # different conditioning. Raw replay uses X X^T; MGS uses Q Q^T.
 89    d = 8
 90    wstar = rng.normal(size=d)
 91    wstar /= np.linalg.norm(wstar)
 92    U, _ = np.linalg.qr(rng.normal(size=(d, d)))
 93    cases = [("well_conditioned", np.ones(d)),
 94             ("ill_conditioned", np.geomspace(1.0, 1e-3, d))]
 95    alpha = 0.8
 96    out = []
 97    for name, sing in cases:
 98        X = U @ np.diag(sing)
 99        y = X.T @ wstar
100        Q, _ = mgs(X)
101        # Normalize raw replay's mean Gramian to make it a fair gradient step;
102        # conditioning, rather than overall scale, determines the rate.
103        G = X @ X.T
104        G = G / np.trace(G) * d
105        eb = rng.normal(size=d); eb /= np.linalg.norm(eb)
106        eo = eb.copy()
107        raw_norms = []; ortho_norms = []
108        for t in range(60):
109            raw_norms.append(float(np.linalg.norm(eb)))
110            ortho_norms.append(float(np.linalg.norm(eo)))
111            eb = eb - alpha * G @ eb
112            eo = eo - alpha * Q @ (Q.T @ eo)
113        raw_rate = float(np.polyfit(np.arange(20, 60), np.log(np.maximum(raw_norms[20:], 1e-300)), 1)[0])
114        ortho_rate = float(np.polyfit(np.arange(20, 60), np.log(np.maximum(ortho_norms[20:], 1e-300)), 1)[0])
115        out.append({"case": name, "sigma_min": float(sing.min()),
116                    "raw_norm_step_10": raw_norms[10], "mgs_norm_step_10": ortho_norms[10],
117                    "raw_log_slope_late": raw_rate, "mgs_log_slope_late": ortho_rate})
118    return {"predicted_mgs_slope": math.log(abs(1-alpha)), "rows": out}
119
120
121def main():
122    result = {"seed": SEED,
123              "orthogonality": orthogonality_check(),
124              "stability_sweep": stability_sweep(),
125              "excitation_transition": excitation_transition(),
126              "replay_comparison": replay_comparison()}
127    Path("results.json").write_text(json.dumps(result, indent=2))
128    print(json.dumps(result, indent=2))
129
130if __name__ == "__main__":
131    main()