Adaptive CUR Neural Layer / adaptive_cur.py
Mechanism failed
1import json
2import numpy as np
3
4
5def pinv_truncated(C, rel_tol=1e-10):
6 u, s, vt = np.linalg.svd(C, full_matrices=False)
7 if len(s) == 0 or s[0] == 0:
8 return np.zeros((C.shape[1], C.shape[0]))
9 keep = s > rel_tol * s[0]
10 return (vt[keep].T / s[keep]) @ u[:, keep].T
11
12
13def cur_factors(W, rows, cols, rel_tol=1e-8):
14 rows, cols = np.asarray(rows, dtype=int), np.asarray(cols, dtype=int)
15 A = W[:, cols]
16 B = W[rows, :]
17 M = pinv_truncated(W[np.ix_(rows, cols)], rel_tol)
18 return A, M, B
19
20
21def reconstruct(factors):
22 A, M, B = factors
23 return A @ M @ B
24
25
26def _top_with_retention(scores, old, k, retain):
27 old = np.asarray(old, dtype=int)
28 nold = max(1, int(round(retain * k)))
29 retained = list(old[np.argsort(scores[old])[::-1][:nold]])
30 chosen = retained[:]
31 for idx in np.argsort(scores)[::-1]:
32 if idx not in chosen:
33 chosen.append(int(idx))
34 if len(chosen) == k:
35 break
36 return np.asarray(chosen, dtype=int)
37
38
39def residual_leverage_refresh(W, old_rows, old_cols, probes=8, retain=0.25,
40 rng=None, rel_tol=1e-8):
41 rng = np.random.default_rng() if rng is None else rng
42 A, M, B = cur_factors(W, old_rows, old_cols, rel_tol)
43 R = W - A @ M @ B
44 # Row scores from RP and column scores from P^T R.
45 Pcol = rng.standard_normal((W.shape[1], probes))
46 Qr, _ = np.linalg.qr(R @ Pcol, mode='reduced')
47 Prow = rng.standard_normal((W.shape[0], probes))
48 Qc, _ = np.linalg.qr(R.T @ Prow, mode='reduced')
49 row_scores = np.sum(Qr * Qr, axis=1)
50 col_scores = np.sum(Qc * Qc, axis=1)
51 rows = _top_with_retention(row_scores, old_rows, len(old_rows), retain)
52 cols = _top_with_retention(col_scores, old_cols, len(old_cols), retain)
53 return cur_factors(W, rows, cols, rel_tol), rows, cols, np.linalg.norm(R) / np.linalg.norm(W)
54
55
56def rel_error(W, factors):
57 return np.linalg.norm(W - reconstruct(factors)) / np.linalg.norm(W)
58
59
60def factor_param_count(W, factors):
61 A, M, B = factors
62 return A.size + M.size + B.size
63
64
65def run(seed=7):
66 rng = np.random.default_rng(seed)
67 m, n, r, k = 64, 80, 8, 12
68 # A changing low-rank-plus-noise matrix: adaptive residual scores should track new rows/cols.
69 U = rng.standard_normal((m, r)); U, _ = np.linalg.qr(U)
70 V = rng.standard_normal((n, r)); V, _ = np.linalg.qr(V)
71 singular = np.linspace(3.0, 0.5, r)
72 W0 = (U * singular) @ V.T
73 initial_rows = np.arange(k)
74 initial_cols = np.arange(k)
75 fixed = cur_factors(W0, initial_rows, initial_cols)
76 random = cur_factors(W0, rng.choice(m, k, replace=False), rng.choice(n, k, replace=False))
77 exact = rel_error(W0, cur_factors(W0, np.arange(m), np.arange(n)))
78 initial_err = rel_error(W0, fixed)
79 random_err = rel_error(W0, random)
80
81 # Perturb a previously unselected block, then compare stale, random refresh, adaptive refresh.
82 W1 = W0.copy()
83 new_rows = np.arange(40, 52); new_cols = np.arange(55, 67)
84 W1[np.ix_(new_rows, new_cols)] += 0.8 * rng.standard_normal((len(new_rows), len(new_cols)))
85 stale_err = rel_error(W1, fixed)
86 _, rr, cc, _ = residual_leverage_refresh(W1, initial_rows, initial_cols, probes=16, retain=.25, rng=np.random.default_rng(seed+1))
87 adaptive = cur_factors(W1, rr, cc)
88 adaptive_err = rel_error(W1, adaptive)
89 random2 = cur_factors(W1, rng.choice(m, k, replace=False), rng.choice(n, k, replace=False))
90 random2_err = rel_error(W1, random2)
91 # Formula/evaluation equivalence on a batch.
92 X = rng.standard_normal((11, m))
93 y1 = X @ reconstruct(adaptive)
94 A, M, B = adaptive
95 y2 = ((X @ A) @ M) @ B
96 formula_gap = np.max(np.abs(y1-y2))
97 return dict(seed=seed, shape=[m,n], rank=r, cross_rank=k,
98 exact_dense_error=exact, initial_error=initial_err,
99 random_initial_error=random_err, stale_after_change=stale_err,
100 adaptive_after_change=adaptive_err, random_after_change=random2_err,
101 adaptive_rows=rr.tolist(), adaptive_cols=cc.tolist(),
102 formula_max_abs_gap=float(formula_gap), dense_params=m*n,
103 cur_params=factor_param_count(W1, adaptive), compression=(m*n)/factor_param_count(W1, adaptive))
104
105
106if __name__ == '__main__':
107 print(json.dumps(run(), indent=2))