Randomized Linear Algebra¶
Randomized factorisations that touch a matrix only through a few columns or
matvecs. They take an explicit PRNG key; key=None means
jax.random.PRNGKey(0).
Range finder, QB, SVD and eigh¶
Randomized subspace iteration (Halko, Martinsson & Tropp, 2011) finds an orthonormal basis \(Q\) for the dominant range of a matrix-free operator from a block of \(\ell = k + p\) matvecs. For a Gaussian test matrix,
The tail term hurts for slowly decaying spectra (Matérn-½ Gram matrices, most
geophysical fields); n_power_iter=q applies the bound to \((AA^\top)^qA\),
whose singular values are \(\sigma_j^{2q+1}\), for \(2q\) more passes. Use
n_power_iter >= 2 there. These methods target the top of the spectrum;
the small end (e.g. the smallest Laplacian eigenvalues) is Lanczos / LOBPCG
territory.
qb returns \(Q\) and \(B = Q^\top A\), randomized_svd lifts the SVD of \(B\),
and randomized_eigh is the Rayleigh–Ritz projection \(Q^\top A Q\) for
symmetric, possibly indefinite, operators. svd(op, rank=k,
method="randomized") and eig(op, rank=k, method="randomized") route here;
Lanczos stays the default.
# 50 EOFs of a (100k pixels × 3650 days) anomaly matrix, available only as a matvec
U, s, Vt = gx.randomized_svd(anomalies_op, 50, oversample=10, n_power_iter=2, key=key)
eofs, pcs = U, einx.multiply("k, k t -> k t", s, Vt)
Structured linear algebra and Gaussian primitives for JAX.
range_finder(op: lx.AbstractLinearOperator, rank: int, *, oversample: int = 10, n_power_iter: int = 2, sketch: AbstractSketch | None = None, key: jax.Array | None = None) -> Float[Array, 'm l']
¶
Orthonormal basis \(Q\) for the dominant range of \(A\).
Randomized subspace iteration (Halko, Martinsson & Tropp, 2011,
Algorithm 4.4): draw a test matrix \(\Omega \in \mathbb R^{n\times\ell}\)
with \(\ell\) = rank + oversample, set \(Q = \operatorname{orth}(A\Omega)\),
then repeat n_power_iter times
re-orthonormalising after every half-step so that small directions are not lost to round-off. For a Gaussian \(\Omega\) (HMT 2011, Thm 10.6),
with \(k\) = rank and \(p\) = oversample. The tail term dominates
when the spectrum decays slowly (Matérn-½ Gram matrices, most
geophysical fields); \(q\) power iterations apply the bound to
\((AA^\top)^q A\), whose singular values are \(\sigma_j^{2q+1}\), at the cost
of \(2q\) more passes over \(A\). Use n_power_iter >= 2 for slowly
decaying spectra.
Randomized methods target the top of the spectrum (the largest singular values). For the small end, use Lanczos or LOBPCG.
\(A\) is touched only through mv (and the transpose's mv for power
steps), vmapped over the \(\ell\) columns.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
op
|
AbstractLinearOperator
|
Operator \(A\) of shape |
required |
rank
|
int
|
Target rank \(k\). |
required |
oversample
|
int
|
Extra columns \(p\); \(\ell = k + p\), capped at
|
10
|
n_power_iter
|
int
|
Number of power iterations \(q\). |
2
|
sketch
|
AbstractSketch | None
|
Optional test matrix as a sketch \(S\) with |
None
|
key
|
Array | None
|
PRNG key for the Gaussian test matrix. |
None
|
Returns:
| Type | Description |
|---|---|
Float[Array, 'm l']
|
\(Q\) with orthonormal columns, shape |
Raises:
| Type | Description |
|---|---|
ValueError
|
On a non-positive |
Examples:
>>> import einx, jax.numpy as jnp, jax.random as jr, lineax as lx
>>> import gaussx as gx
>>> A = jr.normal(jr.key(0), (100, 5)) @ jr.normal(jr.key(1), (5, 80))
>>> Q = gx.range_finder(lx.MatrixLinearOperator(A), 5, key=jr.key(2))
>>> Q.shape
(100, 15)
>>> QtA = einx.dot("m l, m n -> l n", Q, A)
>>> bool(jnp.allclose(Q @ QtA, A, atol=1e-4))
True
Source code in src/gaussx/_randomized/_range_finder.py
36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 | |
qb(op: lx.AbstractLinearOperator, rank: int, *, oversample: int = 10, n_power_iter: int = 2, key: jax.Array | None = None) -> tuple[Float[Array, 'm l'], Float[Array, 'l n']]
¶
Randomized QB factorisation \(A \approx QB\) with \(B = Q^\top A\).
\(Q\) comes from range_finder; \(B\) costs \(\ell\) more transpose-matvecs,
\(B = (A^\top Q)^\top\), so \(A\) is never formed. \(\|A - QB\|\) is the
range-finder error (see range_finder for its bound and the advice on
n_power_iter >= 2 for slowly decaying spectra). Randomized methods
target the top of the spectrum.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
op
|
AbstractLinearOperator
|
Operator \(A\) of shape |
required |
rank
|
int
|
Target rank \(k\). |
required |
oversample
|
int
|
Extra columns \(p\); \(\ell = k + p\), capped at
|
10
|
n_power_iter
|
int
|
Number of power iterations \(q\). |
2
|
key
|
Array | None
|
PRNG key for the Gaussian test matrix. |
None
|
Returns:
| Type | Description |
|---|---|
tuple[Float[Array, 'm l'], Float[Array, 'l n']]
|
|
Examples:
>>> import jax.numpy as jnp, jax.random as jr, lineax as lx
>>> import gaussx as gx
>>> A = jr.normal(jr.key(0), (60, 4)) @ jr.normal(jr.key(1), (4, 40))
>>> Q, B = gx.qb(lx.MatrixLinearOperator(A), 4, oversample=4)
>>> Q.shape, B.shape
((60, 8), (8, 40))
>>> bool(jnp.allclose(Q @ B, A, atol=1e-4))
True
Source code in src/gaussx/_randomized/_range_finder.py
randomized_svd(op: lx.AbstractLinearOperator, rank: int, *, oversample: int = 10, n_power_iter: int = 2, key: jax.Array | None = None) -> tuple[Float[Array, 'm k'], Float[Array, ' k'], Float[Array, 'k n']]
¶
Truncated SVD \(A \approx U \operatorname{diag}(s) V^\top\) by randomized QB.
Computes \(A \approx QB\) with qb, the small SVD \(B = U_B \Sigma V^\top\),
and lifts \(U = Q U_B\), keeping the top rank triplets (Halko,
Martinsson & Tropp, 2011, Algorithm 5.1). The cost is
\((2q + 2)\ell\) matvecs with \(A\) or \(A^\top\), \(\ell\) = rank +
oversample, plus \(O((m+n)\ell^2)\) flops; \(A\) is never formed.
Randomized methods target the top of the spectrum: the leading
singular triplets are accurate, the trailing ones are not. Use
n_power_iter >= 2 for slowly decaying spectra (Matérn-½ Gram
matrices, most geophysical fields); see range_finder for the error
bound.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
op
|
AbstractLinearOperator
|
Operator \(A\) of shape |
required |
rank
|
int
|
Number of singular triplets \(k\) to return (at most
|
required |
oversample
|
int
|
Extra columns \(p\) in the range finder. |
10
|
n_power_iter
|
int
|
Number of power iterations \(q\). |
2
|
key
|
Array | None
|
PRNG key for the Gaussian test matrix. |
None
|
Returns:
| Type | Description |
|---|---|
Float[Array, 'm k']
|
|
Float[Array, ' k']
|
|
Examples:
50 EOFs of an anomaly matrix available only as a matvec, here a small dense stand-in:
>>> import einx, jax.random as jr, lineax as lx
>>> import gaussx as gx
>>> X = jr.normal(jr.key(0), (300, 8)) @ jr.normal(jr.key(1), (8, 120))
>>> U, s, Vt = gx.randomized_svd(lx.MatrixLinearOperator(X), 5, key=jr.key(2))
>>> U.shape, s.shape, Vt.shape
((300, 5), (5,), (5, 120))
>>> eofs, pcs = U, einx.multiply("k, k t -> k t", s, Vt)
Source code in src/gaussx/_randomized/_svd.py
randomized_eigh(op: lx.AbstractLinearOperator, rank: int, *, oversample: int = 10, n_power_iter: int = 2, which: Literal['largest', 'magnitude'] = 'largest', key: jax.Array | None = None) -> tuple[Float[Array, ' k'], Float[Array, 'n k']]
¶
Partial eigendecomposition of a symmetric operator by Rayleigh-Ritz.
Finds an orthonormal \(Q\) for the dominant range of \(A\) with
range_finder, forms the Rayleigh-Ritz matrix \(T = Q^\top A Q\)
(\(\ell\) more matvecs), and lifts the eigenpairs of \(T\):
\(A \approx (QW)\Lambda(QW)^\top\) with \(T = W\Lambda W^\top\). \(A\) may be
indefinite. The range finder captures the eigenvalues of largest
magnitude (randomized methods target the top of the spectrum),
and rank Ritz pairs are kept by which:
"largest": the algebraically largest Ritz values;"magnitude": the Ritz values of largest absolute value.
For an indefinite \(A\) whose large negative eigenvalues dominate,
"largest" is only as good as the subspace, so prefer
"magnitude" there. The small end of a spectrum (e.g. the smallest
eigenvalues of a graph Laplacian) is Lanczos / LOBPCG territory.
Use n_power_iter >= 2 for slowly decaying spectra (Matérn-½ Gram
matrices, most geophysical fields). With n_power_iter=0 this is the
one-pass randomized Rayleigh-Ritz projection, the successor of the
algorithm behind NystromPreconditioner up to gaussx 0.4; it projects
onto \(\operatorname{orth}(A\Omega)\) rather than
\(\operatorname{orth}(\Omega)\), which is more accurate, so it is not
numerically identical. For PSD operators,
randomized_nystrom is strictly more accurate for the same number of
matvecs (Tropp et al., 2017).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
op
|
AbstractLinearOperator
|
Symmetric operator \(A\) of shape |
required |
rank
|
int
|
Number of eigenpairs \(k\) to return. |
required |
oversample
|
int
|
Extra columns \(p\) in the range finder. |
10
|
n_power_iter
|
int
|
Number of power iterations \(q\). |
2
|
which
|
Literal['largest', 'magnitude']
|
|
'largest'
|
key
|
Array | None
|
PRNG key for the Gaussian test matrix. |
None
|
Returns:
| Type | Description |
|---|---|
Float[Array, ' k']
|
|
Float[Array, 'n k']
|
eigenvalues in ascending order (as |
tuple[Float[Array, ' k'], Float[Array, 'n k']]
|
eigenvectors orthonormal. |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Examples:
>>> import einx, jax.numpy as jnp, jax.random as jr, lineax as lx
>>> import gaussx as gx
>>> lam = jnp.array([-9.0, 5.0, 1.0, 0.1, 0.01, 0.0])
>>> Q, _ = jnp.linalg.qr(jr.normal(jr.key(0), (6, 6)))
>>> A = einx.dot("i k, j k -> i j", einx.multiply("i k, k -> i k", Q, lam), Q)
>>> A = lx.MatrixLinearOperator(A, lx.symmetric_tag)
>>> vals, vecs = gx.randomized_eigh(A, 2, oversample=2, which="magnitude")
>>> bool(jnp.allclose(vals, jnp.array([-9.0, 5.0]), atol=1e-4))
True
>>> vals, _ = gx.randomized_eigh(A, 2, oversample=2, which="largest")
>>> bool(jnp.allclose(vals, jnp.array([1.0, 5.0]), atol=1e-4))
True
Source code in src/gaussx/_randomized/_svd.py
71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 | |
Randomized Nyström¶
For a PSD operator, randomized_nystrom returns the Nyström approximation
\(\hat A = (A\Omega)(\Omega^\top A\Omega)^{+}(A\Omega)^\top\) from one pass of
\(\ell\) matvecs (Tropp, Yurtsever, Udell & Cevher, 2017, Algorithm 3). It
satisfies \(0 \preceq \hat A \preceq A\) and, for the same \(\ell\), is more
accurate than the Rayleigh–Ritz projection of randomized_eigh. The result
is an orthonormal LowRankUpdate
\(U\hat\Lambda U^\top\), so the same factors on a \(\sigma^2 I\) base solve and
take log-determinants of \(\hat A + \sigma^2 I\) through the Woodbury rules. It
is also the sketch behind
NystromPreconditioner.
Structured linear algebra and Gaussian primitives for JAX.
randomized_nystrom(op: lx.AbstractLinearOperator, rank: int, *, oversample: int = 0, key: jax.Array | None = None) -> LowRankUpdate
¶
Randomized Nyström approximation \(\hat A = U\hat\Lambda U^\top\) of a PSD \(A\).
For a test matrix \(\Omega\) the Nyström approximation is
It costs one pass (\(\ell\) = rank + oversample matvecs) and, for the
same \(\ell\), is more accurate than the Rayleigh-Ritz approximation
\(QQ^\top AQQ^\top\) of randomized_eigh (Tropp, Yurtsever, Udell &
Cevher, 2017). The algorithm is their Algorithm 3:
- \(\Omega = \operatorname{qr}(\text{randn}(n, \ell))\);
- \(Y = A\Omega\) and the shift \(\nu = \sqrt n\,\varepsilon\,\|Y\|_2\);
- \(Y_\nu = Y + \nu\Omega\), \(C = \operatorname{chol}(\Omega^\top Y_\nu)\), \(B = Y_\nu C^{-\top}\);
- \(U, \Sigma, \_ = \operatorname{svd}(B)\), \(\hat\Lambda = \max(\Sigma^2 - \nu, 0)\).
The shift \(\nu\) only stabilises the small Cholesky (it keeps float32
finite) and is subtracted again in step 4. With oversample > 0 the
top rank eigenpairs of the rank-\(\ell\) approximation are kept.
The result is an orthonormal LowRankUpdate with a zero diagonal base,
so the same factors on a \(\sigma^2 I\) base
(gaussx.svd_low_rank_plus_diag) give \(\hat A + \sigma^2 I\), whose
gaussx.solve, gaussx.logdet and the rest dispatch through the
Woodbury rules. Randomized methods target the top of the spectrum;
\(A\) must be PSD (for symmetric indefinite \(A\) use randomized_eigh).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
op
|
AbstractLinearOperator
|
PSD operator \(A\) of shape |
required |
rank
|
int
|
Number of eigenpairs \(k\) to return. |
required |
oversample
|
int
|
Extra columns \(p\); \(\ell = k + p\), capped at \(n\). |
0
|
key
|
Array | None
|
PRNG key for the Gaussian test matrix. |
None
|
Returns:
| Type | Description |
|---|---|
LowRankUpdate
|
|
LowRankUpdate
|
symmetric and PSD, with |
LowRankUpdate
|
descending. |
Raises:
| Type | Description |
|---|---|
ValueError
|
On a non-positive |
Examples:
>>> import einx, jax.numpy as jnp, jax.random as jr, lineax as lx
>>> import gaussx as gx
>>> W = jr.normal(jr.key(0), (50, 4))
>>> A = lx.MatrixLinearOperator(
... einx.dot("i r, j r -> i j", W, W), lx.positive_semidefinite_tag
... )
>>> A_hat = gx.randomized_nystrom(A, 4, oversample=2, key=jr.key(1))
>>> A_hat.U.shape, A_hat.d.shape
((50, 4), (4,))
>>> bool(jnp.allclose(A_hat.as_matrix(), A.as_matrix(), atol=1e-3))
True
Add the noise and solve with the Woodbury identity:
>>> noisy = gx.svd_low_rank_plus_diag(
... jnp.full(50, 0.1), A_hat.U, A_hat.d, A_hat.U, psd=True
... )
>>> x = gx.solve(noisy, jnp.ones(50))
Source code in src/gaussx/_randomized/_nystrom.py
16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 | |
Randomly pivoted Cholesky¶
rp_cholesky builds a partial Cholesky factor from the diagonal and a
column(j) callable, picking each pivot with probability proportional to the
residual diagonal (Chen, Epperly, Tropp & Webber, 2023). It returns the
pivots too, so they serve as landmark indices for Nyström, Falkon or SVGP
inducing points without ever forming the kernel matrix. pivoting="greedy"
is the classic pivoted Cholesky behind
PartialCholeskyPreconditioner.
Structured linear algebra and Gaussian primitives for JAX.
rp_cholesky(diagonal: Float[Array, ' N'], column: Callable[[Int[Array, '']], Float[Array, ' N']], rank: int, *, pivoting: Literal['random', 'greedy'] = 'random', block_size: int = 1, key: jax.Array | None = None) -> tuple[Float[Array, 'N k'], Int[Array, ' k']]
¶
Randomly pivoted partial Cholesky of a PSD matrix A.
Builds F with F Fᵀ ≈ A one pivot at a time, touching A only
through its diagonal and rank of its columns (Chen, Epperly, Tropp &
Webber, 2023). At step i the pivot s is drawn with probability
proportional to the residual diagonal \(d_s = [A - FF^\top]_{ss}\), the
variance not yet explained, and then
With \(k \ge r/\varepsilon + r\log(1/(\varepsilon\eta))\) pivots,
\(\mathbb E\,\operatorname{tr}(A - FF^\top) \le
(1+\varepsilon)\operatorname{tr}(A - [\![A]\!]_r)\), where
\(\eta = \operatorname{tr}(A - [\![A]\!]_r)/\operatorname{tr}A\). Greedy
pivoting (argmax of the residual diagonal) has no such guarantee and
chases outliers. The cost is rank column evaluations and
\(O(N k^2)\) flops.
The returned pivots S make F Fᵀ = A[:, S] A[S, S]⁺ A[S, :], the
column Nyström approximation on those columns, so they double as
landmark (inducing-point) indices. The diagonal and column can come
from a kernel evaluated on the fly, so A is never formed: for 10⁶
points and rank=1000 this is 1000 kernel columns.
The residual is guarded like LAPACK ?pstrf: once the chosen pivot
falls below N · eps · max|diag A| the numerical rank is exhausted,
and that and every later column of F is exactly zero, with pivot
-1 (gh-236, gh-237). Filter with pivots[pivots >= 0].
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
diagonal
|
Float[Array, ' N']
|
Diagonal of the PSD matrix |
required |
column
|
Callable[[Int[Array, '']], Float[Array, ' N']]
|
Callable returning column |
required |
rank
|
int
|
Number of pivots |
required |
pivoting
|
Literal['random', 'greedy']
|
|
'random'
|
block_size
|
int
|
Pivots per step. Only |
1
|
key
|
Array | None
|
PRNG key for |
None
|
Returns:
| Type | Description |
|---|---|
Float[Array, 'N k']
|
|
Int[Array, ' k']
|
shape |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
NotImplementedError
|
If |
Examples:
Pick 20 landmarks from 1000 points without forming the kernel matrix.
>>> import jax.numpy as jnp, jax.random as jr, gaussx
>>> X = jr.normal(jr.key(0), (1000,))
>>> def column(j):
... return jnp.exp(-0.5 * (X - X[j]) ** 2)
>>> F, pivots = gaussx.rp_cholesky(jnp.ones(1000), column, 20, key=jr.key(1))
>>> F.shape, pivots.shape
((1000, 20), (20,))
>>> Z = X[pivots] # landmarks for Nyström / Falkon / SVGP
Source code in src/gaussx/_randomized/_rpcholesky.py
15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 | |