Attention¶
attention
¶
Kind
¶
Bases: Enum
Discrete head-pattern label.
centroid_compensated_attention
¶
centroid_compensated_attention(q: array, k: array, v: array, *, q_labels: array, k_labels: array, block_mask: array, scale: float | None = None) -> array
Block-sparse attention with skipped blocks compensated by key centroids (SVG-EAR).
Queries and keys are grouped into clusters by integer labels. For every
(query cluster, key cluster) block kept by block_mask, attention is
exact. For a skipped block, each key and value is replaced by the mean
key and mean value of its cluster. The result is exactly dense attention
over that modified key set, so every output row is still a convex
softmax average — skipped blocks are approximated, not dropped.
Keys of one cluster share a single centroid logit, so the skipped part
collapses to one extra key per cluster carrying a log n_c bias
(n_c = cluster size). The implementation therefore runs one
mx.fast.scaled_dot_product_attention over [K; K̄] / [V; V̄] with an
additive (Sq, Sk + Ck) mask. It is dense — no speedup — and meant as a
quality tool and a reference for a future block-sparse kernel. The mask
is built once when labels are (S,) and block_mask has no per-batch or
per-head dims; per-head labels or masks materialize it B·H times.
With block_mask all 0 this is exact dense attention; with it all
-inf, attention over the cluster centroids only.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
v
|
array
|
|
required |
q_labels
|
array
|
|
required |
k_labels
|
array
|
|
required |
block_mask
|
array
|
|
required |
scale
|
float | None
|
Logit scale. Defaults to |
None
|
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
|
array
|
when they share it, their promoted type otherwise (e.g. bfloat16 |
array
|
with float16 |
Note
The mask is built in q.dtype (SDPA rejects a mask wider than its
output). In bfloat16 the log n_c bias is therefore rounded to half
a bfloat16 ulp: at most 1/64 in logit units for 55 ≤ n_c < 2981
(4 ≤ log n_c < 8), i.e. up to ~1.6% on that centroid's weight — the
same order as the rounding of the bfloat16 q·k logits themselves.
Pass float32 inputs when this tool is used as a numerical reference.
probe_residual_correction
¶
probe_residual_correction(o_sparse: array, o_probe: array, probe_idx: array, *, rank: int = 16, ridge: float = 0.1, weights: array | None = None) -> array
Repair an approximate attention output from a few exact probe rows (SparsePR).
The residual R = O_dense - O_sparse is measured exactly on the probe
rows, then predicted on every row by a weighted ridge regression on the
standardized sparse output, projected onto the top-rank principal
directions of the probe residuals:
B = (X̄ᵀ W X̄ + ridge·I)⁻¹ X̄ᵀ W R̄ X̄, R̄: standardized / centered probe rows
Ψ = top-`rank` right singular vectors of W^½ R̄
R̂ = μ_R + ((O_sparse - μ_X) / σ_X) B Ψ Ψᵀ
out = O_sparse + R̂ probe rows replaced by `o_probe`
Works with any approximate attention (hard-dropped block-sparse, masked, centroid-compensated). The correction is additive: rows are not renormalized, so the output is no longer a convex softmax average. The fit is stateless — SparsePR refits on every call — and runs per (batch, head) in float32; the solve and SVD use the CPU stream (MLX has no GPU kernels for them). Get the exact probe rows with one dense call:
Deviations from the paper: weights are normalized to sum to 1 (so
ridge does not scale with the number of probes), and the optional
blend factor and norm cap of the reference code (off by default there)
are not exposed.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
o_sparse
|
array
|
|
required |
o_probe
|
array
|
|
required |
probe_idx
|
array
|
|
required |
rank
|
int
|
Rank of the residual projection, in |
16
|
ridge
|
float
|
Ridge strength, |
0.1
|
weights
|
array | None
|
|
None
|
Returns:
| Type | Description |
|---|---|
array
|
|
select_probe_rows
¶
Pick probe query rows for :func:probe_residual_correction, spread across clusters.
Round-robin over the non-empty clusters in increasing label order: pass
r takes, from every cluster that still has unused members, its member
at middle-out rank r (positions sorted, middle first, then alternating
outward), until num_probes rows are taken. When the last pass cannot
serve every cluster, it takes clusters evenly spaced over the label
range, so fewer probes than clusters still span the whole sequence.
Deviation from SparsePR, which takes the row nearest to each query-group
centroid: that needs q and per-head groups. This selection depends only
on the labels, so a single probe_idx serves every head. Pass your own
probe_idx to :func:probe_residual_correction for another policy.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q_labels
|
array
|
|
required |
num_probes
|
int
|
Number of rows to pick, in |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
float32 weights |
tuple[array, array]
|
that cluster). When every cluster gets at least one probe, the |
tuple[array, array]
|
weights sum to |
tuple[array, array]
|
clusters, only the sampled clusters are weighted: the weights sum to |
tuple[array, array]
|
the total size of those clusters, and unsampled clusters have no |
tuple[array, array]
|
say in the fit. |
tile_labels
¶
Cluster label of every token for a (tt, th, tw) tiling of a video grid.
The label of token (t, h, w) is the row-major index of its tile in the
(T // tt, H // th, W // tw) tile grid; tokens are in T-major order.
Labels line up with :func:~mlx_arsenal.attention.sliding_tile_block_mask
evaluated at tile resolution, which yields the matching cluster-level
block mask:
labels = tile_labels(T, H, W, tile=(tt, th, tw))
block_mask = sliding_tile_block_mask(T // tt, H // th, W // tw, tile=(1, 1, 1), window=window)
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. Must be divisible by |
required |
H
|
int
|
Latent height. Must be divisible by |
required |
W
|
int
|
Latent width. Must be divisible by |
required |
tile
|
tuple[int, int, int]
|
|
required |
Returns:
| Type | Description |
|---|---|
array
|
|
antidiagonal_block_scores
¶
antidiagonal_block_scores(q: array, k: array, *, block_size: int = 128, stride: int = 16, scale: float | None = None) -> array
Estimate the attention mass of every (query block, key block) pair.
Tokens are grouped into contiguous blocks of block_size. Each
stride × stride sub-block of QKᵀ is summarized by the sum along
its anti-diagonal, Σ_r q_{iS+S-1-r} · k_{jS+r}, which touches every
query and key of the sub-block once. These strided logits (scaled by
scale / stride) are softmaxed per row in float32, summed over each
(block_size/stride)² tile, and normalized so every query-block row
sums to 1. XAttention keeps the raw tile sums; the normalization does not
change :func:top_p_block_mask, which is relative to the row total.
The strided matmul costs about 1/stride of a full QKᵀ in FLOPs;
the logits and softmax are 1/stride² of the full attention matrix.
All B·H heads are scored at once: transient memory is roughly
2 · B·H · (Nq/stride)·(Nk/stride) · 4 bytes (logits + softmax) —
≈16 MB per head at N = 32k, stride = 16, but ≈14 GB for a Wan
720p layer (N ≈ 75.8k, H = 40, CFG B = 2). For large models,
call it per head slice (q[:, h0:h1]) and concatenate.
Sequence lengths must be multiples of block_size: pad beforehand.
Zero-padded keys still take softmax mass (as in XAttention's non-causal
mode); mask their blocks out afterwards if that matters. Permute tokens
first (e.g. :func:~mlx_arsenal.attention.block_contiguous_permutation)
if blocks should follow another order.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
block_size
|
int
|
Block size in tokens, a multiple of |
128
|
stride
|
int
|
Anti-diagonal sampling stride, |
16
|
scale
|
float | None
|
Logit scale. Defaults to |
None
|
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
each row summing to 1. |
block_self_similarity
¶
Mean pairwise cosine similarity of the tokens inside each contiguous block.
SPADE's query-cohesion summary (arXiv 2608.03335, "SICS"): for each block
of block_size consecutive tokens, the mean cosine similarity over all
pairs of distinct tokens, ignoring zero-norm (padding) tokens. Computed
in O(block_size · D) as (‖Σ x̂‖² − Σ‖x̂‖²) / (n(n−1)) with x̂ the
normalized valid tokens and n their count; blocks with fewer than two
valid tokens score 0. High values mean the block's queries point the same
way, so a per-block summary (e.g. :func:minmax_block_scores) represents
them well.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
array
|
|
required |
block_size
|
int
|
Tokens per block; |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
minmax_block_scores
¶
Cheap per-block attention estimate from element-wise min/max summaries.
SPADE's estimator ("DSA"): each block of block_size consecutive
tokens is summarized by its element-wise max and min over tokens, and the
(query block, key block) score is
max((q_max + q_min) · k_maxᵀ, (q_max + q_min) · k_minᵀ) — a
(Nq/bs) × (Nk/bs) matmul pair instead of Nq × Nk. Scores are
unscaled logit estimates (as in the reference): rank them with
:func:top_k_block_mask, or softmax each row before
:func:top_p_block_mask. Reorder tokens first (e.g. by the tiling from
:func:select_tiling) so blocks are meaningful.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
block_size
|
int
|
Tokens per block; |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
radius_bounded_block_scores
¶
radius_bounded_block_scores(q: array, k: array, *, block_size: int, low: float = 0.5, high: float = 0.9) -> tuple[array, array]
RBS-Attention block scores: centroid and radius-bounded ("rescue").
A block centroid c_b can hide one strongly matching key among keys that
cancel it ("mean dilution"). RBS (arXiv 2609.20971) adds the key block's
radius r_b = max ‖k − c_b‖₂: by Cauchy-Schwarz
qᵀk ≤ qᵀc_b + ‖q‖·r_b for every key of the block. Per query token::
ℓ_base = qᵀc_b / √D
ℓ_rescue = (qᵀc_b + ‖q‖·r_b·β_b) / √D
with β_b = clip((r_b − r_low) / (r_high − r_low), 0, 1) from the
low / high quantiles of the radii of each (batch, head), so only
unusually spread blocks get the bound. Each (query block, key block) score
is the logsumexp of the logits over the query block's tokens — the log of
the paper's Σ exp(ℓ). Select with :func:relative_block_mask per
branch and take the union (mx.maximum of the additive masks), plus any
forced blocks (sinks, diagonal band).
β_b is 1 above r_high and 0 below r_low; when the two
quantiles coincide it is 1 for blocks strictly wider than r_low.
Unlike the other estimators it scores every query token: an
Nq × Nk / block_size matmul, cheaper than QKᵀ by block_size
in FLOPs, but it holds about three (B, H, Nq, Nk / block_size)
float32 tensors at once — ≈1.1 GB per head for a Wan 720p layer
(N ≈ 75.8k, block_size = 64). Call it per head (or per group of
heads) at video scale.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
block_size
|
int
|
Tokens per block; |
required |
low
|
float
|
Radius quantile below which |
0.5
|
high
|
float
|
Radius quantile above which |
0.9
|
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
float32 log scores. Use them with :func: |
tuple[array, array]
|
func: |
tuple[array, array]
|
masses, i.e. |
relative_block_mask
¶
Keep key blocks whose score is at least alpha times the row maximum.
RBS-Attention's selection rule, on log scores (as returned by
:func:radius_bounded_block_scores): keep b when
log S_b ≥ max_b' log S_b' + log α. The paper applies it to each branch
separately (α = 0.22 base, 0.18 rescue, tuned offline on LLM
prompts) and unions the masks with mx.maximum.
Scores must be in the log domain (RBS, or logit-like estimates such as
:func:minmax_block_scores); take mx.log of mass estimates such as
:func:antidiagonal_block_scores first. A row whose scores are all equal
(including all -inf) keeps every block; a NaN row keeps none.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
log_scores
|
array
|
|
required |
alpha
|
float
|
Relative threshold in |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
select_tiling
¶
Pick, per head, the 3D tiling under which the queries are most cohesive.
SPADE's input-adaptive blocking: for each candidate (tt, th, tw) tile,
tokens (T-major, N = T·H·W) are reordered so every tile is contiguous
and the tiles' :func:block_self_similarity is averaged. Each (batch,
head) takes the candidate with the highest mean cohesion (the first one
on ties). To use a choice, reorder q/k/v with
mx.argsort(tile_labels(T, H, W, tile=tiles[c])) along the token axis
and build block masks with block_size = tt·th·tw.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
grid
|
tuple[int, int, int]
|
|
required |
tiles
|
Sequence[tuple[int, int, int]]
|
Candidate tile shapes, each dividing |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
top_k_block_mask
¶
Keep the k highest-scoring key blocks of every query-block row.
SPADE's selection rule (a fixed block budget, e.g. 17 % of the key
blocks on Wan 2.1). Ties keep the lower block index; k larger than
the number of key blocks keeps them all.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
scores
|
array
|
|
required |
k
|
int
|
Blocks kept per row, |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
top_p_block_mask
¶
Keep, per query block, the key blocks covering a fraction τ of the mass.
Key blocks are ranked by score (descending, stable: lower index first on
ties) and a block is kept while the mass ranked before it is below
τ · row_total — so the block crossing the threshold is kept, and every
row keeps at least one block (also when a row is all zeros). This is
XAttention's non-causal selection rule, except that XAttention's scatter
also re-marks key block 0 in every row as a side effect; this function
keeps only the top-p set.
threshold may be one value per head, e.g. the per-head table produced
by an offline calibration such as HEART's EBC.
Scores must be finite: a NaN makes its row total NaN, and only the forced
top-ranked block survives. Validating a threshold array reads it on
the host (one sync per call), negligible next to the attention it gates.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
scores
|
array
|
|
required |
threshold
|
float | array
|
|
required |
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
|
array
|
compensation, expand it with |
array
|
by |
block_causal_mask
¶
Create a block-causal attention mask (block diffusion LLMs).
Positions are grouped into blocks of block_len aligned on absolute
positions (position // block_len). Attention is bidirectional inside
a block and causal across blocks: the query at absolute position
p = offset + i sees key j iff j // block_len <= p // block_len.
This is the mask block-diffusion models such as LLaDA2.x and SDAR are
trained with (LLaDA2.x's reference generate builds exactly this
tril over blocks), and it makes a prefix KV cache exact for them.
block_len=1 reduces to :func:causal_mask.
The prompt is split into blocks like the rest of the sequence. Models with another layout need their own mask: Nemotron-Labs Diffusion, for one, encodes its prefix causally token by token and only the current block bidirectionally.
The mask is additive (0 attend, -inf blocked), as
mx.fast.scaled_dot_product_attention expects. Runtimes that take a
boolean or 0/1 keep mask (mlx-vlm's mask=, Hugging Face
attention_mask) read it inverted; pass block_causal_mask(...) == 0
to them.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
seq_len
|
int
|
Number of query positions. |
required |
block_len
|
int
|
Block size in tokens. |
required |
offset
|
int
|
Offset for KV cache (total KV length = offset + seq_len). |
0
|
dtype
|
Dtype
|
Output dtype. Masked positions are -inf. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Mask of shape (1, 1, seq_len, offset + seq_len). |
Raises:
| Type | Description |
|---|---|
ValueError
|
if |
causal_mask
¶
Create a causal (lower-triangular) attention mask.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
seq_len
|
int
|
Sequence length. |
required |
offset
|
int
|
Offset for KV cache (total KV length = offset + seq_len). |
0
|
dtype
|
Dtype
|
Output dtype. Masked positions are -inf. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Mask of shape (1, 1, seq_len, offset + seq_len). |
Raises:
| Type | Description |
|---|---|
ValueError
|
if |
sliding_window_mask
¶
sliding_window_mask(seq_len: int, window_size: int, offset: int = 0, dtype: Dtype = float32) -> array
Create a sliding window causal attention mask.
Each position can attend to at most window_size previous positions
(including itself).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
seq_len
|
int
|
Sequence length. |
required |
window_size
|
int
|
Size of the attention window. |
required |
offset
|
int
|
Offset for KV cache. |
0
|
dtype
|
Dtype
|
Output dtype. Masked positions are -inf. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Mask of shape (1, 1, seq_len, offset + seq_len). |
Raises:
| Type | Description |
|---|---|
ValueError
|
if |
block_contiguous_permutation
¶
block_contiguous_permutation(scores: array, *, block_size: int, descending: bool = True) -> tuple[array, array]
Sort tokens by score so high-importance ones cluster into early blocks.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
scores
|
array
|
|
required |
block_size
|
int
|
Block size of the downstream sparse kernel.
Informational only — this function does not pad |
required |
descending
|
bool
|
If |
True
|
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
|
tuple[array, array]
|
|
tuple[array, array]
|
Tie-breaking among equal scores is stable (preserves original order), |
tuple[array, array]
|
which is the MLX |
invert_permutation
¶
Compute the inverse of a 1D permutation.
Equivalent to mx.argsort(perm). Caller is responsible for ensuring
perm is a valid permutation of [0, S); misuse silently produces
wrong results.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
perm
|
array
|
|
required |
Returns:
| Type | Description |
|---|---|
array
|
|
classify
¶
classify(scores: array, *, spatial_threshold: float = 0.5, temporal_threshold: float = 0.5) -> list[Kind]
Convert raw per-head mass scores to discrete labels.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
scores
|
array
|
|
required |
spatial_threshold
|
float
|
Min column-0 mass to label a head |
0.5
|
temporal_threshold
|
float
|
Min column-1 mass to label a head |
0.5
|
Returns:
| Type | Description |
|---|---|
list[Kind]
|
List of |
list[Kind]
|
exceed their threshold, |
classify_heads_from_probs
¶
Per-head attention-mass fractions on same-frame and same-position keys.
Uses all queries (no sampling) — assumes the caller has already paid the
cost of materializing (B, num_heads, S, S) softmaxed probabilities.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
probs
|
array
|
|
required |
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
column 1 = mass on same-spatial-position keys. Averaged over batch |
array
|
and queries. |
classify_heads_from_qk
¶
classify_heads_from_qk(q: array, k: array, T: int, H: int, W: int, *, n_samples: int = 64, key: array | None = None) -> array
Per-head attention-mass fractions, sampled from Q,K.
Avoids materializing the full (B, num_heads, S, S) attention by
sampling n_samples queries uniformly per call. Reproducible: with a
fixed key, returns identical results.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
n_samples
|
int
|
How many queries to sample uniformly per call. Must satisfy
|
64
|
key
|
array | None
|
|
None
|
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
column 1 = mass on same-spatial-position keys. |
frame_stride_diagonal_mask
¶
frame_stride_diagonal_mask(T: int, H: int, W: int, *, num_diagonals: int, dtype: Dtype = float32) -> array
Multi-diagonal mask at frame-stride offsets (Sparse-vDiT M3).
Token at flat index i attends to token at flat index j iff
(j - i) is a multiple of the per-frame stride H*W in
{-(k-1)*HW, ..., -HW, 0, HW, ..., (k-1)*HW} where k = num_diagonals.
Captures the "multi-diagonal" head pattern from Sparse-vDiT
(Chen et al. 2025): same (h, w) across nearby frames.
Setting num_diagonals=1 reduces to the main diagonal (self-attention
only). Setting num_diagonals=T is equivalent to temporal_only_mask.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
num_diagonals
|
int
|
Strictly positive number of diagonal bands (counting the main diagonal once). |
required |
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
radial_box_mask
¶
radial_box_mask(T: int, H: int, W: int, *, radius_t: int, radius_s: float, dtype: Dtype = float32) -> array
Hard-cutoff radial spatiotemporal mask.
Query (t, h, w) attends to (t', h', w') iff
|t-t'| <= radius_t AND sqrt((h-h')**2 + (w-w')**2) <= radius_s.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
radius_t
|
int
|
Non-negative temporal radius (frames, inclusive). |
required |
radius_s
|
float
|
Non-negative Euclidean spatial radius (latent units). |
required |
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
radial_gaussian_mask
¶
radial_gaussian_mask(T: int, H: int, W: int, *, sigma_t: float, sigma_s: float, cutoff: float = -6.0, dtype: Dtype = float32) -> array
Exponential-decay radial mask (dense log-weights).
Value at (i, j) is -(dt**2 / (2 sigma_t**2) + ds**2 / (2 sigma_s**2))
where ds**2 = (h-h')**2 + (w-w')**2. Values below cutoff are clamped
to -inf so the mask is usable in fp16 without underflow.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
sigma_t
|
float
|
Temporal scale, strictly positive. |
required |
sigma_s
|
float
|
Spatial scale, strictly positive. |
required |
cutoff
|
float
|
Strictly negative log-weight floor; values below are replaced
by |
-6.0
|
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
sliding_tile_block_mask
¶
sliding_tile_block_mask(T: int, H: int, W: int, *, tile: tuple[int, int, int], window: tuple[int, int, int] = (1, 1, 1), dtype: Dtype = float32) -> array
Tile-block sliding attention (STA, ICML 2025).
Tokens are grouped into non-overlapping tiles of shape
tile = (tt, th, tw). Every query in a tile attends to all keys in the
±window neighboring tiles (window in tile units, inclusive).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. Must be divisible by |
required |
H
|
int
|
Latent height. Must be divisible by |
required |
W
|
int
|
Latent width. Must be divisible by |
required |
tile
|
tuple[int, int, int]
|
|
required |
window
|
tuple[int, int, int]
|
|
(1, 1, 1)
|
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
sliding_tile_centered_mask
¶
sliding_tile_centered_mask(T: int, H: int, W: int, *, window: tuple[int, int, int], dtype: Dtype = float32) -> array
Per-query centered spatiotemporal window mask.
Token (t, h, w) attends to (t', h', w') iff
|t-t'| <= window[0] AND |h-h'| <= window[1] AND |w-w'| <= window[2].
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
window
|
tuple[int, int, int]
|
|
required |
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
spatial_only_mask
¶
Mask that restricts attention to tokens in the same frame.
Each token at frame t attends only to other tokens whose frame index
equals t. Captures the "spatial-locality" head pattern from Sparse
VideoGen.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
temporal_only_mask
¶
Mask that restricts attention to tokens at the same spatial position.
Each token at (h, w) attends only to tokens whose (h, w) matches,
across all frames. Captures the "temporal-locality" head pattern from
Sparse VideoGen.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
vertical_stripe_mask
¶
vertical_stripe_mask(T: int, H: int, W: int, *, key_indices: array, dtype: Dtype = float32) -> array
Anchor-column mask (Sparse-vDiT M4).
Every query attends only to a fixed set of "sink" key tokens identified
by key_indices (flat indices into the T-major sequence). Captures
the "vertical-stripe" head pattern from Sparse-vDiT, where a small set
of anchor positions act as global memory.
The set must be non-empty and contain unique in-range indices. The
main diagonal is not added automatically — include it in
key_indices if self-attention is desired.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
key_indices
|
array
|
1-D |
required |
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
compensation
¶
Sparse-block compensation: dense reference implementations.
Block-sparse attention drops (query cluster, key cluster) blocks. Two training-free methods recover part of the dropped signal:
- SVG-EAR (arXiv 2603.08982): each skipped key is replaced by its
cluster centroid, so the output stays a softmax average over the full
key set. See :func:
centroid_compensated_attention. - SparsePR (arXiv 2608.18484): a few exact "probe" rows are computed
densely and a ridge regression predicts the residual on the other rows.
See :func:
probe_residual_correctionand :func:select_probe_rows.
MLX has no block-sparse kernel, so everything here runs dense: these are
quality tools and numerical references, not speedups. Clusters are
caller-provided integer labels; :func:tile_labels gives the labels that
match the shipped Sliding Tile Attention masks.
tile_labels
¶
Cluster label of every token for a (tt, th, tw) tiling of a video grid.
The label of token (t, h, w) is the row-major index of its tile in the
(T // tt, H // th, W // tw) tile grid; tokens are in T-major order.
Labels line up with :func:~mlx_arsenal.attention.sliding_tile_block_mask
evaluated at tile resolution, which yields the matching cluster-level
block mask:
labels = tile_labels(T, H, W, tile=(tt, th, tw))
block_mask = sliding_tile_block_mask(T // tt, H // th, W // tw, tile=(1, 1, 1), window=window)
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. Must be divisible by |
required |
H
|
int
|
Latent height. Must be divisible by |
required |
W
|
int
|
Latent width. Must be divisible by |
required |
tile
|
tuple[int, int, int]
|
|
required |
Returns:
| Type | Description |
|---|---|
array
|
|
centroid_compensated_attention
¶
centroid_compensated_attention(q: array, k: array, v: array, *, q_labels: array, k_labels: array, block_mask: array, scale: float | None = None) -> array
Block-sparse attention with skipped blocks compensated by key centroids (SVG-EAR).
Queries and keys are grouped into clusters by integer labels. For every
(query cluster, key cluster) block kept by block_mask, attention is
exact. For a skipped block, each key and value is replaced by the mean
key and mean value of its cluster. The result is exactly dense attention
over that modified key set, so every output row is still a convex
softmax average — skipped blocks are approximated, not dropped.
Keys of one cluster share a single centroid logit, so the skipped part
collapses to one extra key per cluster carrying a log n_c bias
(n_c = cluster size). The implementation therefore runs one
mx.fast.scaled_dot_product_attention over [K; K̄] / [V; V̄] with an
additive (Sq, Sk + Ck) mask. It is dense — no speedup — and meant as a
quality tool and a reference for a future block-sparse kernel. The mask
is built once when labels are (S,) and block_mask has no per-batch or
per-head dims; per-head labels or masks materialize it B·H times.
With block_mask all 0 this is exact dense attention; with it all
-inf, attention over the cluster centroids only.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
v
|
array
|
|
required |
q_labels
|
array
|
|
required |
k_labels
|
array
|
|
required |
block_mask
|
array
|
|
required |
scale
|
float | None
|
Logit scale. Defaults to |
None
|
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
|
array
|
when they share it, their promoted type otherwise (e.g. bfloat16 |
array
|
with float16 |
Note
The mask is built in q.dtype (SDPA rejects a mask wider than its
output). In bfloat16 the log n_c bias is therefore rounded to half
a bfloat16 ulp: at most 1/64 in logit units for 55 ≤ n_c < 2981
(4 ≤ log n_c < 8), i.e. up to ~1.6% on that centroid's weight — the
same order as the rounding of the bfloat16 q·k logits themselves.
Pass float32 inputs when this tool is used as a numerical reference.
select_probe_rows
¶
Pick probe query rows for :func:probe_residual_correction, spread across clusters.
Round-robin over the non-empty clusters in increasing label order: pass
r takes, from every cluster that still has unused members, its member
at middle-out rank r (positions sorted, middle first, then alternating
outward), until num_probes rows are taken. When the last pass cannot
serve every cluster, it takes clusters evenly spaced over the label
range, so fewer probes than clusters still span the whole sequence.
Deviation from SparsePR, which takes the row nearest to each query-group
centroid: that needs q and per-head groups. This selection depends only
on the labels, so a single probe_idx serves every head. Pass your own
probe_idx to :func:probe_residual_correction for another policy.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q_labels
|
array
|
|
required |
num_probes
|
int
|
Number of rows to pick, in |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
float32 weights |
tuple[array, array]
|
that cluster). When every cluster gets at least one probe, the |
tuple[array, array]
|
weights sum to |
tuple[array, array]
|
clusters, only the sampled clusters are weighted: the weights sum to |
tuple[array, array]
|
the total size of those clusters, and unsampled clusters have no |
tuple[array, array]
|
say in the fit. |
probe_residual_correction
¶
probe_residual_correction(o_sparse: array, o_probe: array, probe_idx: array, *, rank: int = 16, ridge: float = 0.1, weights: array | None = None) -> array
Repair an approximate attention output from a few exact probe rows (SparsePR).
The residual R = O_dense - O_sparse is measured exactly on the probe
rows, then predicted on every row by a weighted ridge regression on the
standardized sparse output, projected onto the top-rank principal
directions of the probe residuals:
B = (X̄ᵀ W X̄ + ridge·I)⁻¹ X̄ᵀ W R̄ X̄, R̄: standardized / centered probe rows
Ψ = top-`rank` right singular vectors of W^½ R̄
R̂ = μ_R + ((O_sparse - μ_X) / σ_X) B Ψ Ψᵀ
out = O_sparse + R̂ probe rows replaced by `o_probe`
Works with any approximate attention (hard-dropped block-sparse, masked, centroid-compensated). The correction is additive: rows are not renormalized, so the output is no longer a convex softmax average. The fit is stateless — SparsePR refits on every call — and runs per (batch, head) in float32; the solve and SVD use the CPU stream (MLX has no GPU kernels for them). Get the exact probe rows with one dense call:
Deviations from the paper: weights are normalized to sum to 1 (so
ridge does not scale with the number of probes), and the optional
blend factor and norm cap of the reference code (off by default there)
are not exposed.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
o_sparse
|
array
|
|
required |
o_probe
|
array
|
|
required |
probe_idx
|
array
|
|
required |
rank
|
int
|
Rank of the residual projection, in |
16
|
ridge
|
float
|
Ridge strength, |
0.1
|
weights
|
array | None
|
|
None
|
Returns:
| Type | Description |
|---|---|
array
|
|
dynamic_masks
¶
Content-dependent block-sparse attention masks (XAttention-style).
The shipped video masks are static patterns. Dynamic predictors instead score every (query block, key block) pair from the actual Q/K and keep the blocks carrying most of the attention mass. This module implements the XAttention estimator (Xu et al., arXiv 2503.16428) as dense MLX array math, written from the paper's description:
- :func:
antidiagonal_block_scores— cheap per-block attention mass estimate from strided anti-diagonal sums (~1/stride ofQKᵀFLOPs). - :func:
top_p_block_mask— per query block, keep the smallest set of key blocks covering a fractionτof the mass;τmay be per head.
SPADE (Liu et al., arXiv 2608.03335) adds input-adaptive blocking:
:func:block_self_similarity (query cohesion per block),
:func:select_tiling (per-head choice among candidate 3D tilings),
:func:minmax_block_scores (min/max block summaries) and
:func:top_k_block_mask (fixed per-row budget).
The masks are block-level and additive, consumable as block_mask by
:func:~mlx_arsenal.attention.centroid_compensated_attention with
labels = mx.arange(N) // block_size. For reuse across denoising steps
see :class:mlx_arsenal.diffusion.HeadMaskCache.
RBS-Attention (Song et al., arXiv 2609.20971) guards centroid estimates
against mean dilution: :func:radius_bounded_block_scores (centroid and
radius-bounded log scores) and :func:relative_block_mask (keep blocks
within α of the row maximum).
antidiagonal_block_scores
¶
antidiagonal_block_scores(q: array, k: array, *, block_size: int = 128, stride: int = 16, scale: float | None = None) -> array
Estimate the attention mass of every (query block, key block) pair.
Tokens are grouped into contiguous blocks of block_size. Each
stride × stride sub-block of QKᵀ is summarized by the sum along
its anti-diagonal, Σ_r q_{iS+S-1-r} · k_{jS+r}, which touches every
query and key of the sub-block once. These strided logits (scaled by
scale / stride) are softmaxed per row in float32, summed over each
(block_size/stride)² tile, and normalized so every query-block row
sums to 1. XAttention keeps the raw tile sums; the normalization does not
change :func:top_p_block_mask, which is relative to the row total.
The strided matmul costs about 1/stride of a full QKᵀ in FLOPs;
the logits and softmax are 1/stride² of the full attention matrix.
All B·H heads are scored at once: transient memory is roughly
2 · B·H · (Nq/stride)·(Nk/stride) · 4 bytes (logits + softmax) —
≈16 MB per head at N = 32k, stride = 16, but ≈14 GB for a Wan
720p layer (N ≈ 75.8k, H = 40, CFG B = 2). For large models,
call it per head slice (q[:, h0:h1]) and concatenate.
Sequence lengths must be multiples of block_size: pad beforehand.
Zero-padded keys still take softmax mass (as in XAttention's non-causal
mode); mask their blocks out afterwards if that matters. Permute tokens
first (e.g. :func:~mlx_arsenal.attention.block_contiguous_permutation)
if blocks should follow another order.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
block_size
|
int
|
Block size in tokens, a multiple of |
128
|
stride
|
int
|
Anti-diagonal sampling stride, |
16
|
scale
|
float | None
|
Logit scale. Defaults to |
None
|
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
each row summing to 1. |
top_p_block_mask
¶
Keep, per query block, the key blocks covering a fraction τ of the mass.
Key blocks are ranked by score (descending, stable: lower index first on
ties) and a block is kept while the mass ranked before it is below
τ · row_total — so the block crossing the threshold is kept, and every
row keeps at least one block (also when a row is all zeros). This is
XAttention's non-causal selection rule, except that XAttention's scatter
also re-marks key block 0 in every row as a side effect; this function
keeps only the top-p set.
threshold may be one value per head, e.g. the per-head table produced
by an offline calibration such as HEART's EBC.
Scores must be finite: a NaN makes its row total NaN, and only the forced
top-ranked block survives. Validating a threshold array reads it on
the host (one sync per call), negligible next to the attention it gates.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
scores
|
array
|
|
required |
threshold
|
float | array
|
|
required |
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
|
array
|
compensation, expand it with |
array
|
by |
block_self_similarity
¶
Mean pairwise cosine similarity of the tokens inside each contiguous block.
SPADE's query-cohesion summary (arXiv 2608.03335, "SICS"): for each block
of block_size consecutive tokens, the mean cosine similarity over all
pairs of distinct tokens, ignoring zero-norm (padding) tokens. Computed
in O(block_size · D) as (‖Σ x̂‖² − Σ‖x̂‖²) / (n(n−1)) with x̂ the
normalized valid tokens and n their count; blocks with fewer than two
valid tokens score 0. High values mean the block's queries point the same
way, so a per-block summary (e.g. :func:minmax_block_scores) represents
them well.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
array
|
|
required |
block_size
|
int
|
Tokens per block; |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
select_tiling
¶
Pick, per head, the 3D tiling under which the queries are most cohesive.
SPADE's input-adaptive blocking: for each candidate (tt, th, tw) tile,
tokens (T-major, N = T·H·W) are reordered so every tile is contiguous
and the tiles' :func:block_self_similarity is averaged. Each (batch,
head) takes the candidate with the highest mean cohesion (the first one
on ties). To use a choice, reorder q/k/v with
mx.argsort(tile_labels(T, H, W, tile=tiles[c])) along the token axis
and build block masks with block_size = tt·th·tw.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
grid
|
tuple[int, int, int]
|
|
required |
tiles
|
Sequence[tuple[int, int, int]]
|
Candidate tile shapes, each dividing |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
minmax_block_scores
¶
Cheap per-block attention estimate from element-wise min/max summaries.
SPADE's estimator ("DSA"): each block of block_size consecutive
tokens is summarized by its element-wise max and min over tokens, and the
(query block, key block) score is
max((q_max + q_min) · k_maxᵀ, (q_max + q_min) · k_minᵀ) — a
(Nq/bs) × (Nk/bs) matmul pair instead of Nq × Nk. Scores are
unscaled logit estimates (as in the reference): rank them with
:func:top_k_block_mask, or softmax each row before
:func:top_p_block_mask. Reorder tokens first (e.g. by the tiling from
:func:select_tiling) so blocks are meaningful.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
block_size
|
int
|
Tokens per block; |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
top_k_block_mask
¶
Keep the k highest-scoring key blocks of every query-block row.
SPADE's selection rule (a fixed block budget, e.g. 17 % of the key
blocks on Wan 2.1). Ties keep the lower block index; k larger than
the number of key blocks keeps them all.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
scores
|
array
|
|
required |
k
|
int
|
Blocks kept per row, |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
radius_bounded_block_scores
¶
radius_bounded_block_scores(q: array, k: array, *, block_size: int, low: float = 0.5, high: float = 0.9) -> tuple[array, array]
RBS-Attention block scores: centroid and radius-bounded ("rescue").
A block centroid c_b can hide one strongly matching key among keys that
cancel it ("mean dilution"). RBS (arXiv 2609.20971) adds the key block's
radius r_b = max ‖k − c_b‖₂: by Cauchy-Schwarz
qᵀk ≤ qᵀc_b + ‖q‖·r_b for every key of the block. Per query token::
ℓ_base = qᵀc_b / √D
ℓ_rescue = (qᵀc_b + ‖q‖·r_b·β_b) / √D
with β_b = clip((r_b − r_low) / (r_high − r_low), 0, 1) from the
low / high quantiles of the radii of each (batch, head), so only
unusually spread blocks get the bound. Each (query block, key block) score
is the logsumexp of the logits over the query block's tokens — the log of
the paper's Σ exp(ℓ). Select with :func:relative_block_mask per
branch and take the union (mx.maximum of the additive masks), plus any
forced blocks (sinks, diagonal band).
β_b is 1 above r_high and 0 below r_low; when the two
quantiles coincide it is 1 for blocks strictly wider than r_low.
Unlike the other estimators it scores every query token: an
Nq × Nk / block_size matmul, cheaper than QKᵀ by block_size
in FLOPs, but it holds about three (B, H, Nq, Nk / block_size)
float32 tensors at once — ≈1.1 GB per head for a Wan 720p layer
(N ≈ 75.8k, block_size = 64). Call it per head (or per group of
heads) at video scale.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
block_size
|
int
|
Tokens per block; |
required |
low
|
float
|
Radius quantile below which |
0.5
|
high
|
float
|
Radius quantile above which |
0.9
|
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
float32 log scores. Use them with :func: |
tuple[array, array]
|
func: |
tuple[array, array]
|
masses, i.e. |
relative_block_mask
¶
Keep key blocks whose score is at least alpha times the row maximum.
RBS-Attention's selection rule, on log scores (as returned by
:func:radius_bounded_block_scores): keep b when
log S_b ≥ max_b' log S_b' + log α. The paper applies it to each branch
separately (α = 0.22 base, 0.18 rescue, tuned offline on LLM
prompts) and unions the masks with mx.maximum.
Scores must be in the log domain (RBS, or logit-like estimates such as
:func:minmax_block_scores); take mx.log of mass estimates such as
:func:antidiagonal_block_scores first. A row whose scores are all equal
(including all -inf) keeps every block; a NaN row keeps none.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
log_scores
|
array
|
|
required |
alpha
|
float
|
Relative threshold in |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
masks
¶
Attention mask utilities.
causal_mask
¶
Create a causal (lower-triangular) attention mask.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
seq_len
|
int
|
Sequence length. |
required |
offset
|
int
|
Offset for KV cache (total KV length = offset + seq_len). |
0
|
dtype
|
Dtype
|
Output dtype. Masked positions are -inf. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Mask of shape (1, 1, seq_len, offset + seq_len). |
Raises:
| Type | Description |
|---|---|
ValueError
|
if |
block_causal_mask
¶
Create a block-causal attention mask (block diffusion LLMs).
Positions are grouped into blocks of block_len aligned on absolute
positions (position // block_len). Attention is bidirectional inside
a block and causal across blocks: the query at absolute position
p = offset + i sees key j iff j // block_len <= p // block_len.
This is the mask block-diffusion models such as LLaDA2.x and SDAR are
trained with (LLaDA2.x's reference generate builds exactly this
tril over blocks), and it makes a prefix KV cache exact for them.
block_len=1 reduces to :func:causal_mask.
The prompt is split into blocks like the rest of the sequence. Models with another layout need their own mask: Nemotron-Labs Diffusion, for one, encodes its prefix causally token by token and only the current block bidirectionally.
The mask is additive (0 attend, -inf blocked), as
mx.fast.scaled_dot_product_attention expects. Runtimes that take a
boolean or 0/1 keep mask (mlx-vlm's mask=, Hugging Face
attention_mask) read it inverted; pass block_causal_mask(...) == 0
to them.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
seq_len
|
int
|
Number of query positions. |
required |
block_len
|
int
|
Block size in tokens. |
required |
offset
|
int
|
Offset for KV cache (total KV length = offset + seq_len). |
0
|
dtype
|
Dtype
|
Output dtype. Masked positions are -inf. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Mask of shape (1, 1, seq_len, offset + seq_len). |
Raises:
| Type | Description |
|---|---|
ValueError
|
if |
sliding_window_mask
¶
sliding_window_mask(seq_len: int, window_size: int, offset: int = 0, dtype: Dtype = float32) -> array
Create a sliding window causal attention mask.
Each position can attend to at most window_size previous positions
(including itself).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
seq_len
|
int
|
Sequence length. |
required |
window_size
|
int
|
Size of the attention window. |
required |
offset
|
int
|
Offset for KV cache. |
0
|
dtype
|
Dtype
|
Output dtype. Masked positions are -inf. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Mask of shape (1, 1, seq_len, offset + seq_len). |
Raises:
| Type | Description |
|---|---|
ValueError
|
if |
permute
¶
Block-contiguous token permutation (SVG2 semantic permutation).
Reorders a sequence of tokens so that high-importance ones fall into the
first contiguous blocks, which is what block-sparse attention kernels
actually need to realize their savings. Pair with mx.take(x, perm, axis=...)
to permute Q/K/V tensors and mx.take(y, inv_perm, axis=...) to undo.
block_contiguous_permutation
¶
block_contiguous_permutation(scores: array, *, block_size: int, descending: bool = True) -> tuple[array, array]
Sort tokens by score so high-importance ones cluster into early blocks.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
scores
|
array
|
|
required |
block_size
|
int
|
Block size of the downstream sparse kernel.
Informational only — this function does not pad |
required |
descending
|
bool
|
If |
True
|
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
|
tuple[array, array]
|
|
tuple[array, array]
|
Tie-breaking among equal scores is stable (preserves original order), |
tuple[array, array]
|
which is the MLX |
invert_permutation
¶
Compute the inverse of a 1D permutation.
Equivalent to mx.argsort(perm). Caller is responsible for ensuring
perm is a valid permutation of [0, S); misuse silently produces
wrong results.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
perm
|
array
|
|
required |
Returns:
| Type | Description |
|---|---|
array
|
|
profile
¶
Head-pattern profiler for video DiTs.
Classify each attention head as SPATIAL (mass concentrated on same-frame
keys), TEMPORAL (same-position cross-frame), or OTHER (neither).
All functions assume T-major token flattening — same convention as
mlx_arsenal.attention.video_masks: tokens flatten as
[t0(h0w0..hHwW), t1(...), ..., tT(...)] to a sequence of length
S = T*H*W.
Kind
¶
Bases: Enum
Discrete head-pattern label.
classify
¶
classify(scores: array, *, spatial_threshold: float = 0.5, temporal_threshold: float = 0.5) -> list[Kind]
Convert raw per-head mass scores to discrete labels.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
scores
|
array
|
|
required |
spatial_threshold
|
float
|
Min column-0 mass to label a head |
0.5
|
temporal_threshold
|
float
|
Min column-1 mass to label a head |
0.5
|
Returns:
| Type | Description |
|---|---|
list[Kind]
|
List of |
list[Kind]
|
exceed their threshold, |
classify_heads_from_probs
¶
Per-head attention-mass fractions on same-frame and same-position keys.
Uses all queries (no sampling) — assumes the caller has already paid the
cost of materializing (B, num_heads, S, S) softmaxed probabilities.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
probs
|
array
|
|
required |
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
column 1 = mass on same-spatial-position keys. Averaged over batch |
array
|
and queries. |
classify_heads_from_qk
¶
classify_heads_from_qk(q: array, k: array, T: int, H: int, W: int, *, n_samples: int = 64, key: array | None = None) -> array
Per-head attention-mass fractions, sampled from Q,K.
Avoids materializing the full (B, num_heads, S, S) attention by
sampling n_samples queries uniformly per call. Reproducible: with a
fixed key, returns identical results.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
n_samples
|
int
|
How many queries to sample uniformly per call. Must satisfy
|
64
|
key
|
array | None
|
|
None
|
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
column 1 = mass on same-spatial-position keys. |
video_masks
¶
Spatiotemporal attention masks for video diffusion transformers.
All functions in this module assume T-major token flattening: a video
tensor of shape (T, H, W) is flattened to S = T*H*W tokens in the order
[t0(h0w0..hHwW), t1(...), ..., tT(...)]. This matches LTX-Video,
CogVideoX, and the convention used by mlx_arsenal.spatial.patchify.
Each function returns a mask of shape (1, 1, S, S) with float values:
0.0 means the query is allowed to attend to the key, -inf means it is
blocked. The shape broadcasts over batch and head axes expected by
mx.fast.scaled_dot_product_attention.
For typical LTX latents (T=8, H=32, W=32 → S=8192) the mask is
S² ≈ 67M entries. Use dtype=mx.float16 to halve memory.
spatial_only_mask
¶
Mask that restricts attention to tokens in the same frame.
Each token at frame t attends only to other tokens whose frame index
equals t. Captures the "spatial-locality" head pattern from Sparse
VideoGen.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
temporal_only_mask
¶
Mask that restricts attention to tokens at the same spatial position.
Each token at (h, w) attends only to tokens whose (h, w) matches,
across all frames. Captures the "temporal-locality" head pattern from
Sparse VideoGen.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
sliding_tile_centered_mask
¶
sliding_tile_centered_mask(T: int, H: int, W: int, *, window: tuple[int, int, int], dtype: Dtype = float32) -> array
Per-query centered spatiotemporal window mask.
Token (t, h, w) attends to (t', h', w') iff
|t-t'| <= window[0] AND |h-h'| <= window[1] AND |w-w'| <= window[2].
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
window
|
tuple[int, int, int]
|
|
required |
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
sliding_tile_block_mask
¶
sliding_tile_block_mask(T: int, H: int, W: int, *, tile: tuple[int, int, int], window: tuple[int, int, int] = (1, 1, 1), dtype: Dtype = float32) -> array
Tile-block sliding attention (STA, ICML 2025).
Tokens are grouped into non-overlapping tiles of shape
tile = (tt, th, tw). Every query in a tile attends to all keys in the
±window neighboring tiles (window in tile units, inclusive).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. Must be divisible by |
required |
H
|
int
|
Latent height. Must be divisible by |
required |
W
|
int
|
Latent width. Must be divisible by |
required |
tile
|
tuple[int, int, int]
|
|
required |
window
|
tuple[int, int, int]
|
|
(1, 1, 1)
|
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
radial_box_mask
¶
radial_box_mask(T: int, H: int, W: int, *, radius_t: int, radius_s: float, dtype: Dtype = float32) -> array
Hard-cutoff radial spatiotemporal mask.
Query (t, h, w) attends to (t', h', w') iff
|t-t'| <= radius_t AND sqrt((h-h')**2 + (w-w')**2) <= radius_s.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
radius_t
|
int
|
Non-negative temporal radius (frames, inclusive). |
required |
radius_s
|
float
|
Non-negative Euclidean spatial radius (latent units). |
required |
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
frame_stride_diagonal_mask
¶
frame_stride_diagonal_mask(T: int, H: int, W: int, *, num_diagonals: int, dtype: Dtype = float32) -> array
Multi-diagonal mask at frame-stride offsets (Sparse-vDiT M3).
Token at flat index i attends to token at flat index j iff
(j - i) is a multiple of the per-frame stride H*W in
{-(k-1)*HW, ..., -HW, 0, HW, ..., (k-1)*HW} where k = num_diagonals.
Captures the "multi-diagonal" head pattern from Sparse-vDiT
(Chen et al. 2025): same (h, w) across nearby frames.
Setting num_diagonals=1 reduces to the main diagonal (self-attention
only). Setting num_diagonals=T is equivalent to temporal_only_mask.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
num_diagonals
|
int
|
Strictly positive number of diagonal bands (counting the main diagonal once). |
required |
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
vertical_stripe_mask
¶
vertical_stripe_mask(T: int, H: int, W: int, *, key_indices: array, dtype: Dtype = float32) -> array
Anchor-column mask (Sparse-vDiT M4).
Every query attends only to a fixed set of "sink" key tokens identified
by key_indices (flat indices into the T-major sequence). Captures
the "vertical-stripe" head pattern from Sparse-vDiT, where a small set
of anchor positions act as global memory.
The set must be non-empty and contain unique in-range indices. The
main diagonal is not added automatically — include it in
key_indices if self-attention is desired.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
key_indices
|
array
|
1-D |
required |
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |
radial_gaussian_mask
¶
radial_gaussian_mask(T: int, H: int, W: int, *, sigma_t: float, sigma_s: float, cutoff: float = -6.0, dtype: Dtype = float32) -> array
Exponential-decay radial mask (dense log-weights).
Value at (i, j) is -(dt**2 / (2 sigma_t**2) + ds**2 / (2 sigma_s**2))
where ds**2 = (h-h')**2 + (w-w')**2. Values below cutoff are clamped
to -inf so the mask is usable in fp16 without underflow.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
T
|
int
|
Number of frames. |
required |
H
|
int
|
Latent height. |
required |
W
|
int
|
Latent width. |
required |
sigma_t
|
float
|
Temporal scale, strictly positive. |
required |
sigma_s
|
float
|
Spatial scale, strictly positive. |
required |
cutoff
|
float
|
Strictly negative log-weight floor; values below are replaced
by |
-6.0
|
dtype
|
Dtype
|
Output dtype. |
float32
|
Returns:
| Type | Description |
|---|---|
array
|
Additive mask of shape |