Skip to content

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,

\[ \mathbb E\,\|A - QQ^\top A\|_2 \le \Big(1+\sqrt{\tfrac{k}{p-1}}\Big)\sigma_{k+1} + \frac{e\sqrt{k+p}}{p}\Big(\sum_{j>k}\sigma_j^2\Big)^{1/2}. \]

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

\[ \hat Q = \operatorname{orth}(A^\top Q),\qquad Q = \operatorname{orth}(A\hat Q), \]

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),

\[ \mathbb E\,\|A - QQ^\top A\|_2 \le \Big(1+\sqrt{\tfrac{k}{p-1}}\Big)\sigma_{k+1} + \frac{e\sqrt{k+p}}{p}\Big(\sum_{j>k}\sigma_j^2\Big)^{1/2}, \]

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 (m, n); may be matrix-free.

required
rank int

Target rank \(k\).

required
oversample int

Extra columns \(p\); \(\ell = k + p\), capped at min(m, n). Ignored when sketch is given.

10
n_power_iter int

Number of power iterations \(q\).

2
sketch AbstractSketch | None

Optional test matrix as a sketch \(S\) with in_size == n (e.g. SparseSignSketch, SRHTSketch), applied as its transpose, \(\Omega = S^\top \in \mathbb R^{n\times\ell}\) with \(\ell\) = sketch.out_size. None draws a Gaussian \(\Omega\).

None
key Array | None

PRNG key for the Gaussian test matrix. None means jax.random.PRNGKey(0). Ignored when sketch is given.

None

Returns:

Type Description
Float[Array, 'm l']

\(Q\) with orthonormal columns, shape (m, l).

Raises:

Type Description
ValueError

On a non-positive rank, negative oversample or n_power_iter, a sketch whose in_size is not n, or a sketch with fewer than rank rows.

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
def 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"]:
    r"""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

    $$
    \hat Q = \operatorname{orth}(A^\top Q),\qquad Q = \operatorname{orth}(A\hat Q),
    $$

    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),

    $$
    \mathbb E\,\|A - QQ^\top A\|_2 \le
    \Big(1+\sqrt{\tfrac{k}{p-1}}\Big)\sigma_{k+1}
    + \frac{e\sqrt{k+p}}{p}\Big(\sum_{j>k}\sigma_j^2\Big)^{1/2},
    $$

    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.

    Args:
        op: Operator $A$ of shape ``(m, n)``; may be matrix-free.
        rank: Target rank $k$.
        oversample: Extra columns $p$; $\ell = k + p$, capped at
            ``min(m, n)``. Ignored when ``sketch`` is given.
        n_power_iter: Number of power iterations $q$.
        sketch: Optional test matrix as a sketch $S$ with ``in_size == n``
            (e.g. `SparseSignSketch`, `SRHTSketch`), applied as its transpose,
            $\Omega = S^\top \in \mathbb R^{n\times\ell}$ with $\ell$ =
            ``sketch.out_size``. ``None`` draws a Gaussian $\Omega$.
        key: PRNG key for the Gaussian test matrix. ``None`` means
            ``jax.random.PRNGKey(0)``. Ignored when ``sketch`` is given.

    Returns:
        $Q$ with orthonormal columns, shape ``(m, l)``.

    Raises:
        ValueError: On a non-positive ``rank``, negative ``oversample`` or
            ``n_power_iter``, a sketch whose ``in_size`` is not ``n``, or a
            sketch with fewer than ``rank`` rows.

    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
    """
    _check_args(rank, oversample, n_power_iter)
    m, n = op.out_size(), op.in_size()
    dtype = op.in_structure().dtype
    if sketch is None:
        if key is None:
            key = jax.random.PRNGKey(0)
        ell = min(rank + oversample, m, n)
        omega = jax.random.normal(key, (n, ell), dtype=dtype)
    else:
        if sketch.in_size != n:
            raise ValueError(
                f"sketch.in_size={sketch.in_size} must equal the operator's "
                f"in_size={n}."
            )
        if sketch.out_size < rank:
            raise ValueError(
                f"sketch.out_size={sketch.out_size} must be at least rank={rank}."
            )
        omega = sketch.apply_transpose(jnp.eye(sketch.out_size, dtype=dtype))

    Q = _orth(_matmat(op, omega))
    if n_power_iter > 0:
        op_t = op.transpose()
        for _ in range(n_power_iter):
            Q = _orth(_matmat(op, _orth(_matmat(op_t, Q))))
    return Q

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 (m, n); may be matrix-free.

required
rank int

Target rank \(k\).

required
oversample int

Extra columns \(p\); \(\ell = k + p\), capped at min(m, n).

10
n_power_iter int

Number of power iterations \(q\).

2
key Array | None

PRNG key for the Gaussian test matrix. None means jax.random.PRNGKey(0).

None

Returns:

Type Description
tuple[Float[Array, 'm l'], Float[Array, 'l n']]

(Q, B) of shapes (m, l) and (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
def 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"]]:
    r"""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.

    Args:
        op: Operator $A$ of shape ``(m, n)``; may be matrix-free.
        rank: Target rank $k$.
        oversample: Extra columns $p$; $\ell = k + p$, capped at
            ``min(m, n)``.
        n_power_iter: Number of power iterations $q$.
        key: PRNG key for the Gaussian test matrix. ``None`` means
            ``jax.random.PRNGKey(0)``.

    Returns:
        ``(Q, B)`` of shapes ``(m, l)`` and ``(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
    """
    Q = range_finder(
        op, rank, oversample=oversample, n_power_iter=n_power_iter, key=key
    )
    B = rearrange(_matmat(op.transpose(), Q), "n l -> l n")
    return Q, B

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 (m, n); may be matrix-free.

required
rank int

Number of singular triplets \(k\) to return (at most min(m, n)).

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 means jax.random.PRNGKey(0).

None

Returns:

Type Description
Float[Array, 'm k']

(U, s, Vt) of shapes (m, k), (k,) and (k, n), with

Float[Array, ' k']

s descending.

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
def 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"]]:
    r"""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.

    Args:
        op: Operator $A$ of shape ``(m, n)``; may be matrix-free.
        rank: Number of singular triplets $k$ to return (at most
            ``min(m, n)``).
        oversample: Extra columns $p$ in the range finder.
        n_power_iter: Number of power iterations $q$.
        key: PRNG key for the Gaussian test matrix. ``None`` means
            ``jax.random.PRNGKey(0)``.

    Returns:
        ``(U, s, Vt)`` of shapes ``(m, k)``, ``(k,)`` and ``(k, n)``, with
        ``s`` descending.

    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)
    """
    Q, B = qb(op, rank, oversample=oversample, n_power_iter=n_power_iter, key=key)
    U_b, s, Vt = jnp.linalg.svd(B, full_matrices=False)
    k = min(rank, s.shape[0])
    U = einsum(Q, U_b[:, :k], "m l, l k -> m k")
    return U, s[:k], Vt[:k]

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 (n, n); may be matrix-free.

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" or "magnitude".

'largest'
key Array | None

PRNG key for the Gaussian test matrix. None means jax.random.PRNGKey(0).

None

Returns:

Type Description
Float[Array, ' k']

(eigenvalues, eigenvectors) of shapes (k,) and (n, k),

Float[Array, 'n k']

eigenvalues in ascending order (as jax.numpy.linalg.eigh),

tuple[Float[Array, ' k'], Float[Array, 'n k']]

eigenvectors orthonormal.

Raises:

Type Description
ValueError

If which is invalid or the operator is not square.

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
def 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"]]:
    r"""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).

    Args:
        op: Symmetric operator $A$ of shape ``(n, n)``; may be matrix-free.
        rank: Number of eigenpairs $k$ to return.
        oversample: Extra columns $p$ in the range finder.
        n_power_iter: Number of power iterations $q$.
        which: ``"largest"`` or ``"magnitude"``.
        key: PRNG key for the Gaussian test matrix. ``None`` means
            ``jax.random.PRNGKey(0)``.

    Returns:
        ``(eigenvalues, eigenvectors)`` of shapes ``(k,)`` and ``(n, k)``,
        eigenvalues in ascending order (as `jax.numpy.linalg.eigh`),
        eigenvectors orthonormal.

    Raises:
        ValueError: If ``which`` is invalid or the operator is not square.

    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
    """
    if which not in ("largest", "magnitude"):
        raise ValueError(f"which must be 'largest' or 'magnitude', got {which!r}.")
    if op.in_size() != op.out_size():
        raise ValueError(
            f"randomized_eigh needs a square operator, got "
            f"{op.out_size()}x{op.in_size()}."
        )
    Q = range_finder(
        op, rank, oversample=oversample, n_power_iter=n_power_iter, key=key
    )
    T = symmetrize(einsum(Q, _matmat(op, Q), "n a, n b -> a b"))
    vals, W = jnp.linalg.eigh(T)
    k = min(rank, vals.shape[0])
    if which == "largest":
        keep = jnp.arange(vals.shape[0] - k, vals.shape[0])
    else:
        keep = jnp.sort(jnp.argsort(-jnp.abs(vals))[:k])
    vecs = einsum(Q, W[:, keep], "n l, l k -> n k")
    return vals[keep], vecs

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

\[ \hat A = (A\Omega)\,(\Omega^\top A\Omega)^{+}\,(A\Omega)^\top, \qquad 0 \preceq \hat A \preceq A. \]

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:

  1. \(\Omega = \operatorname{qr}(\text{randn}(n, \ell))\);
  2. \(Y = A\Omega\) and the shift \(\nu = \sqrt n\,\varepsilon\,\|Y\|_2\);
  3. \(Y_\nu = Y + \nu\Omega\), \(C = \operatorname{chol}(\Omega^\top Y_\nu)\), \(B = Y_\nu C^{-\top}\);
  4. \(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 (n, n); may be matrix-free (touched only through mv, vmapped over the \(\ell\) columns).

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 means jax.random.PRNGKey(0).

None

Returns:

Type Description
LowRankUpdate

LowRankUpdate(base=0, U=U, d=Λ̂, V=U, orthonormal=True), tagged

LowRankUpdate

symmetric and PSD, with U of shape (n, k) and Λ̂

LowRankUpdate

descending.

Raises:

Type Description
ValueError

On a non-positive rank, a negative oversample, or a non-square operator.

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
def randomized_nystrom(
    op: lx.AbstractLinearOperator,
    rank: int,
    *,
    oversample: int = 0,
    key: jax.Array | None = None,
) -> LowRankUpdate:
    r"""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

    $$
    \hat A = (A\Omega)\,(\Omega^\top A\Omega)^{+}\,(A\Omega)^\top,
    \qquad 0 \preceq \hat A \preceq A.
    $$

    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:

    1. $\Omega = \operatorname{qr}(\text{randn}(n, \ell))$;
    2. $Y = A\Omega$ and the shift $\nu = \sqrt n\,\varepsilon\,\|Y\|_2$;
    3. $Y_\nu = Y + \nu\Omega$, $C = \operatorname{chol}(\Omega^\top Y_\nu)$,
       $B = Y_\nu C^{-\top}$;
    4. $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`).

    Args:
        op: PSD operator $A$ of shape ``(n, n)``; may be matrix-free (touched
            only through ``mv``, vmapped over the $\ell$ columns).
        rank: Number of eigenpairs $k$ to return.
        oversample: Extra columns $p$; $\ell = k + p$, capped at $n$.
        key: PRNG key for the Gaussian test matrix. ``None`` means
            ``jax.random.PRNGKey(0)``.

    Returns:
        ``LowRankUpdate(base=0, U=U, d=Λ̂, V=U, orthonormal=True)``, tagged
        symmetric and PSD, with ``U`` of shape ``(n, k)`` and ``Λ̂``
        descending.

    Raises:
        ValueError: On a non-positive ``rank``, a negative ``oversample``, or
            a non-square operator.

    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))
    """
    if rank < 1:
        raise ValueError(f"rank must be a positive integer, got {rank}.")
    if oversample < 0:
        raise ValueError(f"oversample must be non-negative, got {oversample}.")
    n = op.in_size()
    if op.out_size() != n:
        raise ValueError(
            f"randomized_nystrom needs a square operator, got {op.out_size()}x{n}."
        )
    if key is None:
        key = jax.random.PRNGKey(0)
    dtype = op.in_structure().dtype
    ell = min(rank + oversample, n)

    omega, _ = jnp.linalg.qr(jax.random.normal(key, (n, ell), dtype=dtype))
    Y = _matmat(op, omega)
    # The shift only stabilises the small Cholesky; it is removed again below.
    nu = jnp.sqrt(jnp.asarray(n, dtype)) * jnp.finfo(dtype).eps * jnp.linalg.norm(Y, 2)
    Y_nu = Y + nu * omega
    C = jnp.linalg.cholesky(symmetrize(einsum(omega, Y_nu, "n a, n b -> a b")))
    # B = Y_nu C⁻ᵀ, so B Bᵀ = Y_nu (Ωᵀ Y_nu)⁻¹ Y_nuᵀ.
    Bt = jax.scipy.linalg.solve_triangular(C, rearrange(Y_nu, "n l -> l n"), lower=True)
    U, s, _ = jnp.linalg.svd(rearrange(Bt, "l n -> n l"), full_matrices=False)
    k = min(rank, ell)
    U, eigenvalues = U[:, :k], jnp.maximum(s[:k] ** 2 - nu, 0.0)
    # V is U (the same object), so the update is symmetric and PSD-tagged.
    zeros = jnp.zeros(n, dtype=dtype)
    return svd_low_rank_plus_diag(zeros, U, eigenvalues, U, psd=True)

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

\[ g = A_{:,s} - F F_{s,:}^\top,\qquad F_{:,i} = g / \sqrt{g_s}. \]

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 A, shape (N,).

required
column Callable[[Int[Array, '']], Float[Array, ' N']]

Callable returning column s of A, shape (N,), for a scalar integer index s (a traced value inside the loop).

required
rank int

Number of pivots k.

required
pivoting Literal['random', 'greedy']

"random" samples s ∝ max(d, 0); "greedy" takes argmax(d), the classic pivoted Cholesky.

'random'
block_size int

Pivots per step. Only 1 is implemented; the blocked variant (Epperly, Tropp & Webber, 2024) is a follow-up.

1
key Array | None

PRNG key for "random" pivoting. None means jax.random.PRNGKey(0). Ignored by "greedy".

None

Returns:

Type Description
Float[Array, 'N k']

(F, pivots): the factor, shape (N, k), and the pivot indices,

Int[Array, ' k']

shape (k,), in the order chosen (-1 past the numerical rank).

Raises:

Type Description
ValueError

If pivoting is not "random" or "greedy".

NotImplementedError

If block_size != 1.

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
def 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"]]:
    r"""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

    $$
    g = A_{:,s} - F F_{s,:}^\top,\qquad F_{:,i} = g / \sqrt{g_s}.
    $$

    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]``.

    Args:
        diagonal: Diagonal of the PSD matrix ``A``, shape ``(N,)``.
        column: Callable returning column ``s`` of ``A``, shape ``(N,)``, for
            a scalar integer index ``s`` (a traced value inside the loop).
        rank: Number of pivots ``k``.
        pivoting: ``"random"`` samples ``s ∝ max(d, 0)``; ``"greedy"`` takes
            ``argmax(d)``, the classic pivoted Cholesky.
        block_size: Pivots per step. Only ``1`` is implemented; the blocked
            variant (Epperly, Tropp & Webber, 2024) is a follow-up.
        key: PRNG key for ``"random"`` pivoting. ``None`` means
            ``jax.random.PRNGKey(0)``. Ignored by ``"greedy"``.

    Returns:
        ``(F, pivots)``: the factor, shape ``(N, k)``, and the pivot indices,
        shape ``(k,)``, in the order chosen (``-1`` past the numerical rank).

    Raises:
        ValueError: If ``pivoting`` is not ``"random"`` or ``"greedy"``.
        NotImplementedError: If ``block_size != 1``.

    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
    """
    if pivoting not in ("random", "greedy"):
        raise ValueError(f"pivoting must be 'random' or 'greedy', got {pivoting!r}")
    if block_size != 1:
        raise NotImplementedError("rp_cholesky supports only block_size=1 for now")
    if key is None:
        key = jax.random.PRNGKey(0)

    # LAPACK ?pstrf stopping criterion: n * eps * max diagonal entry.
    tol = diagonal.shape[0] * jnp.finfo(diagonal.dtype).eps * jnp.max(jnp.abs(diagonal))

    def body(i, carry):
        F, pivots = carry
        residual = diagonal - reduce(F * F, "n k -> n", "sum")
        if pivoting == "greedy":
            s = jnp.argmax(residual)
        else:
            # Entries at or below the guard (including chosen pivots, whose
            # residual is rounding noise) get probability zero.
            usable = residual > tol
            log_weights = jnp.where(
                usable, jnp.log(jnp.where(usable, residual, 1.0)), -jnp.inf
            )
            s = jax.random.categorical(jax.random.fold_in(key, i), log_weights)
        pivot = residual[s]
        ok = pivot > tol
        # Double-where keeps the sqrt's gradient finite when guarded.
        denom = jnp.sqrt(jnp.where(ok, pivot, 1.0))
        col = (column(s) - F @ F[s, :]) / denom
        F = F.at[:, i].set(jnp.where(ok, col, 0.0))
        pivots = pivots.at[i].set(jnp.where(ok, s, -1).astype(pivots.dtype))
        return F, pivots

    F0 = jnp.zeros((diagonal.shape[0], rank), dtype=diagonal.dtype)
    pivots0 = jnp.full((rank,), -1, dtype=jnp.int32)
    return jax.lax.fori_loop(0, rank, body, (F0, pivots0))