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
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
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
apply(A: Float[Array, 'm ...']) -> Float[Array, 'd ...']
abstractmethod
¶
apply_transpose(Y: Float[Array, 'd ...']) -> Float[Array, 'm ...']
abstractmethod
¶
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 |
required |
Returns:
| Type | Description |
|---|---|
Float[Array, 'd n']
|
The dense sketch \(SA\), shape |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in src/gaussx/_sketching/_base.py
as_operator() -> lx.AbstractLinearOperator
¶
Return \(S\) as a matrix-free (d, m) lineax operator.
Source code in src/gaussx/_sketching/_base.py
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 |
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
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. |
required |
d
|
int
|
Sketch size (rows of \(S\)). |
required |
m
|
int
|
Input size (columns of \(S\)). |
required |
Returns:
| Type | Description |
|---|---|
GaussianSketch
|
The sampled |
Source code in src/gaussx/_sketching/_dense.py
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 |
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
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. |
required |
d
|
int
|
Sketch size (rows of \(S\)); must satisfy |
required |
m
|
int
|
Input size (columns of \(S\)). |
required |
Returns:
| Type | Description |
|---|---|
OrthonormalSketch
|
The sampled |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in src/gaussx/_sketching/_dense.py
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 |
signs |
Float[Array, 'nnz m']
|
Value of each non-zero, \(\pm 1/\sqrt{\text{nnz}}\), shape
|
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
14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 | |
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. |
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 |
8
|
Returns:
| Type | Description |
|---|---|
SparseSignSketch
|
The sampled |
Source code in src/gaussx/_sketching/_sparse_sign.py
SRHTSketch
¶
Bases: AbstractSketch
Subsampled randomized Hadamard transform sketch.
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 |
signs |
Float[Array, ' m']
|
Diagonal of \(D\), \(\pm 1\), shape |
rows |
Int[Array, ' d']
|
The \(d\) distinct rows of the padded transform kept by \(R\),
shape |
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
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. |
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 |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in src/gaussx/_sketching/_srht.py
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 |
weights |
Float[Array, ' d']
|
Row weights \(1/\sqrt{d\, p_{i_k}}\), shape |
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
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. |
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 |
None
|
Returns:
| Type | Description |
|---|---|
RowSamplingSketch
|
The sampled |
Source code in src/gaussx/_sketching/_sampling.py
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 |
required |
Returns:
| Type | Description |
|---|---|
Float[Array, '... d']
|
\(H_d x\), same shape as |
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]