Adaptive CUR Neural Layer / adaptive_cur.py

Mechanism failed

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