Sparse-block compensation¶
Research notes on training-free methods that recover the signal lost by
block-sparse attention in video DiTs, and how mlx-arsenal exposes them.
This page documents a pattern and its fit with the library; it does not
open an ADR. It extends the sparse-attention work of
ADR-0001.
The problem¶
Block-sparse attention (STA, SVG2, Radial…) groups queries and keys into clusters and computes only some (query cluster, key cluster) blocks. The skipped blocks are usually hard-dropped: the softmax is renormalized over the kept keys, and the probability mass of the dropped keys is lost. Two 2026 methods recover part of it without training.
SVG-EAR — centroid compensation¶
SVG-EAR: Parameter-Free Linear Compensation for Sparse Video Generation
via Error-aware Routing (arXiv 2603.08982, ECCV 2026; code Apache-2.0 in
svg-project/Sparse-VideoGen).
Every skipped key k_j is replaced by its cluster centroid k̄_c, and its
value by v̄_c, inside one softmax. Keys of a cluster then share one
logit, so the skipped part collapses to one term per cluster with a
log-count prior:
The output stays an exact softmax average over a modified key set. The paper merges the sparse kernel's log-sum-exp with the centroid terms in a custom Triton kernel, and adds an error-aware router choosing which blocks to compute exactly.
SparsePR — probe-fitted residual repair¶
Partition the Support, Reconstruct the Residual: Training-Free Sparse
Attention for Video Generation and World Models (arXiv 2608.18484; code
Apache-2.0 at PardisTaghavi/SparsePR).
The sparse output is hard-dropped. A few probe rows (64 by default)
get exact dense attention; the residual R = O_dense − O_sparse on those
rows fits a weighted ridge regression from the standardized sparse output,
projected on the top-16 principal residual directions, which then predicts
the residual on every other row. The fit happens on every attention call,
per (batch, head); nothing is calibrated offline. The paper's ablation
attributes most of its quality gain to this repair step.
What mlx-arsenal ships¶
mlx_arsenal.attention:
| Function | Role |
|---|---|
centroid_compensated_attention |
SVG-EAR compensation for caller-given clusters and a cluster-level 0 / -inf block mask |
probe_residual_correction |
SparsePR repair of any approximate attention output |
select_probe_rows |
Head-agnostic probe choice, round-robin over query clusters, with SparsePR's \|G_a\| / m_a weights |
tile_labels |
Cluster labels matching Sliding Tile Attention tiles |
No LSE needed¶
mx.fast.scaled_dot_product_attention does not return the log-sum-exp,
which the paper's merge relies on. The identity above removes the need:
the compensated output equals one SDPA call over the extended keys
[K; K̄], values [V; V̄], with an additive (Sq, Sk + Ck) mask that is
0 / -inf on token columns (kept / skipped block) and log n_c / -inf on
centroid columns (skipped / kept block).
Dense by design¶
MLX has no block-sparse attention kernel, and ADR-0001 keeps writing one
out of scope. Everything here runs dense — the compensated call costs
slightly more than dense attention. With (S,) labels and a
(Cq, Ck) block mask, the (Sq, Sk + Ck) mask is built once and
broadcast over batch and heads; per-head labels or masks build it per
head, which at video sizes (S ≈ 32k, 40 heads) is out of reach. These functions are quality tools:
measure what a sparse pattern costs in a port, with and without
compensation, and serve as the numerical reference for a future kernel.
Recipe: STA with both repairs¶
import mlx.core as mx
from mlx_arsenal.attention import (
centroid_compensated_attention, probe_residual_correction,
select_probe_rows, sliding_tile_block_mask, tile_labels,
)
tile = (tt, th, tw)
labels = tile_labels(T, H, W, tile=tile)
block_mask = sliding_tile_block_mask(T // tt, H // th, W // tw, tile=(1, 1, 1), window=window)
out = centroid_compensated_attention(q, k, v, q_labels=labels, k_labels=labels, block_mask=block_mask)
probe_idx, weights = select_probe_rows(labels, 64)
o_probe = mx.fast.scaled_dot_product_attention(q[:, :, probe_idx], k, v, scale=q.shape[-1] ** -0.5)
out = probe_residual_correction(out, o_probe, probe_idx, weights=weights)
Evaluating sliding_tile_block_mask at tile resolution gives exactly the
cluster-level mask: expanding it through tile_labels reproduces the
token-level STA mask (pinned by a test).
Deviations from the papers¶
- Clusters are inputs. Both papers cluster Q and K with k-means (SVG2 semantic permutation, or SparsePR's response-coupled embedding) and warm start it across denoising steps. That state belongs to the caller's loop (see the doctrine: the caller orchestrates); pass any integer labels.
- Probe selection takes the middle-out member of each query cluster
instead of the row nearest to the cluster centroid, so one
probe_idxserves all heads without needingq. - Probe weights are normalized to sum to 1, so
ridgedoes not scale with the probe count. The reference code's blend factor and norm cap (disabled by default there) are not exposed. - SparsePR on top of centroid compensation is allowed; the paper applies its repair to a hard-dropped output.
Out of scope, and why¶
- Error-aware routing (SVG-EAR's block selector): a policy on top of the primitive, cheap for a caller to add once a port needs it.
- k-means / cluster permutation: stateful across steps, see above.
- GQA: per-head labels are ambiguous with shared KV heads.
- A Metal block-sparse kernel: a project of its own (ADR-0001).
Reported results (from the papers)¶
| Method | Model | Density | Speedup | PSNR / LPIPS |
|---|---|---|---|---|
| SVG-EAR | HunyuanVideo-13B | 22.2% | 1.93× | 31.04 / 0.092 |
| SVG-EAR | Wan2.2-14B T2V 720p | 26.0% | 1.59× | 25.00 / 0.153 |
| SparsePR | HunyuanVideo | 21.9% | 2.61× | 31.84 / 0.087 |
| SparsePR | Wan2.2-I2V-A14B | 22.0% | 1.80× | 30.66 / 0.044 |
Speedups are CUDA figures with custom kernels; they do not carry over to this dense implementation.