Skip to content

Sketching

Random sketching operators \(S \in \mathbb{R}^{d \times m}\) that compress a tall matrix \(A \in \mathbb{R}^{m \times n}\) to \(SA \in \mathbb{R}^{d \times n}\) while approximately preserving the geometry of its range. \(S\) is an \(\varepsilon\)-subspace embedding for \(\operatorname{range}(A)\) if

\[ (1-\varepsilon)\|Ax\| \le \|SAx\| \le (1+\varepsilon)\|Ax\| \qquad \forall x . \]

Sketches are the foundation of the randomized linear-algebra stack (range finders, randomized SVD, sketch-and-precondition least squares).

Sketch Size d for an \(\varepsilon\)-embedding Cost to apply
GaussianSketch \(O(n/\varepsilon^2)\) \(O(dmn)\)
SparseSignSketch (default for tall problems) \(O(n\log n/\varepsilon^2)\) \(O(\text{nnz}\cdot mn)\)
SRHTSketch \(O((n+\log m)\log n/\varepsilon^2)\) \(O(mn\log m)\)

A sketch is sampled once, with an explicit PRNG key (key=None means jax.random.PRNGKey(0)), and its random draws live in the module: apply (\(SA\)) and apply_transpose (\(S^\top Y\)) always refer to the same \(S\). sketch_operator sketches a matrix-free lineax operator with \(d\) transpose-matvecs, and as_operator returns \(S\) itself as a lineax operator.

# Sketch a tall Jacobian (10⁶ residuals × 200 parameters) down to 800 rows
S = gx.SparseSignSketch.sample(key, d=800, m=1_000_000, nnz=8)
SJ = S.sketch_operator(J_op)  # (800, 200), matrix-free
sv = jnp.linalg.svd(SJ, compute_uv=False)  # J's singular values, within (1 ± ε)

Abstract interface

Structured linear algebra and Gaussian primitives for JAX.

AbstractSketch

Bases: Module

A random sketching matrix \(S \in \mathbb{R}^{d \times m}\), sampled once.

A sketch compresses a tall matrix \(A \in \mathbb{R}^{m \times n}\) to \(SA \in \mathbb{R}^{d \times n}\) with \(d \ll m\) while approximately preserving the geometry of \(\operatorname{range}(A)\). \(S\) is an \(\varepsilon\)-subspace embedding for \(\operatorname{range}(A)\) if

\[ (1-\varepsilon)\|Ax\| \le \|SAx\| \le (1+\varepsilon)\|Ax\| \qquad \forall x . \]

The random draws live in the module, so apply and apply_transpose always refer to the same \(S\). Concrete sketches are built with their sample classmethod, which takes a PRNG key (key=None means jax.random.PRNGKey(0)).

Sketches cast their stored values to the dtype of the array they are applied to, so a float32 input never meets a float64 sketch.

Attributes:

Name Type Description
in_size AbstractVar[int]

Number of columns \(m\) of \(S\) (rows of the sketched input).

out_size AbstractVar[int]

Number of rows \(d\) of \(S\) (the sketch size).

Source code in src/gaussx/_sketching/_base.py
class AbstractSketch(eqx.Module):
    r"""A random sketching matrix $S \in \mathbb{R}^{d \times m}$, sampled once.

    A sketch compresses a tall matrix $A \in \mathbb{R}^{m \times n}$ to
    $SA \in \mathbb{R}^{d \times n}$ with $d \ll m$ while approximately
    preserving the geometry of $\operatorname{range}(A)$. $S$ is an
    $\varepsilon$-subspace embedding for $\operatorname{range}(A)$ if

    $$
    (1-\varepsilon)\|Ax\| \le \|SAx\| \le (1+\varepsilon)\|Ax\|
    \qquad \forall x .
    $$

    The random draws live in the module, so `apply` and `apply_transpose`
    always refer to the same $S$. Concrete sketches are built with their
    `sample` classmethod, which takes a PRNG key (``key=None`` means
    ``jax.random.PRNGKey(0)``).

    Sketches cast their stored values to the dtype of the array they are
    applied to, so a float32 input never meets a float64 sketch.

    Attributes:
        in_size: Number of columns $m$ of $S$ (rows of the sketched input).
        out_size: Number of rows $d$ of $S$ (the sketch size).
    """

    in_size: eqx.AbstractVar[int]
    out_size: eqx.AbstractVar[int]

    @abc.abstractmethod
    def apply(self, A: Float[Array, "m ..."]) -> Float[Array, "d ..."]:
        """Compute $S A$ along the leading axis of ``A``."""

    @abc.abstractmethod
    def apply_transpose(self, Y: Float[Array, "d ..."]) -> Float[Array, "m ..."]:
        r"""Compute $S^\top Y$ along the leading axis of ``Y``."""

    def sketch_operator(self, op: lx.AbstractLinearOperator) -> Float[Array, "d n"]:
        r"""Sketch a (possibly matrix-free) operator: $S A$.

        A `lineax.MatrixLinearOperator` is sketched directly with `apply`.
        Any other operator is sketched with $d$ transpose-matvecs,
        $SA = (A^\top S^\top)^\top$, vmapped over the rows of $S$; this
        materialises $S^\top$ as an $(m, d)$ block but never forms $A$.

        Args:
            op: Operator $A$ of shape ``(m, n)``.

        Returns:
            The dense sketch $SA$, shape ``(d, n)``.

        Raises:
            ValueError: If ``op.out_size()`` is not the sketch's ``in_size``.
        """
        if op.out_size() != self.in_size:
            raise ValueError(
                f"Cannot sketch an operator with {op.out_size()} rows using a "
                f"sketch with in_size={self.in_size}."
            )
        if isinstance(op, lx.MatrixLinearOperator):
            return self.apply(op.matrix)
        dtype = op.out_structure().dtype
        rows_of_s = self.apply_transpose(jnp.eye(self.out_size, dtype=dtype))
        return jax.vmap(op.transpose().mv, in_axes=1)(rows_of_s)

    def as_operator(self) -> lx.AbstractLinearOperator:
        """Return $S$ as a matrix-free ``(d, m)`` lineax operator."""
        dtype = jax.tree.leaves(eqx.filter(self, eqx.is_inexact_array))[0].dtype
        return lx.FunctionLinearOperator(
            self.apply, jax.ShapeDtypeStruct((self.in_size,), dtype)
        )

apply(A: Float[Array, 'm ...']) -> Float[Array, 'd ...'] abstractmethod

Compute \(S A\) along the leading axis of A.

Source code in src/gaussx/_sketching/_base.py
@abc.abstractmethod
def apply(self, A: Float[Array, "m ..."]) -> Float[Array, "d ..."]:
    """Compute $S A$ along the leading axis of ``A``."""

apply_transpose(Y: Float[Array, 'd ...']) -> Float[Array, 'm ...'] abstractmethod

Compute \(S^\top Y\) along the leading axis of Y.

Source code in src/gaussx/_sketching/_base.py
@abc.abstractmethod
def apply_transpose(self, Y: Float[Array, "d ..."]) -> Float[Array, "m ..."]:
    r"""Compute $S^\top Y$ along the leading axis of ``Y``."""

sketch_operator(op: lx.AbstractLinearOperator) -> Float[Array, 'd n']

Sketch a (possibly matrix-free) operator: \(S A\).

A lineax.MatrixLinearOperator is sketched directly with apply. Any other operator is sketched with \(d\) transpose-matvecs, \(SA = (A^\top S^\top)^\top\), vmapped over the rows of \(S\); this materialises \(S^\top\) as an \((m, d)\) block but never forms \(A\).

Parameters:

Name Type Description Default
op AbstractLinearOperator

Operator \(A\) of shape (m, n).

required

Returns:

Type Description
Float[Array, 'd n']

The dense sketch \(SA\), shape (d, n).

Raises:

Type Description
ValueError

If op.out_size() is not the sketch's in_size.

Source code in src/gaussx/_sketching/_base.py
def sketch_operator(self, op: lx.AbstractLinearOperator) -> Float[Array, "d n"]:
    r"""Sketch a (possibly matrix-free) operator: $S A$.

    A `lineax.MatrixLinearOperator` is sketched directly with `apply`.
    Any other operator is sketched with $d$ transpose-matvecs,
    $SA = (A^\top S^\top)^\top$, vmapped over the rows of $S$; this
    materialises $S^\top$ as an $(m, d)$ block but never forms $A$.

    Args:
        op: Operator $A$ of shape ``(m, n)``.

    Returns:
        The dense sketch $SA$, shape ``(d, n)``.

    Raises:
        ValueError: If ``op.out_size()`` is not the sketch's ``in_size``.
    """
    if op.out_size() != self.in_size:
        raise ValueError(
            f"Cannot sketch an operator with {op.out_size()} rows using a "
            f"sketch with in_size={self.in_size}."
        )
    if isinstance(op, lx.MatrixLinearOperator):
        return self.apply(op.matrix)
    dtype = op.out_structure().dtype
    rows_of_s = self.apply_transpose(jnp.eye(self.out_size, dtype=dtype))
    return jax.vmap(op.transpose().mv, in_axes=1)(rows_of_s)

as_operator() -> lx.AbstractLinearOperator

Return \(S\) as a matrix-free (d, m) lineax operator.

Source code in src/gaussx/_sketching/_base.py
def as_operator(self) -> lx.AbstractLinearOperator:
    """Return $S$ as a matrix-free ``(d, m)`` lineax operator."""
    dtype = jax.tree.leaves(eqx.filter(self, eqx.is_inexact_array))[0].dtype
    return lx.FunctionLinearOperator(
        self.apply, jax.ShapeDtypeStruct((self.in_size,), dtype)
    )

Sketches

Structured linear algebra and Gaussian primitives for JAX.

GaussianSketch

Bases: _DenseSketch

Gaussian sketch: i.i.d. entries \(S_{ij} \sim \mathcal{N}(0, 1/d)\).

The scaling gives \(\mathbb{E}[S^\top S] = I_m\). A Gaussian sketch of size \(d = O(n/\varepsilon^2)\) is an \(\varepsilon\)-subspace embedding for any \(n\)-dimensional subspace; applying it costs \(O(dmn)\).

Attributes:

Name Type Description
matrix Float[Array, 'd m']

The sketching matrix, shape (d, m).

in_size int

\(m\).

out_size int

\(d\).

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> import gaussx as gx
>>> S = gx.GaussianSketch.sample(jr.key(0), d=20, m=500)
>>> S.apply(jnp.ones((500, 3))).shape
(20, 3)
Source code in src/gaussx/_sketching/_dense.py
class GaussianSketch(_DenseSketch):
    r"""Gaussian sketch: i.i.d. entries $S_{ij} \sim \mathcal{N}(0, 1/d)$.

    The scaling gives $\mathbb{E}[S^\top S] = I_m$. A Gaussian sketch of size
    $d = O(n/\varepsilon^2)$ is an $\varepsilon$-subspace embedding for any
    $n$-dimensional subspace; applying it costs $O(dmn)$.

    Attributes:
        matrix: The sketching matrix, shape ``(d, m)``.
        in_size: $m$.
        out_size: $d$.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> import gaussx as gx
        >>> S = gx.GaussianSketch.sample(jr.key(0), d=20, m=500)
        >>> S.apply(jnp.ones((500, 3))).shape
        (20, 3)
    """

    @classmethod
    def sample(cls, key: jax.Array | None, d: int, m: int) -> GaussianSketch:
        """Draw a ``(d, m)`` Gaussian sketch.

        Args:
            key: PRNG key. ``None`` means ``jax.random.PRNGKey(0)``.
            d: Sketch size (rows of $S$).
            m: Input size (columns of $S$).

        Returns:
            The sampled `GaussianSketch`.
        """
        if key is None:
            key = jax.random.PRNGKey(0)
        matrix = jax.random.normal(key, (d, m)) / jnp.sqrt(d)
        return cls(matrix=matrix, in_size=m, out_size=d)

sample(key: jax.Array | None, d: int, m: int) -> GaussianSketch classmethod

Draw a (d, m) Gaussian sketch.

Parameters:

Name Type Description Default
key Array | None

PRNG key. None means jax.random.PRNGKey(0).

required
d int

Sketch size (rows of \(S\)).

required
m int

Input size (columns of \(S\)).

required

Returns:

Type Description
GaussianSketch

The sampled GaussianSketch.

Source code in src/gaussx/_sketching/_dense.py
@classmethod
def sample(cls, key: jax.Array | None, d: int, m: int) -> GaussianSketch:
    """Draw a ``(d, m)`` Gaussian sketch.

    Args:
        key: PRNG key. ``None`` means ``jax.random.PRNGKey(0)``.
        d: Sketch size (rows of $S$).
        m: Input size (columns of $S$).

    Returns:
        The sampled `GaussianSketch`.
    """
    if key is None:
        key = jax.random.PRNGKey(0)
    matrix = jax.random.normal(key, (d, m)) / jnp.sqrt(d)
    return cls(matrix=matrix, in_size=m, out_size=d)

OrthonormalSketch

Bases: _DenseSketch

Gaussian sketch with orthonormal rows, \(S S^\top = I_d\).

Built from the thin QR of an \((m, d)\) Gaussian block, so the row space of \(S\) is a uniformly random \(d\)-dimensional subspace of \(\mathbb{R}^m\). Rows are orthonormal, so \(\mathbb{E}[S^\top S] = (d/m)\, I_m\): rescale by \(\sqrt{m/d}\) when an isotropic embedding is needed. Requires \(d \le m\).

Attributes:

Name Type Description
matrix Float[Array, 'd m']

The sketching matrix with orthonormal rows, shape (d, m).

in_size int

\(m\).

out_size int

\(d\).

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> import gaussx as gx
>>> S = gx.OrthonormalSketch.sample(jr.key(0), d=4, m=50)
>>> SSt = S.apply(S.apply_transpose(jnp.eye(4)))
>>> bool(jnp.allclose(SSt, jnp.eye(4), atol=1e-5))
True
Source code in src/gaussx/_sketching/_dense.py
class OrthonormalSketch(_DenseSketch):
    r"""Gaussian sketch with orthonormal rows, $S S^\top = I_d$.

    Built from the thin QR of an $(m, d)$ Gaussian block, so the row space of
    $S$ is a uniformly random $d$-dimensional subspace of $\mathbb{R}^m$.
    Rows are orthonormal, so $\mathbb{E}[S^\top S] = (d/m)\, I_m$: rescale by
    $\sqrt{m/d}$ when an isotropic embedding is needed. Requires $d \le m$.

    Attributes:
        matrix: The sketching matrix with orthonormal rows, shape ``(d, m)``.
        in_size: $m$.
        out_size: $d$.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> import gaussx as gx
        >>> S = gx.OrthonormalSketch.sample(jr.key(0), d=4, m=50)
        >>> SSt = S.apply(S.apply_transpose(jnp.eye(4)))
        >>> bool(jnp.allclose(SSt, jnp.eye(4), atol=1e-5))
        True
    """

    @classmethod
    def sample(cls, key: jax.Array | None, d: int, m: int) -> OrthonormalSketch:
        """Draw a ``(d, m)`` sketch with orthonormal rows.

        Args:
            key: PRNG key. ``None`` means ``jax.random.PRNGKey(0)``.
            d: Sketch size (rows of $S$); must satisfy ``d <= m``.
            m: Input size (columns of $S$).

        Returns:
            The sampled `OrthonormalSketch`.

        Raises:
            ValueError: If ``d > m``.
        """
        if d > m:
            raise ValueError(f"OrthonormalSketch needs d <= m, got d={d}, m={m}.")
        if key is None:
            key = jax.random.PRNGKey(0)
        q, _ = jnp.linalg.qr(jax.random.normal(key, (m, d)))
        return cls(matrix=rearrange(q, "m d -> d m"), in_size=m, out_size=d)

sample(key: jax.Array | None, d: int, m: int) -> OrthonormalSketch classmethod

Draw a (d, m) sketch with orthonormal rows.

Parameters:

Name Type Description Default
key Array | None

PRNG key. None means jax.random.PRNGKey(0).

required
d int

Sketch size (rows of \(S\)); must satisfy d <= m.

required
m int

Input size (columns of \(S\)).

required

Returns:

Type Description
OrthonormalSketch

The sampled OrthonormalSketch.

Raises:

Type Description
ValueError

If d > m.

Source code in src/gaussx/_sketching/_dense.py
@classmethod
def sample(cls, key: jax.Array | None, d: int, m: int) -> OrthonormalSketch:
    """Draw a ``(d, m)`` sketch with orthonormal rows.

    Args:
        key: PRNG key. ``None`` means ``jax.random.PRNGKey(0)``.
        d: Sketch size (rows of $S$); must satisfy ``d <= m``.
        m: Input size (columns of $S$).

    Returns:
        The sampled `OrthonormalSketch`.

    Raises:
        ValueError: If ``d > m``.
    """
    if d > m:
        raise ValueError(f"OrthonormalSketch needs d <= m, got d={d}, m={m}.")
    if key is None:
        key = jax.random.PRNGKey(0)
    q, _ = jnp.linalg.qr(jax.random.normal(key, (m, d)))
    return cls(matrix=rearrange(q, "m d -> d m"), in_size=m, out_size=d)

SparseSignSketch

Bases: AbstractSketch

Sparse sign sketch (sparse Johnson-Lindenstrauss transform).

Each column of \(S \in \mathbb{R}^{d \times m}\) has exactly nnz non-zeros, \(\pm 1/\sqrt{\text{nnz}}\) with independent random signs, at nnz distinct uniformly random rows. With \(O(\log n)\) non-zeros per column, \(d = O(n \log n / \varepsilon^2)\) suffices for an \(\varepsilon\)-subspace embedding of an \(n\)-dimensional subspace (Cohen, 2016), and applying \(S\) costs \(O(\text{nnz} \cdot m \cdot n)\): the default sketch for tall problems.

apply is a single segment_sum of signs * A[column] into rows; no sparse-matrix library is involved.

Attributes:

Name Type Description
rows Int[Array, 'nnz m']

Row index of each non-zero, shape (nnz, m); the nnz entries of each column are distinct.

signs Float[Array, 'nnz m']

Value of each non-zero, \(\pm 1/\sqrt{\text{nnz}}\), shape (nnz, m).

in_size int

\(m\).

out_size int

\(d\).

Examples:

Sketch a tall Jacobian (10⁶ residuals × 200 parameters), available only as a matrix-free J_op, down to 800 rows:

S = gx.SparseSignSketch.sample(key, d=800, m=1_000_000, nnz=8)
SJ = S.sketch_operator(J_op)  # (800, 200), matrix-free
sv = jnp.linalg.svd(SJ, compute_uv=False)  # J's singular values, within (1 ± ε)

A small runnable version:

>>> import jax.numpy as jnp, jax.random as jr
>>> import gaussx as gx
>>> S = gx.SparseSignSketch.sample(jr.key(0), d=40, m=1000, nnz=4)
>>> S.apply(jnp.ones((1000, 5))).shape
(40, 5)
Source code in src/gaussx/_sketching/_sparse_sign.py
class SparseSignSketch(AbstractSketch):
    r"""Sparse sign sketch (sparse Johnson-Lindenstrauss transform).

    Each column of $S \in \mathbb{R}^{d \times m}$ has exactly ``nnz``
    non-zeros, $\pm 1/\sqrt{\text{nnz}}$ with independent random signs, at
    ``nnz`` distinct uniformly random rows. With $O(\log n)$ non-zeros per
    column, $d = O(n \log n / \varepsilon^2)$ suffices for an
    $\varepsilon$-subspace embedding of an $n$-dimensional subspace (Cohen,
    2016), and applying $S$ costs $O(\text{nnz} \cdot m \cdot n)$: the default
    sketch for tall problems.

    `apply` is a single ``segment_sum`` of ``signs * A[column]`` into
    ``rows``; no sparse-matrix library is involved.

    Attributes:
        rows: Row index of each non-zero, shape ``(nnz, m)``; the ``nnz``
            entries of each column are distinct.
        signs: Value of each non-zero, $\pm 1/\sqrt{\text{nnz}}$, shape
            ``(nnz, m)``.
        in_size: $m$.
        out_size: $d$.

    Examples:
        Sketch a tall Jacobian (10⁶ residuals × 200 parameters), available
        only as a matrix-free ``J_op``, down to 800 rows:

        ```python
        S = gx.SparseSignSketch.sample(key, d=800, m=1_000_000, nnz=8)
        SJ = S.sketch_operator(J_op)  # (800, 200), matrix-free
        sv = jnp.linalg.svd(SJ, compute_uv=False)  # J's singular values, within (1 ± ε)
        ```

        A small runnable version:

        >>> import jax.numpy as jnp, jax.random as jr
        >>> import gaussx as gx
        >>> S = gx.SparseSignSketch.sample(jr.key(0), d=40, m=1000, nnz=4)
        >>> S.apply(jnp.ones((1000, 5))).shape
        (40, 5)
    """

    rows: Int[Array, "nnz m"]
    signs: Float[Array, "nnz m"]
    in_size: int = eqx.field(static=True)
    out_size: int = eqx.field(static=True)

    @classmethod
    def sample(
        cls, key: jax.Array | None, d: int, m: int, *, nnz: int = 8
    ) -> SparseSignSketch:
        r"""Draw a ``(d, m)`` sparse sign sketch.

        The ``nnz`` distinct rows of each column are drawn with Floyd's
        algorithm, vectorised over columns: $O(\text{nnz}^2 m)$ work, never
        $O(dm)$.

        Args:
            key: PRNG key. ``None`` means ``jax.random.PRNGKey(0)``.
            d: Sketch size (rows of $S$).
            m: Input size (columns of $S$).
            nnz: Non-zeros per column; clipped to ``d``.

        Returns:
            The sampled `SparseSignSketch`.
        """
        if key is None:
            key = jax.random.PRNGKey(0)
        nnz = min(nnz, d)
        row_key, sign_key = jax.random.split(key)
        # Floyd: for j = d - nnz, ..., d - 1 draw t ~ U{0..j}; keep t unless it
        # is already taken, in which case keep j (never taken yet).
        rows: list[Array] = []
        for k, step_key in enumerate(jax.random.split(row_key, nnz)):
            j = d - nnz + k
            t = jax.random.randint(step_key, (m,), 0, j + 1)
            taken = jnp.zeros((m,), dtype=bool)
            for r in rows:
                taken = taken | (r == t)
            rows.append(jnp.where(taken, j, t))
        signs = jax.random.rademacher(sign_key, (nnz, m), dtype=float) / jnp.sqrt(nnz)
        return cls(rows=jnp.stack(rows), signs=signs, in_size=m, out_size=d)

    def apply(self, A: Float[Array, "m ..."]) -> Float[Array, "d ..."]:
        contributions = einsum(self.signs.astype(A.dtype), A, "k m, m ... -> k m ...")
        return jax.ops.segment_sum(
            rearrange(contributions, "k m ... -> (k m) ..."),
            rearrange(self.rows, "k m -> (k m)"),
            num_segments=self.out_size,
        )

    def apply_transpose(self, Y: Float[Array, "d ..."]) -> Float[Array, "m ..."]:
        # Accumulate one gather per non-zero: O(m · ncols) memory, not O(nnz · m).
        signs = self.signs.astype(Y.dtype)
        out = einsum(signs[0], Y[self.rows[0]], "m, m ... -> m ...")
        for k in range(1, signs.shape[0]):
            out = out + einsum(signs[k], Y[self.rows[k]], "m, m ... -> m ...")
        return out

sample(key: jax.Array | None, d: int, m: int, *, nnz: int = 8) -> SparseSignSketch classmethod

Draw a (d, m) sparse sign sketch.

The nnz distinct rows of each column are drawn with Floyd's algorithm, vectorised over columns: \(O(\text{nnz}^2 m)\) work, never \(O(dm)\).

Parameters:

Name Type Description Default
key Array | None

PRNG key. None means jax.random.PRNGKey(0).

required
d int

Sketch size (rows of \(S\)).

required
m int

Input size (columns of \(S\)).

required
nnz int

Non-zeros per column; clipped to d.

8

Returns:

Type Description
SparseSignSketch

The sampled SparseSignSketch.

Source code in src/gaussx/_sketching/_sparse_sign.py
@classmethod
def sample(
    cls, key: jax.Array | None, d: int, m: int, *, nnz: int = 8
) -> SparseSignSketch:
    r"""Draw a ``(d, m)`` sparse sign sketch.

    The ``nnz`` distinct rows of each column are drawn with Floyd's
    algorithm, vectorised over columns: $O(\text{nnz}^2 m)$ work, never
    $O(dm)$.

    Args:
        key: PRNG key. ``None`` means ``jax.random.PRNGKey(0)``.
        d: Sketch size (rows of $S$).
        m: Input size (columns of $S$).
        nnz: Non-zeros per column; clipped to ``d``.

    Returns:
        The sampled `SparseSignSketch`.
    """
    if key is None:
        key = jax.random.PRNGKey(0)
    nnz = min(nnz, d)
    row_key, sign_key = jax.random.split(key)
    # Floyd: for j = d - nnz, ..., d - 1 draw t ~ U{0..j}; keep t unless it
    # is already taken, in which case keep j (never taken yet).
    rows: list[Array] = []
    for k, step_key in enumerate(jax.random.split(row_key, nnz)):
        j = d - nnz + k
        t = jax.random.randint(step_key, (m,), 0, j + 1)
        taken = jnp.zeros((m,), dtype=bool)
        for r in rows:
            taken = taken | (r == t)
        rows.append(jnp.where(taken, j, t))
    signs = jax.random.rademacher(sign_key, (nnz, m), dtype=float) / jnp.sqrt(nnz)
    return cls(rows=jnp.stack(rows), signs=signs, in_size=m, out_size=d)

SRHTSketch

Bases: AbstractSketch

Subsampled randomized Hadamard transform sketch.

\[ S = \sqrt{m_2 / d}\; R\, \tfrac{1}{\sqrt{m_2}} H\, D\, P , \]

where \(P\) is a random permutation of the \(m\) input rows, \(D\) a diagonal of random signs, the result zero-padded to \(m_2 = 2^{\lceil \log_2 m \rceil}\) rows, \(H/\sqrt{m_2}\) the orthonormal Walsh-Hadamard transform (hadamard_transform) and \(R\) selects \(d\) distinct rows uniformly at random. \(HD\) spreads each vector's mass evenly over the coordinates (flattens the leverage), so uniform row sampling afterwards is safe; \(\mathbb{E}[S^\top S] = I_m\). A size \(d = O((n + \log m)\log n / \varepsilon^2)\) suffices for an \(\varepsilon\)-subspace embedding (Tropp, 2011), and applying \(S\) costs \(O(m n \log m)\).

The padding to a power of two costs up to 2× the memory of the input while the transform runs.

Attributes:

Name Type Description
permutation Int[Array, ' m']

Input row read by each position, \((Px)_i = x_{\text{permutation}_i}\), shape (m,).

signs Float[Array, ' m']

Diagonal of \(D\), \(\pm 1\), shape (m,).

rows Int[Array, ' d']

The \(d\) distinct rows of the padded transform kept by \(R\), shape (d,).

in_size int

\(m\).

out_size int

\(d\).

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> import gaussx as gx
>>> S = gx.SRHTSketch.sample(jr.key(0), d=16, m=100)
>>> S.apply(jnp.ones((100, 3))).shape
(16, 3)
Source code in src/gaussx/_sketching/_srht.py
class SRHTSketch(AbstractSketch):
    r"""Subsampled randomized Hadamard transform sketch.

    $$
    S = \sqrt{m_2 / d}\; R\, \tfrac{1}{\sqrt{m_2}} H\, D\, P ,
    $$

    where $P$ is a random permutation of the $m$ input rows, $D$ a diagonal
    of random signs, the result zero-padded to $m_2 = 2^{\lceil \log_2 m
    \rceil}$ rows, $H/\sqrt{m_2}$ the orthonormal Walsh-Hadamard transform
    (`hadamard_transform`) and $R$ selects $d$ distinct rows uniformly at
    random. $HD$ spreads each vector's mass evenly over the coordinates
    (flattens the leverage), so uniform row sampling afterwards is safe;
    $\mathbb{E}[S^\top S] = I_m$. A size $d = O((n + \log m)\log n /
    \varepsilon^2)$ suffices for an $\varepsilon$-subspace embedding (Tropp,
    2011), and applying $S$ costs $O(m n \log m)$.

    The padding to a power of two costs up to 2× the memory of the input
    while the transform runs.

    Attributes:
        permutation: Input row read by each position, $(Px)_i =
            x_{\text{permutation}_i}$, shape ``(m,)``.
        signs: Diagonal of $D$, $\pm 1$, shape ``(m,)``.
        rows: The $d$ distinct rows of the padded transform kept by $R$,
            shape ``(d,)``.
        in_size: $m$.
        out_size: $d$.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> import gaussx as gx
        >>> S = gx.SRHTSketch.sample(jr.key(0), d=16, m=100)
        >>> S.apply(jnp.ones((100, 3))).shape
        (16, 3)
    """

    permutation: Int[Array, " m"]
    signs: Float[Array, " m"]
    rows: Int[Array, " d"]
    in_size: int = eqx.field(static=True)
    out_size: int = eqx.field(static=True)

    @classmethod
    def sample(cls, key: jax.Array | None, d: int, m: int) -> SRHTSketch:
        """Draw a ``(d, m)`` SRHT sketch.

        Args:
            key: PRNG key. ``None`` means ``jax.random.PRNGKey(0)``.
            d: Sketch size (rows of $S$); at most the padded size $m_2$.
            m: Input size (columns of $S$).

        Returns:
            The sampled `SRHTSketch`.

        Raises:
            ValueError: If ``d`` exceeds the padded size $m_2$.
        """
        m2 = _padded_size(m)
        if d > m2:
            raise ValueError(f"SRHTSketch needs d <= {m2} (m={m} padded), got {d}.")
        if key is None:
            key = jax.random.PRNGKey(0)
        perm_key, sign_key, row_key = jax.random.split(key, 3)
        return cls(
            permutation=jax.random.permutation(perm_key, m),
            signs=jax.random.rademacher(sign_key, (m,), dtype=float),
            rows=jax.random.choice(row_key, m2, (d,), replace=False),
            in_size=m,
            out_size=d,
        )

    # Jitted: eagerly, each of the log2(m2) butterfly passes dispatches its own
    # einx rearrangements, which is far slower than one compiled transform.
    @eqx.filter_jit
    def apply(self, A: Float[Array, "m ..."]) -> Float[Array, "d ..."]:
        x = einsum(self.signs.astype(A.dtype), A[self.permutation], "m, m ... -> m ...")
        pad = jnp.zeros(
            (_padded_size(self.in_size) - self.in_size, *A.shape[1:]), dtype=A.dtype
        )
        x = jnp.concatenate([x, pad])
        y = hadamard_transform(rearrange(x, "m ... -> ... m"))[..., self.rows]
        # √(m₂/d) · H/√m₂ = H/√d with the unnormalised transform.
        return rearrange(y, "... d -> d ...") / math.sqrt(self.out_size)

    @eqx.filter_jit
    def apply_transpose(self, Y: Float[Array, "d ..."]) -> Float[Array, "m ..."]:
        m2 = _padded_size(self.in_size)
        z = jnp.zeros((m2, *Y.shape[1:]), dtype=Y.dtype).at[self.rows].set(Y)
        x = rearrange(
            hadamard_transform(rearrange(z, "m ... -> ... m")), "... m -> m ..."
        )
        x = einsum(self.signs.astype(Y.dtype), x[: self.in_size], "m, m ... -> m ...")
        out = jnp.zeros_like(x).at[self.permutation].set(x)
        return out / math.sqrt(self.out_size)

sample(key: jax.Array | None, d: int, m: int) -> SRHTSketch classmethod

Draw a (d, m) SRHT sketch.

Parameters:

Name Type Description Default
key Array | None

PRNG key. None means jax.random.PRNGKey(0).

required
d int

Sketch size (rows of \(S\)); at most the padded size \(m_2\).

required
m int

Input size (columns of \(S\)).

required

Returns:

Type Description
SRHTSketch

The sampled SRHTSketch.

Raises:

Type Description
ValueError

If d exceeds the padded size \(m_2\).

Source code in src/gaussx/_sketching/_srht.py
@classmethod
def sample(cls, key: jax.Array | None, d: int, m: int) -> SRHTSketch:
    """Draw a ``(d, m)`` SRHT sketch.

    Args:
        key: PRNG key. ``None`` means ``jax.random.PRNGKey(0)``.
        d: Sketch size (rows of $S$); at most the padded size $m_2$.
        m: Input size (columns of $S$).

    Returns:
        The sampled `SRHTSketch`.

    Raises:
        ValueError: If ``d`` exceeds the padded size $m_2$.
    """
    m2 = _padded_size(m)
    if d > m2:
        raise ValueError(f"SRHTSketch needs d <= {m2} (m={m} padded), got {d}.")
    if key is None:
        key = jax.random.PRNGKey(0)
    perm_key, sign_key, row_key = jax.random.split(key, 3)
    return cls(
        permutation=jax.random.permutation(perm_key, m),
        signs=jax.random.rademacher(sign_key, (m,), dtype=float),
        rows=jax.random.choice(row_key, m2, (d,), replace=False),
        in_size=m,
        out_size=d,
    )

RowSamplingSketch

Bases: AbstractSketch

Weighted row-sampling sketch.

Row \(k\) of \(S\) is \(e_{i_k}^\top / \sqrt{d\, p_{i_k}}\) with \(i_k \sim p\) drawn i.i.d. (with replacement), so \(\mathbb{E}[S^\top S] = I_m\). With leverage-score probabilities this is a subspace embedding; with uniform probabilities it is only safe once the leverage has been flattened (as inside SRHTSketch). Applying \(S\) is a gather: \(O(dn)\).

Attributes:

Name Type Description
rows Int[Array, ' d']

Sampled row indices \(i_k\), shape (d,).

weights Float[Array, ' d']

Row weights \(1/\sqrt{d\, p_{i_k}}\), shape (d,).

in_size int

\(m\).

out_size int

\(d\).

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> import gaussx as gx
>>> S = gx.RowSamplingSketch.sample(jr.key(0), d=10, m=100)
>>> S.apply(jnp.ones((100, 2))).shape
(10, 2)
Source code in src/gaussx/_sketching/_sampling.py
class RowSamplingSketch(AbstractSketch):
    r"""Weighted row-sampling sketch.

    Row $k$ of $S$ is $e_{i_k}^\top / \sqrt{d\, p_{i_k}}$ with
    $i_k \sim p$ drawn i.i.d. (with replacement), so
    $\mathbb{E}[S^\top S] = I_m$. With leverage-score probabilities this is
    a subspace embedding; with uniform probabilities it is only safe once the
    leverage has been flattened (as inside `SRHTSketch`). Applying $S$ is a
    gather: $O(dn)$.

    Attributes:
        rows: Sampled row indices $i_k$, shape ``(d,)``.
        weights: Row weights $1/\sqrt{d\, p_{i_k}}$, shape ``(d,)``.
        in_size: $m$.
        out_size: $d$.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> import gaussx as gx
        >>> S = gx.RowSamplingSketch.sample(jr.key(0), d=10, m=100)
        >>> S.apply(jnp.ones((100, 2))).shape
        (10, 2)
    """

    rows: Int[Array, " d"]
    weights: Float[Array, " d"]
    in_size: int = eqx.field(static=True)
    out_size: int = eqx.field(static=True)

    @classmethod
    def sample(
        cls,
        key: jax.Array | None,
        d: int,
        m: int,
        *,
        probabilities: Float[ArrayLike, " m"] | None = None,
    ) -> RowSamplingSketch:
        """Draw a ``(d, m)`` row-sampling sketch.

        Args:
            key: PRNG key. ``None`` means ``jax.random.PRNGKey(0)``.
            d: Number of sampled rows.
            m: Input size (columns of $S$).
            probabilities: Non-negative sampling weights $p$, shape ``(m,)``,
                normalised internally. ``None`` means uniform.

        Returns:
            The sampled `RowSamplingSketch`.
        """
        if key is None:
            key = jax.random.PRNGKey(0)
        if probabilities is None:
            p = jnp.full((m,), 1.0 / m)
        else:
            p = jnp.asarray(probabilities)
            p = p / jnp.sum(p)
        rows = jax.random.choice(key, m, (d,), replace=True, p=p)
        return cls(
            rows=rows, weights=1.0 / jnp.sqrt(d * p[rows]), in_size=m, out_size=d
        )

    def apply(self, A: Float[Array, "m ..."]) -> Float[Array, "d ..."]:
        return einsum(self.weights.astype(A.dtype), A[self.rows], "d, d ... -> d ...")

    def apply_transpose(self, Y: Float[Array, "d ..."]) -> Float[Array, "m ..."]:
        weighted = einsum(self.weights.astype(Y.dtype), Y, "d, d ... -> d ...")
        return (
            jnp.zeros((self.in_size, *Y.shape[1:]), dtype=Y.dtype)
            .at[self.rows]
            .add(weighted)
        )

sample(key: jax.Array | None, d: int, m: int, *, probabilities: Float[ArrayLike, ' m'] | None = None) -> RowSamplingSketch classmethod

Draw a (d, m) row-sampling sketch.

Parameters:

Name Type Description Default
key Array | None

PRNG key. None means jax.random.PRNGKey(0).

required
d int

Number of sampled rows.

required
m int

Input size (columns of \(S\)).

required
probabilities Float[ArrayLike, ' m'] | None

Non-negative sampling weights \(p\), shape (m,), normalised internally. None means uniform.

None

Returns:

Type Description
RowSamplingSketch

The sampled RowSamplingSketch.

Source code in src/gaussx/_sketching/_sampling.py
@classmethod
def sample(
    cls,
    key: jax.Array | None,
    d: int,
    m: int,
    *,
    probabilities: Float[ArrayLike, " m"] | None = None,
) -> RowSamplingSketch:
    """Draw a ``(d, m)`` row-sampling sketch.

    Args:
        key: PRNG key. ``None`` means ``jax.random.PRNGKey(0)``.
        d: Number of sampled rows.
        m: Input size (columns of $S$).
        probabilities: Non-negative sampling weights $p$, shape ``(m,)``,
            normalised internally. ``None`` means uniform.

    Returns:
        The sampled `RowSamplingSketch`.
    """
    if key is None:
        key = jax.random.PRNGKey(0)
    if probabilities is None:
        p = jnp.full((m,), 1.0 / m)
    else:
        p = jnp.asarray(probabilities)
        p = p / jnp.sum(p)
    rows = jax.random.choice(key, m, (d,), replace=True, p=p)
    return cls(
        rows=rows, weights=1.0 / jnp.sqrt(d * p[rows]), in_size=m, out_size=d
    )

Fast transforms

The unnormalised fast Walsh–Hadamard transform behind SRHTSketch (and kernellib's FastFood features).

Structured linear algebra and Gaussian primitives for JAX.

hadamard_transform(x: Float[Array, '... d']) -> Float[Array, '... d']

Unnormalized Walsh-Hadamard transform along the last axis.

Computes \(H_d x\) for the Sylvester-ordered Hadamard matrix \(H_{2m} = \begin{pmatrix} H_m & H_m \\ H_m & -H_m \end{pmatrix}\), \(H_1 = 1\), with \(\log_2 d\) butterfly passes: \(O(d \log d)\) work, no \(d \times d\) matrix. Applying it twice returns \(d\,x\).

Parameters:

Name Type Description Default
x Float[Array, '... d']

Array whose last axis has a power-of-two length d. Leading axes are batch axes.

required

Returns:

Type Description
Float[Array, '... d']

\(H_d x\), same shape as x.

Raises:

Type Description
ValueError

If the last axis is not a power of two.

Examples:

>>> import jax.numpy as jnp
>>> from gaussx import hadamard_transform
>>> hadamard_transform(jnp.array([1.0, 0.0, 0.0, 0.0])).tolist()
[1.0, 1.0, 1.0, 1.0]
>>> hadamard_transform(jnp.array([1.0, 2.0])).tolist()
[3.0, -1.0]
Source code in src/gaussx/_sketching/_hadamard.py
def hadamard_transform(x: Float[Array, "... d"]) -> Float[Array, "... d"]:
    r"""Unnormalized Walsh-Hadamard transform along the last axis.

    Computes $H_d x$ for the Sylvester-ordered Hadamard matrix
    $H_{2m} = \begin{pmatrix} H_m & H_m \\ H_m & -H_m \end{pmatrix}$,
    $H_1 = 1$, with $\log_2 d$ butterfly passes: $O(d \log d)$ work, no
    $d \times d$ matrix. Applying it twice returns $d\,x$.

    Args:
        x: Array whose last axis has a power-of-two length ``d``. Leading axes
            are batch axes.

    Returns:
        $H_d x$, same shape as ``x``.

    Raises:
        ValueError: If the last axis is not a power of two.

    Examples:
        >>> import jax.numpy as jnp
        >>> from gaussx import hadamard_transform
        >>> hadamard_transform(jnp.array([1.0, 0.0, 0.0, 0.0])).tolist()
        [1.0, 1.0, 1.0, 1.0]
        >>> hadamard_transform(jnp.array([1.0, 2.0])).tolist()
        [3.0, -1.0]
    """
    d = x.shape[-1]
    if not _is_power_of_two(d):
        raise ValueError(f"hadamard_transform needs a power-of-two last axis, got {d}.")
    h = 1
    while h < d:
        # Pair entry i with entry i + h inside each block of 2h.
        y = rearrange(x, "... (m two h) -> ... m two h", two=2, h=h)
        a, b = y[..., 0, :], y[..., 1, :]
        x = rearrange(
            jnp.stack([a + b, a - b], axis=-2), "... m two h -> ... (m two h)"
        )
        h *= 2
    return x