Skip to content

Sparse Operators

SparseOperator is a sparse matrix whose sparsity pattern is a static, host-side, hashable SparsityPattern; only the non-zero values are traced. Every matrix a GMRF / INLA workflow factorises — a prior precision, a graph Laplacian, the Laplace Hessian \(Q + A^\top W A\) — has a pattern known before any value is, so all symbolic work (canonical ordering, transposition, pattern unions, the pattern of a congruence) runs once per pattern on the host and is reused across jit calls, Newton steps, hyperparameter values and vmap-ped datasets. Changing the values never retraces; changing the pattern does.

The pattern of the Laplace Hessian is known in advance:

\[ (A^\top W A)_{ij} = \sum_k A_{ki}\,w_k\,A_{kj} \neq 0 \;\Rightarrow\; i, j \in \operatorname{supp}(A_{k,:}),\qquad \operatorname{pattern}(Q + A^\top W A) = \operatorname{pattern}(Q)\cup \textstyle\bigcup_k \operatorname{supp}(A_{k,:})^2 . \]

A FEM projector row touches the vertices of one triangle, which are already neighbours in \(Q\), so the pattern does not grow; each fixed effect adds one dense row and column.

Storage. Patterns are canonical: sorted row-major, duplicates merged, and the diagonal of a square matrix always present. symmetric=True stores the lower triangle only; from_coo(..., symmetric=True) takes each off-diagonal pair once (edge-once graph storage) and mirrors it.

Dispatch.

Primitive Behaviour
diag Exact, read from the pattern
solve SparseCholeskySolver when the caller passes it (strategies, dispatch_solve); an explicit lineax solver wins; otherwise CG when PSD-tagged and larger than AutoSolver.size_threshold, dense below
logdet SparseCholeskySolver when passed; otherwise SLQLogdet when PSD-tagged and large, dense below
diag_inv Takahashi through the sparse factor with solver=SparseCholeskySolver(...) (any size) or method="cholesky"; "auto" uses it for N ≤ 2048 and Hutchinson above
eig(rank=) Lanczos on the matvec
cholesky A SparseCholeskyFactor (RCM ordering, cached symbolic analysis)

There is no size heuristic for the exact path: the caller knows N and chooses SparseCholeskySolver; AutoSolver keeps CG for large PSD operators.

JacobiPreconditioner works through the exact diagonal, which is usually enough for a graph Laplacian plus a diagonal shift.

Example. An ICAR structure matrix from an edge list, then the Laplace Hessian rebuilt on its precomputed pattern at each Newton step:

import equinox as eqx
import jax.numpy as jnp
import lineax as lx
import numpy as np

import gaussx as gx

# Path graph 0 - 1 - 2 - 3, each edge once, host-side indices
N = 4
senders, receivers = np.array([1, 2, 3]), np.array([0, 1, 2])
w = np.ones(3)
deg = np.bincount(senders, w, N) + np.bincount(receivers, w, N)
R = gx.SparseOperator.from_coo(
    np.r_[np.arange(N), senders],
    np.r_[np.arange(N), receivers],
    jnp.asarray(np.r_[deg, -w]),
    (N, N),
    symmetric=True,
    tags=frozenset({lx.positive_semidefinite_tag}),
)
tau = 2.0
Q = eqx.tree_at(lambda op: op.values, R, tau * R.values)  # same static pattern

# Observation projector A (2 observations) and Newton weights w_t
A = gx.SparseOperator.from_coo(
    np.array([0, 0, 1]), np.array([0, 1, 3]), jnp.array([0.5, 0.5, 1.0]), (2, N)
)
w_t = jnp.array([1.0, 2.0])
H = Q.union(Q.congruence(A, w_t), tags=lx.positive_semidefinite_tag)  # Q + Aᵀ diag(w_t) A
x = gx.solve(H, jnp.ones(N))

Sparse Cholesky

Cholesky is Gaussian elimination: eliminating node \(j\) connects its not-yet-eliminated neighbours, so the column patterns of \(P Q P^\top = L L^\top\) follow the elimination tree,

\[ \operatorname{struct}(L_{:,j}) = \operatorname{struct}(Q_{j:,j})\ \cup \bigcup_{\operatorname{parent}(c)=j}\operatorname{struct}(L_{:,c})\setminus\{c\}, \qquad \operatorname{parent}(j) = \min\{i>j : L_{ij}\neq 0\}. \]

That depends only on the pattern and the ordering, so symbolic_cholesky runs once on the host (NumPy / SciPy) and is cached per SparsityPattern. sparse_cholesky then traces only the values: it jits, vmaps over values and is differentiable.

  • Ordering. "rcm" (reverse Cuthill–McKee, the default) minimises the bandwidth; "natural" keeps the given order; "amd" (approximate minimum degree, through CHOLMOD) minimises fill.
  • Numeric phase. After RCM on a mesh the factor fills a narrow band, so \(L\) is block tridiagonal in bandwidth-sized blocks and dense block kernels do the work. Otherwise a left-looking lax.scan over columns gathers fixed-size windows from the CSC arrays (columns bucketed by length, so a few long separator columns do not pad all the short ones). Either way SparseCholeskyFactor.values holds \(L\) on its exact CSC pattern.
  • Takahashi. The backward recursion \(Z_{ij} = \delta_{ij}/L_{jj}^2 - L_{jj}^{-1}\sum_{k>j,\,k\in\operatorname{struct}(L_{:,j})} L_{kj} Z_{ki}\) evaluates \(Z = Q^{-1}\) exactly on \(\operatorname{pattern}(L + L^\top)\), which contains \(\operatorname{pattern}(Q)\), at about the cost of the factorisation. selected_inverse() returns it as a SparseOperator in the original order; diag_inv() gives the marginal variances.
  • Gradients. \(d\log|Q| = \operatorname{tr}(Q^{-1}dQ)\), so the log-determinant's cotangent is \(Z\) on \(\operatorname{pattern}(Q)\): one Takahashi sweep, never \(Q^{-1}\). For \(x = Q^{-1}b\): \(\bar b = Q^{-1}\bar x\) and \(\bar Q = -\bar b\,x^\top\), symmetrised. With symmetric=True storage an off-diagonal stored value sets \(Q_{ij}\) and \(Q_{ji}\), so its gradient is doubled (\(2Z_{ij}\)); general storage is factored as \(\tfrac12(Q + Q^\top)\) and each stored value gets \(Z_{ij}\). solve_lower_transpose (sampling, \(x = P^\top L^{-\top} z\)) and diag_inv are differentiated by JAX through the factorisation.
  • CHOLMOD backend (opt-in, pip install scikit-sparse, which needs SuiteSparse; not a dependency of gaussx). backend="cholmod" runs only the numeric factorisation in CHOLMOD through jax.pure_callback, on the same symbolic pattern; the solves, Takahashi and the gradients are the same JAX code, so both backends give identical gradients. CPU only, and vmap calls CHOLMOD once per batch element.

Scale. Measured on triangulated square meshes (the P1 FEM 7-point stencil), CPU, float64, jit-compiled, after compilation:

Nodes Ordering / backend nnz(L) Fill vs tril(Q) logdet value_and_grad(logdet) diag_inv
10,000 RCM / JAX (banded) 681,550 17.2 1.0 s 1.1 s 0.9 s
40,000 RCM / JAX (banded) 5,393,100 33.9 1.1 s 1.6 s 1.5 s
99,856 RCM / JAX (banded) 21,185,746 53.2 3.3 s 5.7 s 11 s
10,000 AMD / JAX (windows) 295,884 7.5 1.7 s 3.4 s 3.2 s
10,000 AMD / CHOLMOD 295,884 7.5 0.07 s 1.7 s 1.6 s
40,000 AMD / CHOLMOD 1,569,916 9.9 0.5 s 15 s 17 s
99,856 AMD / CHOLMOD 4,748,667 11.9 1.8 s ~2 min ~2 min

Timings are from a shared 16-core machine and are indicative only. The symbolic analysis is a one-off host cost (about 0.3 s at 10⁴ nodes and 8 s at 10⁵ with RCM).

  • 2-D meshes up to about \(10^5\) nodes: RCM with the JAX backend.
  • Beyond that, or when the fill of RCM is prohibitive: AMD through CHOLMOD (much less fill; its logdet is fast, but gradients and marginal variances run the JAX Takahashi on the AMD pattern, which gathers windows and is the slower path).
  • Beyond that: the iterative backend (CG, SLQ, Hutchinson).

Supernodes and level scheduling (GPU parallelism over independent subtrees) are follow-ups.

Example. A Matérn-like SPDE precision \(Q(\kappa) = \kappa^2 I + G\) on a 40k-node mesh: analyse once, factor for many \(\kappa\), differentiate the log-determinant.

import jax
import jax.numpy as jnp
import numpy as np

import gaussx as gx

# Stiffness-like graph Laplacian G of a triangulated 200 x 200 square
side = 200
n = side * side
i = np.arange(n)
right = i[i % side != side - 1]  # nodes with a right neighbour
down = i[i < n - side]  # nodes with a neighbour below
diagonal = right[right < n - side]  # one diagonal per cell
senders = np.r_[right + 1, down + side, diagonal + side + 1]
receivers = np.r_[right, down, diagonal]
w = np.ones(senders.size)
deg = np.bincount(senders, w, n) + np.bincount(receivers, w, n)
G = gx.SparseOperator.from_coo(
    np.r_[np.arange(n), senders],
    np.r_[np.arange(n), receivers],
    jnp.asarray(np.r_[deg, -w]),
    (n, n),
    symmetric=True,
)
sym = gx.symbolic_cholesky(G.pattern)  # host, once: RCM, banded layout


def logdet(log_kappa):
    Q = G.add_diagonal(jnp.full(n, jnp.exp(2 * log_kappa)))  # same pattern
    return gx.sparse_cholesky(Q, sym).logdet()


jax.vmap(logdet)(jnp.linspace(-1.0, 1.0, 16))  # 16 factorisations, one analysis
jax.grad(logdet)(0.0)  # d log|Q| / d log κ, through one Takahashi sweep

Structured linear algebra and Gaussian primitives for JAX.

SparseOperator

Bases: AbstractLinearOperator

Sparse matrix with a static SparsityPattern and traced values.

Only values is a pytree leaf, so jit, grad and vmap act on the non-zeros while the pattern stays a compile-time constant. Changing the values (eqx.tree_at(lambda op: op.values, Q, new_values)) never retraces; changing the pattern does.

The primitives dispatch on it: gaussx.diag reads the diagonal from the pattern; gaussx.solve uses CG for large positive semidefinite operators (AutoSolver rules) and a dense solve otherwise; gaussx.logdet uses SLQLogdet for large PSD operators; gaussx.eig with rank= runs Lanczos on the matvec; gaussx.cholesky returns a sparse SparseCholeskyFactor. For exact solves, log-determinants and marginal variances through that factor, pass SparseCholeskySolver.

Parameters:

Name Type Description Default
values Float[ArrayLike, ' nnz']

Stored non-zeros in the pattern's canonical order, shape (nnz,). For a symmetric pattern, the lower triangle only.

required
pattern SparsityPattern

The static sparsity pattern.

required
tags object | frozenset[object]

Lineax tags. lx.symmetric_tag is added for a symmetric pattern.

frozenset()
Example
import numpy as np

# ICAR structure matrix of the path graph 0 - 1 - 2 (edges once)
senders, receivers = np.array([1, 2]), np.array([0, 1])
deg = np.bincount(np.r_[senders, receivers], minlength=3)
R = gaussx.SparseOperator.from_coo(
    np.r_[np.arange(3), senders],
    np.r_[np.arange(3), receivers],
    jnp.asarray(np.r_[deg, -np.ones(2)]),
    (3, 3),
    symmetric=True,
)
R.mv(jnp.ones(3))  # [0, 0, 0]: constants are in the null space
Source code in src/gaussx/_operators/_sparse.py
class SparseOperator(lx.AbstractLinearOperator):
    r"""Sparse matrix with a static `SparsityPattern` and traced values.

    Only ``values`` is a pytree leaf, so ``jit``, ``grad`` and ``vmap`` act on
    the non-zeros while the pattern stays a compile-time constant. Changing
    the values (``eqx.tree_at(lambda op: op.values, Q, new_values)``) never
    retraces; changing the pattern does.

    The primitives dispatch on it: `gaussx.diag` reads the diagonal from the
    pattern; `gaussx.solve` uses CG for large positive semidefinite operators
    (`AutoSolver` rules) and a dense solve otherwise; `gaussx.logdet` uses
    `SLQLogdet` for large PSD operators; `gaussx.eig` with ``rank=`` runs
    Lanczos on the matvec; `gaussx.cholesky` returns a sparse
    `SparseCholeskyFactor`. For exact solves, log-determinants and marginal
    variances through that factor, pass `SparseCholeskySolver`.

    Args:
        values: Stored non-zeros in the pattern's canonical order, shape
            ``(nnz,)``. For a symmetric pattern, the lower triangle only.
        pattern: The static sparsity pattern.
        tags: Lineax tags. ``lx.symmetric_tag`` is added for a symmetric
            pattern.

    Example:
        ```python
        import numpy as np

        # ICAR structure matrix of the path graph 0 - 1 - 2 (edges once)
        senders, receivers = np.array([1, 2]), np.array([0, 1])
        deg = np.bincount(np.r_[senders, receivers], minlength=3)
        R = gaussx.SparseOperator.from_coo(
            np.r_[np.arange(3), senders],
            np.r_[np.arange(3), receivers],
            jnp.asarray(np.r_[deg, -np.ones(2)]),
            (3, 3),
            symmetric=True,
        )
        R.mv(jnp.ones(3))  # [0, 0, 0]: constants are in the null space
        ```
    """

    values: Float[Array, " nnz"]
    pattern: SparsityPattern = eqx.field(static=True)
    tags: frozenset[object] = eqx.field(static=True)

    def __init__(
        self,
        values: Float[ArrayLike, " nnz"],
        pattern: SparsityPattern,
        *,
        tags: object | frozenset[object] = frozenset(),
    ) -> None:
        values = jnp.asarray(values)
        if values.shape != (pattern.nnz,):
            raise ValueError(
                f"values must have shape ({pattern.nnz},) to match the pattern, "
                f"got {values.shape}."
            )
        if not jnp.issubdtype(values.dtype, jnp.inexact):
            values = values.astype(jnp.result_type(values.dtype, jnp.float32))
        self.values = values
        self.pattern = pattern
        tags = _to_frozenset(tags)
        if pattern.symmetric:
            tags = tags | {lx.symmetric_tag}
        self.tags = tags

    @classmethod
    def from_coo(
        cls,
        rows: ArrayLike,
        cols: ArrayLike,
        values: Float[ArrayLike, " k"],
        shape: tuple[int, int],
        *,
        symmetric: bool = False,
        tags: object | frozenset[object] = frozenset(),
    ) -> SparseOperator:
        """Build from coordinate (COO) triplets.

        ``rows`` and ``cols`` must be concrete (host) integer arrays; only
        ``values`` may be traced. Duplicate coordinates are summed, and the
        diagonal of a square matrix is added with zeros where missing.

        Args:
            rows: Row indices, shape ``(k,)``.
            cols: Column indices, shape ``(k,)``.
            values: Entry values, shape ``(k,)``.
            shape: Matrix shape ``(m, n)``.
            symmetric: The matrix is symmetric and each off-diagonal pair is
                given **once** (either ``(i, j)`` or ``(j, i)``); it is mirrored
                to the other triangle. Giving both would sum them. This is
                the edge-once storage of an undirected graph.
            tags: Lineax tags.

        Returns:
            The operator, with values in the pattern's canonical order.
        """
        rows_, cols_, inverse = _canonicalise(rows, cols, shape, symmetric)
        pattern = SparsityPattern._from_canonical(rows_, cols_, shape, symmetric)
        values = jnp.asarray(values)
        if values.shape != (inverse.shape[0],):
            raise ValueError(
                f"values must have shape ({inverse.shape[0]},) to match rows and "
                f"cols, got {values.shape}."
            )
        values = jax.ops.segment_sum(values, inverse, num_segments=pattern.nnz)
        return cls(values, pattern, tags=tags)

    def _full_values(self) -> Float[Array, " nnz_full"]:
        _, _, index = self.pattern._full
        return self.values if index is None else self.values[index]

    def mv(self, vector: Float[Array, " n"]) -> Float[Array, " m"]:
        # segment_sum over the row-sorted full pattern; it beat the BCOO
        # matvec at every size in the G1 benchmark
        # (tests/operators/test_sparse.py::test_matvec_benchmark).
        rows, cols, _ = self.pattern._full
        return jax.ops.segment_sum(
            self._full_values() * vector[cols],
            rows,
            num_segments=self.pattern.shape[0],
            indices_are_sorted=True,
        )

    def as_matrix(self) -> Float[Array, "m n"]:
        rows, cols, _ = self.pattern._full
        dense = jnp.zeros(self.pattern.shape, dtype=self.values.dtype)
        return dense.at[rows, cols].add(self._full_values())

    def to_bcoo(self) -> jsparse.BCOO:
        """The full matrix as a `jax.experimental.sparse.BCOO` array.

        A symmetric pattern is expanded to both triangles.

        Returns:
            A ``BCOO`` with sorted, unique indices.
        """
        rows, cols, _ = self.pattern._full
        indices = jnp.asarray(np.column_stack([rows, cols]))
        return jsparse.BCOO(
            (self._full_values(), indices),
            shape=self.pattern.shape,
            indices_sorted=True,
            unique_indices=True,
        )

    def transpose(self) -> SparseOperator:
        if self.pattern.symmetric:
            return self
        pattern, perm = _transpose_plan(self.pattern)
        return SparseOperator(
            self.values[perm], pattern, tags=lx.transpose_tags(self.tags)
        )

    def in_structure(self) -> jax.ShapeDtypeStruct:
        return jax.ShapeDtypeStruct((self.pattern.shape[1],), self.values.dtype)

    def out_structure(self) -> jax.ShapeDtypeStruct:
        return jax.ShapeDtypeStruct((self.pattern.shape[0],), self.values.dtype)

    def diagonal(self) -> Float[Array, " k"]:
        """The diagonal, read from the pattern (no densification).

        Returns:
            ``diag(A)``, shape ``(min(m, n),)``.
        """
        pos = self.pattern._diagonal_positions
        out = jnp.zeros(min(self.pattern.shape), dtype=self.values.dtype)
        return out.at[self.pattern.rows[pos]].add(self.values[pos])

    def add_diagonal(
        self,
        d: Float[Array, " n"],
        *,
        tags: object | frozenset[object] | None = None,
    ) -> SparseOperator:
        """``A + diag(d)`` on the same pattern.

        Args:
            d: Diagonal to add, shape ``(n,)``.
            tags: Tags of the result. By default only symmetry is kept,
                since an arbitrary ``d`` can break definiteness; pass
                ``lx.positive_semidefinite_tag`` when you know it holds.

        Returns:
            The shifted operator, with the identical pattern.
        """
        n = self._square_size("add_diagonal")
        d = jnp.asarray(d)
        if d.shape != (n,):
            raise ValueError(f"d must have shape ({n},), got {d.shape}.")
        pos = self.pattern._diagonal_positions
        values = self.values.at[pos].add(d)
        if tags is None:
            tags = self.tags & {lx.symmetric_tag}
        return SparseOperator(values, self.pattern, tags=tags)

    def union(
        self,
        other: SparseOperator,
        *,
        tags: object | frozenset[object] | None = None,
    ) -> SparseOperator:
        """``A + B`` on the union of the two patterns.

        The union pattern and the scatter positions are computed on the host
        once per pair of patterns (and cached); only the values are added in
        JAX. Two symmetric patterns stay symmetric; otherwise the symmetric
        operand is expanded to both triangles.

        Args:
            other: Operator of the same shape.
            tags: Tags of the result. Defaults to the tags both operands share
                (a sum of PSD operators is PSD), minus ``unit_diagonal_tag``.

        Returns:
            The sum, on the union pattern.
        """
        if self.pattern.shape != other.pattern.shape:
            raise ValueError(
                f"Shapes differ: {self.pattern.shape} vs {other.pattern.shape}."
            )
        pattern, pos_a, pos_b = _union_plan(self.pattern, other.pattern)
        values_a = self._storage_values(pattern.symmetric)
        values_b = other._storage_values(pattern.symmetric)
        dtype = jnp.result_type(values_a, values_b)
        values = (
            jnp.zeros(pattern.nnz, dtype=dtype)
            .at[pos_a]
            .add(values_a)
            .at[pos_b]
            .add(values_b)
        )
        if tags is None:
            tags = (self.tags & other.tags) - {lx.unit_diagonal_tag}
        return SparseOperator(values, pattern, tags=tags)

    def congruence(
        self,
        A: SparseOperator,
        w: Float[Array, " m"],
        *,
        tags: object | frozenset[object] | None = None,
    ) -> SparseOperator:
        r"""``Aᵀ diag(w) A`` on the union of ``self``'s pattern and ``AᵀA``'s.

        ``(AᵀWA)_{ij} = Σ_k A_{ki} w_k A_{kj}`` is non-zero only where columns
        ``i`` and ``j`` share a row of ``A``. That pattern, unioned with this
        operator's own, and the index triples ``(p, q, k)`` feeding each entry
        are computed on the host once per ``(pattern, A.pattern)``; the values
        are one ``segment_sum`` in JAX. Because the result already lives on
        ``self``'s pattern (padded with zeros), ``self.union(result)`` is an
        aligned add, and a projector whose rows only touch neighbours in
        ``self`` (a FEM projector on a mesh precision) leaves the pattern
        unchanged.

        Args:
            A: Operator of shape ``(m, n)``, with ``n`` the size of ``self``.
            w: Row weights, shape ``(m,)``.
            tags: Tags of the result. Defaults to symmetry only (``w`` may
                have either sign).

        Returns:
            The congruence, symmetric storage iff ``self``'s pattern is.
        """
        n = self._square_size("congruence")
        m, n_a = A.pattern.shape
        if n_a != n:
            raise ValueError(f"A must have {n} columns, got shape {A.pattern.shape}.")
        w = jnp.asarray(w)
        if w.shape != (m,):
            raise ValueError(f"w must have shape ({m},), got {w.shape}.")
        pattern, p, q, k, target = _congruence_plan(self.pattern, A.pattern)
        full = A._full_values()
        contributions = full[p] * w[k] * full[q]
        values = jax.ops.segment_sum(contributions, target, num_segments=pattern.nnz)
        if tags is None:
            tags = frozenset()
        return SparseOperator(values, pattern, tags=tags)

    def _storage_values(self, symmetric: bool) -> Float[Array, " nnz"]:
        """Values in symmetric (stored) or general (full) storage."""
        return self.values if symmetric else self._full_values()

    def _square_size(self, method: str) -> int:
        m, n = self.pattern.shape
        if m != n:
            raise ValueError(f"{method} needs a square operator, got {(m, n)}.")
        return n

from_coo(rows: ArrayLike, cols: ArrayLike, values: Float[ArrayLike, ' k'], shape: tuple[int, int], *, symmetric: bool = False, tags: object | frozenset[object] = frozenset()) -> SparseOperator classmethod

Build from coordinate (COO) triplets.

rows and cols must be concrete (host) integer arrays; only values may be traced. Duplicate coordinates are summed, and the diagonal of a square matrix is added with zeros where missing.

Parameters:

Name Type Description Default
rows ArrayLike

Row indices, shape (k,).

required
cols ArrayLike

Column indices, shape (k,).

required
values Float[ArrayLike, ' k']

Entry values, shape (k,).

required
shape tuple[int, int]

Matrix shape (m, n).

required
symmetric bool

The matrix is symmetric and each off-diagonal pair is given once (either (i, j) or (j, i)); it is mirrored to the other triangle. Giving both would sum them. This is the edge-once storage of an undirected graph.

False
tags object | frozenset[object]

Lineax tags.

frozenset()

Returns:

Type Description
SparseOperator

The operator, with values in the pattern's canonical order.

Source code in src/gaussx/_operators/_sparse.py
@classmethod
def from_coo(
    cls,
    rows: ArrayLike,
    cols: ArrayLike,
    values: Float[ArrayLike, " k"],
    shape: tuple[int, int],
    *,
    symmetric: bool = False,
    tags: object | frozenset[object] = frozenset(),
) -> SparseOperator:
    """Build from coordinate (COO) triplets.

    ``rows`` and ``cols`` must be concrete (host) integer arrays; only
    ``values`` may be traced. Duplicate coordinates are summed, and the
    diagonal of a square matrix is added with zeros where missing.

    Args:
        rows: Row indices, shape ``(k,)``.
        cols: Column indices, shape ``(k,)``.
        values: Entry values, shape ``(k,)``.
        shape: Matrix shape ``(m, n)``.
        symmetric: The matrix is symmetric and each off-diagonal pair is
            given **once** (either ``(i, j)`` or ``(j, i)``); it is mirrored
            to the other triangle. Giving both would sum them. This is
            the edge-once storage of an undirected graph.
        tags: Lineax tags.

    Returns:
        The operator, with values in the pattern's canonical order.
    """
    rows_, cols_, inverse = _canonicalise(rows, cols, shape, symmetric)
    pattern = SparsityPattern._from_canonical(rows_, cols_, shape, symmetric)
    values = jnp.asarray(values)
    if values.shape != (inverse.shape[0],):
        raise ValueError(
            f"values must have shape ({inverse.shape[0]},) to match rows and "
            f"cols, got {values.shape}."
        )
    values = jax.ops.segment_sum(values, inverse, num_segments=pattern.nnz)
    return cls(values, pattern, tags=tags)

to_bcoo() -> jsparse.BCOO

The full matrix as a jax.experimental.sparse.BCOO array.

A symmetric pattern is expanded to both triangles.

Returns:

Type Description
BCOO

A BCOO with sorted, unique indices.

Source code in src/gaussx/_operators/_sparse.py
def to_bcoo(self) -> jsparse.BCOO:
    """The full matrix as a `jax.experimental.sparse.BCOO` array.

    A symmetric pattern is expanded to both triangles.

    Returns:
        A ``BCOO`` with sorted, unique indices.
    """
    rows, cols, _ = self.pattern._full
    indices = jnp.asarray(np.column_stack([rows, cols]))
    return jsparse.BCOO(
        (self._full_values(), indices),
        shape=self.pattern.shape,
        indices_sorted=True,
        unique_indices=True,
    )

diagonal() -> Float[Array, ' k']

The diagonal, read from the pattern (no densification).

Returns:

Type Description
Float[Array, ' k']

diag(A), shape (min(m, n),).

Source code in src/gaussx/_operators/_sparse.py
def diagonal(self) -> Float[Array, " k"]:
    """The diagonal, read from the pattern (no densification).

    Returns:
        ``diag(A)``, shape ``(min(m, n),)``.
    """
    pos = self.pattern._diagonal_positions
    out = jnp.zeros(min(self.pattern.shape), dtype=self.values.dtype)
    return out.at[self.pattern.rows[pos]].add(self.values[pos])

add_diagonal(d: Float[Array, ' n'], *, tags: object | frozenset[object] | None = None) -> SparseOperator

A + diag(d) on the same pattern.

Parameters:

Name Type Description Default
d Float[Array, ' n']

Diagonal to add, shape (n,).

required
tags object | frozenset[object] | None

Tags of the result. By default only symmetry is kept, since an arbitrary d can break definiteness; pass lx.positive_semidefinite_tag when you know it holds.

None

Returns:

Type Description
SparseOperator

The shifted operator, with the identical pattern.

Source code in src/gaussx/_operators/_sparse.py
def add_diagonal(
    self,
    d: Float[Array, " n"],
    *,
    tags: object | frozenset[object] | None = None,
) -> SparseOperator:
    """``A + diag(d)`` on the same pattern.

    Args:
        d: Diagonal to add, shape ``(n,)``.
        tags: Tags of the result. By default only symmetry is kept,
            since an arbitrary ``d`` can break definiteness; pass
            ``lx.positive_semidefinite_tag`` when you know it holds.

    Returns:
        The shifted operator, with the identical pattern.
    """
    n = self._square_size("add_diagonal")
    d = jnp.asarray(d)
    if d.shape != (n,):
        raise ValueError(f"d must have shape ({n},), got {d.shape}.")
    pos = self.pattern._diagonal_positions
    values = self.values.at[pos].add(d)
    if tags is None:
        tags = self.tags & {lx.symmetric_tag}
    return SparseOperator(values, self.pattern, tags=tags)

union(other: SparseOperator, *, tags: object | frozenset[object] | None = None) -> SparseOperator

A + B on the union of the two patterns.

The union pattern and the scatter positions are computed on the host once per pair of patterns (and cached); only the values are added in JAX. Two symmetric patterns stay symmetric; otherwise the symmetric operand is expanded to both triangles.

Parameters:

Name Type Description Default
other SparseOperator

Operator of the same shape.

required
tags object | frozenset[object] | None

Tags of the result. Defaults to the tags both operands share (a sum of PSD operators is PSD), minus unit_diagonal_tag.

None

Returns:

Type Description
SparseOperator

The sum, on the union pattern.

Source code in src/gaussx/_operators/_sparse.py
def union(
    self,
    other: SparseOperator,
    *,
    tags: object | frozenset[object] | None = None,
) -> SparseOperator:
    """``A + B`` on the union of the two patterns.

    The union pattern and the scatter positions are computed on the host
    once per pair of patterns (and cached); only the values are added in
    JAX. Two symmetric patterns stay symmetric; otherwise the symmetric
    operand is expanded to both triangles.

    Args:
        other: Operator of the same shape.
        tags: Tags of the result. Defaults to the tags both operands share
            (a sum of PSD operators is PSD), minus ``unit_diagonal_tag``.

    Returns:
        The sum, on the union pattern.
    """
    if self.pattern.shape != other.pattern.shape:
        raise ValueError(
            f"Shapes differ: {self.pattern.shape} vs {other.pattern.shape}."
        )
    pattern, pos_a, pos_b = _union_plan(self.pattern, other.pattern)
    values_a = self._storage_values(pattern.symmetric)
    values_b = other._storage_values(pattern.symmetric)
    dtype = jnp.result_type(values_a, values_b)
    values = (
        jnp.zeros(pattern.nnz, dtype=dtype)
        .at[pos_a]
        .add(values_a)
        .at[pos_b]
        .add(values_b)
    )
    if tags is None:
        tags = (self.tags & other.tags) - {lx.unit_diagonal_tag}
    return SparseOperator(values, pattern, tags=tags)

congruence(A: SparseOperator, w: Float[Array, ' m'], *, tags: object | frozenset[object] | None = None) -> SparseOperator

Aᵀ diag(w) A on the union of self's pattern and AᵀA's.

(AᵀWA)_{ij} = Σ_k A_{ki} w_k A_{kj} is non-zero only where columns i and j share a row of A. That pattern, unioned with this operator's own, and the index triples (p, q, k) feeding each entry are computed on the host once per (pattern, A.pattern); the values are one segment_sum in JAX. Because the result already lives on self's pattern (padded with zeros), self.union(result) is an aligned add, and a projector whose rows only touch neighbours in self (a FEM projector on a mesh precision) leaves the pattern unchanged.

Parameters:

Name Type Description Default
A SparseOperator

Operator of shape (m, n), with n the size of self.

required
w Float[Array, ' m']

Row weights, shape (m,).

required
tags object | frozenset[object] | None

Tags of the result. Defaults to symmetry only (w may have either sign).

None

Returns:

Type Description
SparseOperator

The congruence, symmetric storage iff self's pattern is.

Source code in src/gaussx/_operators/_sparse.py
def congruence(
    self,
    A: SparseOperator,
    w: Float[Array, " m"],
    *,
    tags: object | frozenset[object] | None = None,
) -> SparseOperator:
    r"""``Aᵀ diag(w) A`` on the union of ``self``'s pattern and ``AᵀA``'s.

    ``(AᵀWA)_{ij} = Σ_k A_{ki} w_k A_{kj}`` is non-zero only where columns
    ``i`` and ``j`` share a row of ``A``. That pattern, unioned with this
    operator's own, and the index triples ``(p, q, k)`` feeding each entry
    are computed on the host once per ``(pattern, A.pattern)``; the values
    are one ``segment_sum`` in JAX. Because the result already lives on
    ``self``'s pattern (padded with zeros), ``self.union(result)`` is an
    aligned add, and a projector whose rows only touch neighbours in
    ``self`` (a FEM projector on a mesh precision) leaves the pattern
    unchanged.

    Args:
        A: Operator of shape ``(m, n)``, with ``n`` the size of ``self``.
        w: Row weights, shape ``(m,)``.
        tags: Tags of the result. Defaults to symmetry only (``w`` may
            have either sign).

    Returns:
        The congruence, symmetric storage iff ``self``'s pattern is.
    """
    n = self._square_size("congruence")
    m, n_a = A.pattern.shape
    if n_a != n:
        raise ValueError(f"A must have {n} columns, got shape {A.pattern.shape}.")
    w = jnp.asarray(w)
    if w.shape != (m,):
        raise ValueError(f"w must have shape ({m},), got {w.shape}.")
    pattern, p, q, k, target = _congruence_plan(self.pattern, A.pattern)
    full = A._full_values()
    contributions = full[p] * w[k] * full[q]
    values = jax.ops.segment_sum(contributions, target, num_segments=pattern.nnz)
    if tags is None:
        tags = frozenset()
    return SparseOperator(values, pattern, tags=tags)

SparsityPattern

Static, hashable sparsity pattern of a (m, n) matrix.

The index arrays live on the host (NumPy) and are canonicalised on construction:

  • duplicate (row, col) pairs are merged;
  • entries are sorted row-major (by row, then column);
  • for a square pattern the diagonal is always present, so diag, add_diagonal and a later Cholesky never change the pattern;
  • with symmetric=True only the lower triangle (row >= col) is stored, and an upper-triangle pair (i, j) is stored as (j, i).

The hash is a content hash (SHA-256 of the canonical index arrays, the shape and the symmetry flag), so it is stable across processes and can key a cache of symbolic analyses. The index arrays are read-only.

Parameters:

Name Type Description Default
rows ArrayLike

Row indices, shape (k,).

required
cols ArrayLike

Column indices, shape (k,).

required
shape tuple[int, int]

Matrix shape (m, n).

required
symmetric bool

Store the lower triangle of a symmetric matrix. Requires a square shape.

False

Raises:

Type Description
ValueError

If the index arrays differ in length, are not rank 1, are out of range, or symmetric=True with a non-square shape.

Example
import numpy as np

# Path graph 0 - 1 - 2, each edge given once
p = gaussx.SparsityPattern(
    np.array([1, 2]), np.array([0, 1]), (3, 3), symmetric=True
)
p.rows, p.cols  # ([0, 1, 1, 2, 2], [0, 0, 1, 1, 2]): diagonal added
Source code in src/gaussx/_operators/_sparse.py
class SparsityPattern:
    r"""Static, hashable sparsity pattern of a ``(m, n)`` matrix.

    The index arrays live on the host (NumPy) and are canonicalised on
    construction:

    - duplicate ``(row, col)`` pairs are merged;
    - entries are sorted row-major (by row, then column);
    - for a square pattern the diagonal is always present, so ``diag``,
      ``add_diagonal`` and a later Cholesky never change the pattern;
    - with ``symmetric=True`` only the lower triangle (``row >= col``) is
      stored, and an upper-triangle pair ``(i, j)`` is stored as ``(j, i)``.

    The hash is a content hash (SHA-256 of the canonical index arrays, the
    shape and the symmetry flag), so it is stable across processes and can key
    a cache of symbolic analyses. The index arrays are read-only.

    Args:
        rows: Row indices, shape ``(k,)``.
        cols: Column indices, shape ``(k,)``.
        shape: Matrix shape ``(m, n)``.
        symmetric: Store the lower triangle of a symmetric matrix. Requires a
            square ``shape``.

    Raises:
        ValueError: If the index arrays differ in length, are not rank 1, are
            out of range, or ``symmetric=True`` with a non-square shape.

    Example:
        ```python
        import numpy as np

        # Path graph 0 - 1 - 2, each edge given once
        p = gaussx.SparsityPattern(
            np.array([1, 2]), np.array([0, 1]), (3, 3), symmetric=True
        )
        p.rows, p.cols  # ([0, 1, 1, 2, 2], [0, 0, 1, 1, 2]): diagonal added
        ```
    """

    rows: np.ndarray
    cols: np.ndarray
    shape: tuple[int, int]
    symmetric: bool
    _digest: str

    def __init__(
        self,
        rows: ArrayLike,
        cols: ArrayLike,
        shape: tuple[int, int],
        *,
        symmetric: bool = False,
    ) -> None:
        rows_, cols_, _ = _canonicalise(rows, cols, shape, symmetric)
        self._set(rows_, cols_, shape, symmetric)

    @classmethod
    def _from_canonical(
        cls,
        rows: np.ndarray,
        cols: np.ndarray,
        shape: tuple[int, int],
        symmetric: bool,
    ) -> SparsityPattern:
        pattern = cls.__new__(cls)
        pattern._set(rows, cols, shape, symmetric)
        return pattern

    def _set(
        self,
        rows: np.ndarray,
        cols: np.ndarray,
        shape: tuple[int, int],
        symmetric: bool,
    ) -> None:
        rows = np.ascontiguousarray(rows, dtype=np.int32)
        cols = np.ascontiguousarray(cols, dtype=np.int32)
        rows.flags.writeable = False
        cols.flags.writeable = False
        shape = (int(shape[0]), int(shape[1]))
        digest = hashlib.sha256()
        digest.update(repr((shape, bool(symmetric))).encode())
        digest.update(rows.astype("<i4").tobytes())
        digest.update(cols.astype("<i4").tobytes())
        object.__setattr__(self, "rows", rows)
        object.__setattr__(self, "cols", cols)
        object.__setattr__(self, "shape", shape)
        object.__setattr__(self, "symmetric", bool(symmetric))
        object.__setattr__(self, "_digest", digest.hexdigest())

    def __setattr__(self, name: str, value: Any) -> None:
        raise AttributeError("SparsityPattern is immutable.")

    @property
    def nnz(self) -> int:
        """Number of stored entries (the length of ``values``)."""
        return int(self.rows.shape[0])

    @property
    def digest(self) -> str:
        """Hex SHA-256 content hash, identical across processes."""
        return self._digest

    def __hash__(self) -> int:
        return int(self._digest[:15], 16)  # 60 bits: never truncated by hash()

    def __eq__(self, other: object) -> bool:
        if self is other:
            return True
        if not isinstance(other, SparsityPattern):
            return NotImplemented
        return (
            self._digest == other._digest
            and self.shape == other.shape
            and self.symmetric == other.symmetric
            and np.array_equal(self.rows, other.rows)
            and np.array_equal(self.cols, other.cols)
        )

    def __repr__(self) -> str:
        return (
            f"SparsityPattern(shape={self.shape}, nnz={self.nnz}, "
            f"symmetric={self.symmetric})"
        )

    @ft.cached_property
    def _diagonal_positions(self) -> np.ndarray:
        """Positions of the stored diagonal entries."""
        return np.flatnonzero(self.rows == self.cols).astype(np.int32)

    @ft.cached_property
    def _full(self) -> tuple[np.ndarray, np.ndarray, np.ndarray | None]:
        """The full matrix's entries in row-major order.

        Returns ``(rows, cols, value_index)``: entry ``e`` of the full matrix is
        ``values[value_index[e]]``. ``value_index`` is ``None`` when the stored
        entries already are the full matrix (general storage).
        """
        if not self.symmetric:
            return self.rows, self.cols, None
        off = np.flatnonzero(self.rows != self.cols)
        rows = np.concatenate([self.rows, self.cols[off]])
        cols = np.concatenate([self.cols, self.rows[off]])
        index = np.concatenate([np.arange(self.nnz), off])
        order = np.lexsort((cols, rows))
        return (
            rows[order].astype(np.int32),
            cols[order].astype(np.int32),
            index[order].astype(np.int32),
        )

nnz: int property

Number of stored entries (the length of values).

digest: str property

Hex SHA-256 content hash, identical across processes.

SymbolicCholesky

Symbolic Cholesky factor of a symmetric sparsity pattern.

For the permuted matrix A = P Q Pᵀ (A[k, l] = Q[perm[k], perm[l]]) it holds the elimination tree and the pattern of L (A = L Lᵀ) in compressed sparse column (CSC) form, with each column's rows sorted so the diagonal comes first, together with the static index plans the numeric phase, the triangular solves and the Takahashi recursion gather through. The column patterns follow the elimination tree,

\[ \operatorname{struct}(L_{:,j}) = \operatorname{struct}(A_{j:,j})\ \cup \bigcup_{\operatorname{parent}(c)=j}\operatorname{struct}(L_{:,c})\setminus\{c\}, \qquad \operatorname{parent}(j) = \min\{i>j : L_{ij}\neq 0\}. \]

Build it with gaussx.symbolic_cholesky (cached), not directly. It is hashable and compared by (pattern, ordering, backend, banded), so it sits in a static field and causes no retrace for equal inputs.

Parameters:

Name Type Description Default
pattern SparsityPattern

Square sparsity pattern to analyse.

required
ordering Ordering

Name of the ordering that produced perm.

required
backend Backend

Numeric backend, "jax" or "cholmod".

required
perm ndarray

Fill-reducing permutation, A = Q[perm][:, perm].

required
banded bool | None

Force the banded (True) or windowed (False) layout; None picks by storage cost.

None

Attributes:

Name Type Description
pattern

The pattern it was computed for.

ordering

The fill-reducing ordering used.

backend

The numeric backend ("jax" or "cholmod").

n

Matrix size.

perm

A = Q[perm][:, perm], shape (n,).

iperm

The inverse permutation, iperm[perm] = arange(n).

parent

Elimination tree, parent[j] (-1 at a root).

colptr

CSC column pointers of L, shape (n + 1,).

rowidx

CSC row indices of L, shape (nnz_L,).

colidx

Column of each entry of L, shape (nnz_L,).

nnz

Number of entries of L (diagonal included): the fill.

nnz_lower

Number of entries in the lower triangle of Q.

max_col

Longest column of L (padded window of the column plans).

max_row

Longest row of L below the diagonal (update-list padding).

banded

Whether the numeric phase and Takahashi run on dense block_size × block_size blocks of a block-tridiagonal layout (chosen when the blocks cost at most four times the storage of L, as after RCM on a mesh) rather than on gathered windows.

block_size

The bandwidth of L (at least 1).

Source code in src/gaussx/_sparse/_symbolic.py
class SymbolicCholesky:
    r"""Symbolic Cholesky factor of a symmetric sparsity pattern.

    For the permuted matrix ``A = P Q Pᵀ`` (``A[k, l] = Q[perm[k], perm[l]]``)
    it holds the elimination tree and the pattern of ``L`` (``A = L Lᵀ``) in
    compressed sparse column (CSC) form, with each column's rows sorted so the
    diagonal comes first, together with the static index plans the numeric
    phase, the triangular solves and the Takahashi recursion gather through.
    The column patterns follow the elimination tree,

    $$
    \operatorname{struct}(L_{:,j}) = \operatorname{struct}(A_{j:,j})\ \cup
    \bigcup_{\operatorname{parent}(c)=j}\operatorname{struct}(L_{:,c})\setminus\{c\},
    \qquad \operatorname{parent}(j) = \min\{i>j : L_{ij}\neq 0\}.
    $$

    Build it with `gaussx.symbolic_cholesky` (cached), not directly. It is
    hashable and compared by ``(pattern, ordering, backend, banded)``, so it sits in
    a static field and causes no retrace for equal inputs.

    Args:
        pattern: Square sparsity pattern to analyse.
        ordering: Name of the ordering that produced ``perm``.
        backend: Numeric backend, ``"jax"`` or ``"cholmod"``.
        perm: Fill-reducing permutation, ``A = Q[perm][:, perm]``.
        banded: Force the banded (``True``) or windowed (``False``) layout;
            ``None`` picks by storage cost.

    Attributes:
        pattern: The pattern it was computed for.
        ordering: The fill-reducing ordering used.
        backend: The numeric backend (``"jax"`` or ``"cholmod"``).
        n: Matrix size.
        perm: ``A = Q[perm][:, perm]``, shape ``(n,)``.
        iperm: The inverse permutation, ``iperm[perm] = arange(n)``.
        parent: Elimination tree, ``parent[j]`` (``-1`` at a root).
        colptr: CSC column pointers of ``L``, shape ``(n + 1,)``.
        rowidx: CSC row indices of ``L``, shape ``(nnz_L,)``.
        colidx: Column of each entry of ``L``, shape ``(nnz_L,)``.
        nnz: Number of entries of ``L`` (diagonal included): the fill.
        nnz_lower: Number of entries in the lower triangle of ``Q``.
        max_col: Longest column of ``L`` (padded window of the column plans).
        max_row: Longest row of ``L`` below the diagonal (update-list padding).
        banded: Whether the numeric phase and Takahashi run on dense
            ``block_size × block_size`` blocks of a block-tridiagonal layout
            (chosen when the blocks cost at most four times the storage of
            ``L``, as after RCM on a mesh) rather than on gathered windows.
        block_size: The bandwidth of ``L`` (at least 1).
    """

    def __init__(
        self,
        pattern: SparsityPattern,
        ordering: Ordering,
        backend: Backend,
        perm: np.ndarray,
        *,
        banded: bool | None = None,
    ) -> None:
        n = pattern.shape[0]
        perm = np.asarray(perm, dtype=np.int64)
        iperm = np.empty(n, dtype=np.int64)
        iperm[perm] = np.arange(n)
        self.pattern = pattern
        self.ordering = ordering
        self.backend = backend
        self.n = n
        self.perm = perm.astype(np.int32)
        self.iperm = iperm.astype(np.int32)

        # Lower triangle of A = P Q Pᵀ (structure only), grouped by column.
        r = iperm[pattern.rows]
        c = iperm[pattern.cols]
        lo, hi = np.minimum(r, c), np.maximum(r, c)
        keys = np.unique(lo * n + hi)
        self.nnz_lower = int(keys.shape[0])
        a_cols, a_rows = keys // n, keys % n
        a_ptr = np.searchsorted(a_cols, np.arange(n + 1))

        # Column patterns of L by the elimination-tree union, in column order:
        # every child c < j is finished before its parent j needs it.
        parent = np.full(n, -1, dtype=np.int64)
        children: list[list[int]] = [[] for _ in range(n)]
        structs: list[np.ndarray] = [np.empty(0, dtype=np.int64)] * n
        for j in range(n):
            pieces = [a_rows[a_ptr[j] : a_ptr[j + 1]]]
            pieces.extend(structs[child][1:] for child in children[j])
            struct = pieces[0] if len(pieces) == 1 else _sorted_union(pieces)
            structs[j] = struct
            if struct.shape[0] > 1:
                parent[j] = struct[1]
                children[struct[1]].append(j)
        counts = np.array([s.shape[0] for s in structs], dtype=np.int64)
        colptr = np.concatenate([[0], np.cumsum(counts)])
        rowidx = np.concatenate(structs) if n else np.empty(0, dtype=np.int64)
        colidx = np.repeat(np.arange(n), counts)
        nnz = int(colptr[-1])

        self.parent = parent.astype(np.int32)
        self.colptr = colptr.astype(np.int32)
        self.rowidx = rowidx.astype(np.int32)
        self.colidx = colidx.astype(np.int32)
        self.nnz = nnz
        self.max_col = int(counts.max(initial=1))

        # Row structure of L below the diagonal (the left-looking update list
        # of each column j: the columns k < j with L[j, k] != 0), as CSR over
        # the strictly lower entries: positions of L[j, k] sorted by row j.
        off = np.flatnonzero(rowidx != colidx)
        order = off[np.lexsort((colidx[off], rowidx[off]))]
        row_counts = np.bincount(rowidx[order], minlength=n)
        self.rowptr = np.concatenate([[0], np.cumsum(row_counts)]).astype(np.int32)
        self.rowpos = order.astype(np.int32)
        self.max_row = int(max(row_counts.max(initial=0), 1))

        # Scatter plan from the operator's stored values to the lower
        # triangle of A on L's pattern: A_lower = segment_sum(w · values, target).
        # Full storage contributes ½ from each of (i, j) and (j, i), so the
        # factored matrix is the symmetric part ½(Q + Qᵀ).
        self.value_target = _lower_positions(self, lo, hi)
        weight = np.ones(pattern.nnz)
        if not pattern.symmetric:
            weight[pattern.rows != pattern.cols] = 0.5
        self.value_weight = weight

        # Banded layout: with bandwidth b, A and L are block tridiagonal in
        # b × b blocks, and dense block kernels (BLAS) replace the gathers.
        # RCM keeps a mesh's fill inside a band, so the blocks cost a small
        # multiple of the storage of L; an arrow-shaped pattern (a dense row)
        # does not.
        b = int((self.rowidx - self.colidx).max(initial=0))
        size = max(b, 1)
        blocks = -(-n // size) if n else 0
        if banded is None:
            banded = 2 * blocks * size * size <= 4 * max(nnz, 1)
        self.banded = bool(banded)
        self.block_size = size
        self.num_blocks = blocks
        kr, kc = self.rowidx // size, self.colidx // size
        local = (self.rowidx % size) * size + self.colidx % size
        offset = np.where(kr == kc, 0, blocks * size * size)
        self.block_pos = (offset + kc * size * size + local).astype(np.int64)
        pad = np.arange(n, blocks * size)
        self.pad_pos = ((pad // size) * size * size + (pad % size) * (size + 1)).astype(
            np.int64
        )

    # -- identity -----------------------------------------------------------

    @property
    def _key(self) -> tuple[SparsityPattern, str, str, bool]:
        return (self.pattern, self.ordering, self.backend, self.banded)

    def __hash__(self) -> int:
        return hash(self._key)

    def __eq__(self, other: object) -> bool:
        if not isinstance(other, SymbolicCholesky):
            return NotImplemented
        return self is other or self._key == other._key

    def __repr__(self) -> str:
        return (
            f"SymbolicCholesky(n={self.n}, nnz_L={self.nnz}, "
            f"ordering={self.ordering!r}, backend={self.backend!r}, "
            f"banded={self.banded})"
        )

    @property
    def fill_ratio(self) -> float:
        """``nnz(L) / nnz(tril(Q))``: how much the factor fills in."""
        return self.nnz / max(self.nnz_lower, 1)

    # -- plans for the selected inverse -------------------------------------

    @ft.cached_property
    def inverse_plan(self) -> tuple[SparsityPattern, np.ndarray]:
        """Pattern of ``L + Lᵀ`` in the original order, and the gather into it.

        Returns ``(pattern, index)``: the symmetric (lower-triangle) pattern of
        ``Pᵀ (L + Lᵀ) P``, and ``index`` such that its canonical values are
        ``z[index]`` for ``z`` on ``L``'s CSC pattern.
        """
        r = self.perm[self.rowidx].astype(np.int64)
        c = self.perm[self.colidx].astype(np.int64)
        lower = SparsityPattern(
            np.maximum(r, c), np.minimum(r, c), (self.n, self.n), symmetric=True
        )
        pos = _positions(lower, np.maximum(r, c), np.minimum(r, c))
        index = np.empty(self.nnz, dtype=np.int32)
        index[pos] = np.arange(self.nnz)
        return lower, index

fill_ratio: float property

nnz(L) / nnz(tril(Q)): how much the factor fills in.

inverse_plan: tuple[SparsityPattern, np.ndarray] cached property

Pattern of L + Lᵀ in the original order, and the gather into it.

Returns (pattern, index): the symmetric (lower-triangle) pattern of Pᵀ (L + Lᵀ) P, and index such that its canonical values are z[index] for z on L's CSC pattern.

SparseCholeskyFactor

Bases: Module

Sparse Cholesky factor P Q Pᵀ = L Lᵀ on a static symbolic pattern.

Built by gaussx.sparse_cholesky (or gaussx.cholesky on a SparseOperator). values are the entries of L on the CSC pattern of symbolic; matrix_values are those of the factored matrix (its lower triangle, permuted, on the same pattern), the input that logdet and solve differentiate through their custom VJPs:

  • logdet: \(d\log|Q| = \operatorname{tr}(Q^{-1}dQ)\), so the cotangent is \(Q^{-1}\) on pattern(Q), which Takahashi evaluates in one sweep;
  • solve: \(\bar b = Q^{-1}\bar x\) and \(\bar Q = -\bar b\,x^\top\), symmetrised on the pattern.

Gradients reach the operator's stored values per stored value: with symmetric=True storage an off-diagonal value sets Q_ij and Q_ji, so its cotangent is doubled (2 Z_ij); general storage factors ½(Q + Qᵀ) and each stored value gets Z_ij. The other methods are differentiated by JAX through the factorisation.

Attributes:

Name Type Description
values Float[Array, ' nnz_L']

L on its CSC pattern, shape (nnz_L,).

matrix_values Float[Array, ' nnz_L']

Lower triangle of P Q Pᵀ on L's pattern.

symbolic SymbolicCholesky

The static symbolic analysis.

Examples:

import jax.numpy as jnp
import numpy as np
import gaussx

# Precision of a path graph 0 - 1 - 2 - 3 plus a unit shift
n = 4
Q = gaussx.SparseOperator.from_coo(
    np.r_[np.arange(n), np.arange(1, n)],
    np.r_[np.arange(n), np.arange(n - 1)],
    jnp.r_[jnp.array([2.0, 3.0, 3.0, 2.0]), -jnp.ones(n - 1)],
    (n, n),
    symmetric=True,
)
factor = gaussx.sparse_cholesky(Q)
dense = Q.as_matrix()
assert jnp.allclose(factor.logdet(), jnp.linalg.slogdet(dense)[1])
b = jnp.ones(n)
assert jnp.allclose(factor.solve(b), jnp.linalg.solve(dense, b))
assert jnp.allclose(factor.diag_inv(), jnp.diag(jnp.linalg.inv(dense)))
Source code in src/gaussx/_sparse/_factor.py
class SparseCholeskyFactor(eqx.Module):
    r"""Sparse Cholesky factor ``P Q Pᵀ = L Lᵀ`` on a static symbolic pattern.

    Built by `gaussx.sparse_cholesky` (or `gaussx.cholesky` on a
    `SparseOperator`). ``values`` are the entries of ``L`` on the CSC pattern
    of ``symbolic``; ``matrix_values`` are those of the factored matrix (its
    lower triangle, permuted, on the same pattern), the input that
    ``logdet`` and ``solve`` differentiate through their custom VJPs:

    - ``logdet``: $d\log|Q| = \operatorname{tr}(Q^{-1}dQ)$, so the cotangent is
      $Q^{-1}$ on ``pattern(Q)``, which Takahashi evaluates in one sweep;
    - ``solve``: $\bar b = Q^{-1}\bar x$ and $\bar Q = -\bar b\,x^\top$,
      symmetrised on the pattern.

    Gradients reach the operator's stored values per stored value: with
    ``symmetric=True`` storage an off-diagonal value sets ``Q_ij`` and
    ``Q_ji``, so its cotangent is doubled (``2 Z_ij``); general storage
    factors ``½(Q + Qᵀ)`` and each stored value gets ``Z_ij``. The other
    methods are differentiated by JAX through the factorisation.

    Attributes:
        values: ``L`` on its CSC pattern, shape ``(nnz_L,)``.
        matrix_values: Lower triangle of ``P Q Pᵀ`` on ``L``'s pattern.
        symbolic: The static symbolic analysis.

    Examples:
        ```python
        import jax.numpy as jnp
        import numpy as np
        import gaussx

        # Precision of a path graph 0 - 1 - 2 - 3 plus a unit shift
        n = 4
        Q = gaussx.SparseOperator.from_coo(
            np.r_[np.arange(n), np.arange(1, n)],
            np.r_[np.arange(n), np.arange(n - 1)],
            jnp.r_[jnp.array([2.0, 3.0, 3.0, 2.0]), -jnp.ones(n - 1)],
            (n, n),
            symmetric=True,
        )
        factor = gaussx.sparse_cholesky(Q)
        dense = Q.as_matrix()
        assert jnp.allclose(factor.logdet(), jnp.linalg.slogdet(dense)[1])
        b = jnp.ones(n)
        assert jnp.allclose(factor.solve(b), jnp.linalg.solve(dense, b))
        assert jnp.allclose(factor.diag_inv(), jnp.diag(jnp.linalg.inv(dense)))
        ```
    """

    values: Float[Array, " nnz_L"]
    matrix_values: Float[Array, " nnz_L"]
    symbolic: SymbolicCholesky = eqx.field(static=True)

    def _L(self) -> Array:
        # logdet and solve route their first-order cotangents to
        # ``matrix_values`` and give ``L`` none. ``L`` stays differentiable on
        # the JAX backend so that reverse-over-reverse (a Hessian) sees how
        # the Takahashi / adjoint cotangents in their backward passes move
        # with the values. CHOLMOD's callback has no derivative.
        if self.symbolic.backend == "cholmod":
            return jax.lax.stop_gradient(self.values)
        return self.values

    def solve(self, b: Float[Array, " n"]) -> Float[Array, " n"]:
        """``Q⁻¹ b = Pᵀ L⁻ᵀ L⁻¹ P b``.

        Args:
            b: Right-hand side, shape ``(n,)``.

        Returns:
            The solution, shape ``(n,)``.
        """
        sym = self.symbolic
        y = b[jnp.asarray(sym.perm)]
        z = _vjp.solve(sym, self.matrix_values, self._L(), y)
        return z[jnp.asarray(sym.iperm)]

    def logdet(self) -> Float[Array, ""]:
        """``log|Q| = 2 Σ_j log L_jj``.

        Returns:
            The log-determinant (NaN if ``Q`` is not positive definite).
        """
        return _vjp.logdet(self.symbolic, self.matrix_values, self._L())

    def solve_lower_transpose(self, z: Float[Array, " n"]) -> Float[Array, " n"]:
        """``x = Pᵀ L⁻ᵀ z``: with ``z ~ N(0, I)``, ``x ~ N(0, Q⁻¹)``.

        Args:
            z: Shape ``(n,)``.

        Returns:
            ``x``, shape ``(n,)``.
        """
        sym = self.symbolic
        return solve_upper(sym, self.values, z)[jnp.asarray(sym.iperm)]

    def selected_inverse(self) -> SparseOperator:
        """``Q⁻¹`` on the pattern of ``L + Lᵀ``, in the original order.

        The pattern contains ``pattern(Q)``. Like the block selected inverse,
        the result holds entries of ``Q⁻¹``; it is not an operator equal to
        ``Q⁻¹`` (whose other entries are not zero).

        Returns:
            A symmetric `SparseOperator` holding the selected entries.
        """
        pattern, index = self.symbolic.inverse_plan
        Z = takahashi(self.symbolic, self.values)
        return SparseOperator(Z[jnp.asarray(index)], pattern)

    def diag_inv(self) -> Float[Array, " n"]:
        """``diag(Q⁻¹)``, the marginal variances, by one Takahashi sweep.

        Returns:
            Shape ``(n,)``, in the original order.
        """
        sym = self.symbolic
        Z = takahashi(sym, self.values)
        return Z[jnp.asarray(sym.colptr[sym.iperm])]

solve(b: Float[Array, ' n']) -> Float[Array, ' n']

Q⁻¹ b = Pᵀ L⁻ᵀ L⁻¹ P b.

Parameters:

Name Type Description Default
b Float[Array, ' n']

Right-hand side, shape (n,).

required

Returns:

Type Description
Float[Array, ' n']

The solution, shape (n,).

Source code in src/gaussx/_sparse/_factor.py
def solve(self, b: Float[Array, " n"]) -> Float[Array, " n"]:
    """``Q⁻¹ b = Pᵀ L⁻ᵀ L⁻¹ P b``.

    Args:
        b: Right-hand side, shape ``(n,)``.

    Returns:
        The solution, shape ``(n,)``.
    """
    sym = self.symbolic
    y = b[jnp.asarray(sym.perm)]
    z = _vjp.solve(sym, self.matrix_values, self._L(), y)
    return z[jnp.asarray(sym.iperm)]

logdet() -> Float[Array, '']

log|Q| = 2 Σ_j log L_jj.

Returns:

Type Description
Float[Array, '']

The log-determinant (NaN if Q is not positive definite).

Source code in src/gaussx/_sparse/_factor.py
def logdet(self) -> Float[Array, ""]:
    """``log|Q| = 2 Σ_j log L_jj``.

    Returns:
        The log-determinant (NaN if ``Q`` is not positive definite).
    """
    return _vjp.logdet(self.symbolic, self.matrix_values, self._L())

solve_lower_transpose(z: Float[Array, ' n']) -> Float[Array, ' n']

x = Pᵀ L⁻ᵀ z: with z ~ N(0, I), x ~ N(0, Q⁻¹).

Parameters:

Name Type Description Default
z Float[Array, ' n']

Shape (n,).

required

Returns:

Type Description
Float[Array, ' n']

x, shape (n,).

Source code in src/gaussx/_sparse/_factor.py
def solve_lower_transpose(self, z: Float[Array, " n"]) -> Float[Array, " n"]:
    """``x = Pᵀ L⁻ᵀ z``: with ``z ~ N(0, I)``, ``x ~ N(0, Q⁻¹)``.

    Args:
        z: Shape ``(n,)``.

    Returns:
        ``x``, shape ``(n,)``.
    """
    sym = self.symbolic
    return solve_upper(sym, self.values, z)[jnp.asarray(sym.iperm)]

selected_inverse() -> SparseOperator

Q⁻¹ on the pattern of L + Lᵀ, in the original order.

The pattern contains pattern(Q). Like the block selected inverse, the result holds entries of Q⁻¹; it is not an operator equal to Q⁻¹ (whose other entries are not zero).

Returns:

Type Description
SparseOperator

A symmetric SparseOperator holding the selected entries.

Source code in src/gaussx/_sparse/_factor.py
def selected_inverse(self) -> SparseOperator:
    """``Q⁻¹`` on the pattern of ``L + Lᵀ``, in the original order.

    The pattern contains ``pattern(Q)``. Like the block selected inverse,
    the result holds entries of ``Q⁻¹``; it is not an operator equal to
    ``Q⁻¹`` (whose other entries are not zero).

    Returns:
        A symmetric `SparseOperator` holding the selected entries.
    """
    pattern, index = self.symbolic.inverse_plan
    Z = takahashi(self.symbolic, self.values)
    return SparseOperator(Z[jnp.asarray(index)], pattern)

diag_inv() -> Float[Array, ' n']

diag(Q⁻¹), the marginal variances, by one Takahashi sweep.

Returns:

Type Description
Float[Array, ' n']

Shape (n,), in the original order.

Source code in src/gaussx/_sparse/_factor.py
def diag_inv(self) -> Float[Array, " n"]:
    """``diag(Q⁻¹)``, the marginal variances, by one Takahashi sweep.

    Returns:
        Shape ``(n,)``, in the original order.
    """
    sym = self.symbolic
    Z = takahashi(sym, self.values)
    return Z[jnp.asarray(sym.colptr[sym.iperm])]

symbolic_cholesky(pattern: SparsityPattern, *, ordering: Ordering = 'rcm', backend: Backend = 'jax') -> SymbolicCholesky

Symbolic analysis of a sparse Cholesky factorisation (host, cached).

Computes a fill-reducing permutation, the elimination tree and the pattern of the factor L of P Q Pᵀ = L Lᵀ, plus the static, padded index plans that gaussx.sparse_cholesky and the Takahashi selected inverse gather through. It depends only on the pattern, so it runs once on the host and is cached per (pattern, ordering, backend) (the pattern's content hash): later calls, under jit or not, are a dictionary lookup.

A symmetric pattern stores the lower triangle; a general pattern is symmetrised structurally, and its factor is that of the symmetric part ½(Q + Qᵀ).

Parameters:

Name Type Description Default
pattern SparsityPattern

Square sparsity pattern of the matrix to factor.

required
ordering Ordering

"rcm" (reverse Cuthill-McKee from scipy.sparse.csgraph; minimises bandwidth), "natural" (no permutation) or "amd" (approximate minimum degree, from CHOLMOD; needs scikit-sparse). RCM fill grows like the bandwidth times n; AMD and nested dissection are much sparser on 2-D meshes.

'rcm'
backend Backend

"jax" (numeric factorisation as a lax.scan; jit, grad and vmap work) or "cholmod" (numeric factorisation by CHOLMOD on the host through jax.pure_callback; CPU only, vmap runs sequentially; needs scikit-sparse). Triangular solves, Takahashi and the gradients are the same JAX code for both.

'jax'

Returns:

Type Description
SymbolicCholesky

The cached SymbolicCholesky.

Raises:

Type Description
ValueError

For a non-square pattern or an unknown option.

ImportError

For "amd" or "cholmod" without scikit-sparse.

Examples:

import numpy as np
import gaussx

# Path graph 0 - 1 - 2 - 3 - 4 (lower triangle, edges once)
p = gaussx.SparsityPattern(
    np.arange(1, 5), np.arange(4), (5, 5), symmetric=True
)
sym = gaussx.symbolic_cholesky(p)
assert sym.nnz == 9  # a tridiagonal matrix does not fill in
assert gaussx.symbolic_cholesky(p) is sym  # cached per pattern
Source code in src/gaussx/_sparse/_symbolic.py
def symbolic_cholesky(
    pattern: SparsityPattern,
    *,
    ordering: Ordering = "rcm",
    backend: Backend = "jax",
) -> SymbolicCholesky:
    r"""Symbolic analysis of a sparse Cholesky factorisation (host, cached).

    Computes a fill-reducing permutation, the elimination tree and the
    pattern of the factor ``L`` of ``P Q Pᵀ = L Lᵀ``, plus the static,
    padded index plans that `gaussx.sparse_cholesky` and the Takahashi
    selected inverse gather through. It depends only on the pattern, so it
    runs once on the host and is cached per ``(pattern, ordering, backend)``
    (the pattern's content hash): later calls, under ``jit`` or not, are a
    dictionary lookup.

    A symmetric pattern stores the lower triangle; a general pattern is
    symmetrised structurally, and its factor is that of the symmetric part
    ``½(Q + Qᵀ)``.

    Args:
        pattern: Square sparsity pattern of the matrix to factor.
        ordering: ``"rcm"`` (reverse Cuthill-McKee from
            ``scipy.sparse.csgraph``; minimises bandwidth), ``"natural"`` (no
            permutation) or ``"amd"`` (approximate minimum degree, from
            CHOLMOD; needs ``scikit-sparse``). RCM fill grows like the
            bandwidth times ``n``; AMD and nested dissection are much sparser
            on 2-D meshes.
        backend: ``"jax"`` (numeric factorisation as a ``lax.scan``; ``jit``,
            ``grad`` and ``vmap`` work) or ``"cholmod"`` (numeric
            factorisation by CHOLMOD on the host through
            ``jax.pure_callback``; CPU only, ``vmap`` runs sequentially; needs
            ``scikit-sparse``). Triangular solves, Takahashi and the gradients
            are the same JAX code for both.

    Returns:
        The cached `SymbolicCholesky`.

    Raises:
        ValueError: For a non-square pattern or an unknown option.
        ImportError: For ``"amd"`` or ``"cholmod"`` without ``scikit-sparse``.

    Examples:
        ```python
        import numpy as np
        import gaussx

        # Path graph 0 - 1 - 2 - 3 - 4 (lower triangle, edges once)
        p = gaussx.SparsityPattern(
            np.arange(1, 5), np.arange(4), (5, 5), symmetric=True
        )
        sym = gaussx.symbolic_cholesky(p)
        assert sym.nnz == 9  # a tridiagonal matrix does not fill in
        assert gaussx.symbolic_cholesky(p) is sym  # cached per pattern
        ```
    """
    if ordering not in _ORDERINGS:
        raise ValueError(
            f"Unknown ordering {ordering!r}; expected one of {_ORDERINGS}."
        )
    if backend not in _BACKENDS:
        raise ValueError(f"Unknown backend {backend!r}; expected one of {_BACKENDS}.")
    _check_pattern(pattern)
    return _symbolic_cached(pattern, ordering, backend)

sparse_cholesky(op: SparseOperator, symbolic: SymbolicCholesky | None = None) -> SparseCholeskyFactor

Sparse Cholesky factorisation of a symmetric positive-definite operator.

The symbolic analysis (ordering, elimination tree, pattern of L) is host-side and cached per pattern; pass one from gaussx.symbolic_cholesky to choose the ordering or backend, or to reuse it explicitly. The numeric phase is traced: it jits, vmaps over op.values and is differentiable. With the JAX backend it is a left-looking lax.scan over columns, sequential in n and O(Σ_j |struct(L_{:,j})|²) work.

Parameters:

Name Type Description Default
op SparseOperator

Symmetric positive-definite operator. A general (non-symmetric) pattern is factored as ½(Q + Qᵀ).

required
symbolic SymbolicCholesky | None

Symbolic analysis of op.pattern. Defaults to symbolic_cholesky(op.pattern) (RCM ordering, JAX backend).

None

Returns:

Type Description
SparseCholeskyFactor

The factor.

Raises:

Type Description
TypeError

If op is not a SparseOperator.

ValueError

If symbolic was computed for a different pattern.

Examples:

import equinox as eqx
import jax
import jax.numpy as jnp
import numpy as np
import gaussx

# A 1-D random-walk precision τ R + I: analyse once, factor for many τ
n = 6
R = gaussx.SparseOperator.from_coo(
    np.r_[np.arange(n), np.arange(1, n)],
    np.r_[np.arange(n), np.arange(n - 1)],
    jnp.r_[jnp.array([1.0] + [2.0] * (n - 2) + [1.0]), -jnp.ones(n - 1)],
    (n, n),
    symmetric=True,
)
sym = gaussx.symbolic_cholesky(R.pattern)  # host, once

def logdet(log_tau):
    Q = eqx.tree_at(lambda op: op.values, R, jnp.exp(log_tau) * R.values)
    return gaussx.sparse_cholesky(Q.add_diagonal(jnp.ones(n)), sym).logdet()

values = jax.vmap(logdet)(jnp.linspace(-1.0, 1.0, 4))
slope = jax.grad(logdet)(0.0)  # through one Takahashi sweep
Source code in src/gaussx/_sparse/_factor.py
def sparse_cholesky(
    op: SparseOperator, symbolic: SymbolicCholesky | None = None
) -> SparseCholeskyFactor:
    r"""Sparse Cholesky factorisation of a symmetric positive-definite operator.

    The symbolic analysis (ordering, elimination tree, pattern of ``L``) is
    host-side and cached per pattern; pass one from `gaussx.symbolic_cholesky`
    to choose the ordering or backend, or to reuse it explicitly. The numeric
    phase is traced: it ``jit``s, ``vmap``s over ``op.values`` and is
    differentiable. With the JAX backend it is a left-looking ``lax.scan``
    over columns, sequential in ``n`` and ``O(Σ_j |struct(L_{:,j})|²)`` work.

    Args:
        op: Symmetric positive-definite operator. A general (non-symmetric)
            pattern is factored as ``½(Q + Qᵀ)``.
        symbolic: Symbolic analysis of ``op.pattern``. Defaults to
            ``symbolic_cholesky(op.pattern)`` (RCM ordering, JAX backend).

    Returns:
        The factor.

    Raises:
        TypeError: If ``op`` is not a `SparseOperator`.
        ValueError: If ``symbolic`` was computed for a different pattern.

    Examples:
        ```python
        import equinox as eqx
        import jax
        import jax.numpy as jnp
        import numpy as np
        import gaussx

        # A 1-D random-walk precision τ R + I: analyse once, factor for many τ
        n = 6
        R = gaussx.SparseOperator.from_coo(
            np.r_[np.arange(n), np.arange(1, n)],
            np.r_[np.arange(n), np.arange(n - 1)],
            jnp.r_[jnp.array([1.0] + [2.0] * (n - 2) + [1.0]), -jnp.ones(n - 1)],
            (n, n),
            symmetric=True,
        )
        sym = gaussx.symbolic_cholesky(R.pattern)  # host, once

        def logdet(log_tau):
            Q = eqx.tree_at(lambda op: op.values, R, jnp.exp(log_tau) * R.values)
            return gaussx.sparse_cholesky(Q.add_diagonal(jnp.ones(n)), sym).logdet()

        values = jax.vmap(logdet)(jnp.linspace(-1.0, 1.0, 4))
        slope = jax.grad(logdet)(0.0)  # through one Takahashi sweep
        ```
    """
    if not isinstance(op, SparseOperator):
        raise TypeError(
            f"sparse_cholesky needs a SparseOperator, got {type(op).__name__}."
        )
    if symbolic is None:
        symbolic = symbolic_cholesky(op.pattern)
    elif symbolic.pattern != op.pattern:
        raise ValueError(
            "symbolic was computed for a different sparsity pattern: "
            f"{symbolic.pattern} vs {op.pattern}."
        )
    a = lower_values(symbolic, op.values)
    if symbolic.backend == "cholmod":
        from gaussx._sparse._cholmod import cholmod_cholesky

        L = cholmod_cholesky(symbolic, a)
    else:
        L = numeric_cholesky(symbolic, a)
    return SparseCholeskyFactor(L, a, symbolic)