Finite-Excitation Orthogonal Gradient Memory / experiment.py
Failed on benchmark
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()