Diffusion¶
diffusion
¶
Diffusion primitives: timestep embeddings, schedulers, samplers, caching.
CurvatureAdaptiveStepper
¶
CurvatureAdaptiveStepper(scale: float, *, mode: Literal['ot', 'ov'] = 'ov', beta: float = 0.3, dt_min: float = 0.01, dt_max: float | None = None, warmup_steps: int = 0, t_start: float = 0.0)
Zero-NFE adaptive step-size controller for flow-matching Euler sampling.
Time follows the paper: t runs from 0 (noise) to 1 (data) and the
Euler update is x ← x + dt·u. Both rules are invariant to the sign
of the velocity, so a diffusers-convention output v (σ = 1 − t,
x ← x + (σ_next − σ)·v) can be passed as is. Usage::
stepper = CurvatureAdaptiveStepper(1.75, mode="ov")
while not stepper.done:
sigma = 1.0 - stepper.t
v = model(x, sigma)
dt = stepper.step(v)
x = x - dt * v # σ decreases by dt
Each step size is clipped to [dt_min, min(dt_max, 1 − t)]; the upper
bound wins, so the last step lands exactly on t = 1. A zero norm (no
curvature signal) gives the upper bound. As in the paper, a step that
stops just short of 1 leaves a tiny last step, which still costs one
model evaluation.
Batches share one timeline: norms are taken per sample over all non-batch axes and the smallest step wins (at batch 1 this is the paper's rule).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
scale
|
float
|
|
required |
mode
|
Literal['ot', 'ov']
|
|
'ov'
|
beta
|
float
|
EMA factor of the OV moments (paper: |
0.3
|
dt_min
|
float
|
Smallest step (paper: |
0.01
|
dt_max
|
float | None
|
Optional largest step. |
None
|
warmup_steps
|
int
|
Leading steps forced to |
0
|
t_start
|
float
|
Initial time, e.g. for image-to-image starting mid-way.
The OV bias correction applies to the first step whatever
|
0.0
|
step
¶
Step size to take now from velocity (evaluated at :attr:t).
Advances :attr:t by the returned value. velocity is
(B, ...); its shape must not change during a trajectory. A NaN
velocity raises ValueError (OT: from the step after it).
PerHeadAttentionCache
¶
Stateful per-head attention output cache.
Returns a (num_heads,) bool decision per step. Inputs are assumed to
have the head axis at position 1 — i.e. shape (B, num_heads, ...).
previous_output
property
¶
Last cached attention output. Raises before the first cache_output call.
should_compute
¶
Per-head decide whether to recompute attention at step_index.
Side-effects: advances the stored previous input and per-head summary. Must be called once per step in order.
should_compute_from_summary
¶
Decide per head using a caller-supplied (num_heads,) summary.
summary[h] is the analogue of the per-head delta ratio. Do not
interleave with :meth:should_compute within a single denoising
run: the two methods write semantically different values into the
internal previous-summary slot. Pick one mode per run.
cache_output
¶
Store the full (B, num_heads, ...) attention output for per-head splicing on skip.
PerLayerAttentionCache
¶
Stateful per-layer attention output cache.
previous_output
property
¶
Last cached attention output. Raises before the first cache_output call.
should_compute
¶
Decide whether to recompute attention at step_index.
Side-effects: advances the stored previous input and summary. Must be called once per step in order.
should_compute_from_summary
¶
Decide using a caller-supplied scalar summary instead of a tensor.
summary is the analogue of
mean(abs(input - prev_input)) / mean(abs(prev_input)) — the caller
has already done the math.
Do not interleave with :meth:should_compute within a single
denoising run: the two methods write semantically different values
into the internal previous-summary slot. Pick one mode per run.
cache_output
¶
Store the attention output from the just-computed step for reuse on skip.
CFGSimilarityProfiler
¶
Per-(block, head) running-mean similarity accumulator.
Records the per-head similarity of cond and uncond attention outputs
during a warmup pass and produces a static (num_blocks, num_heads)
skip schedule.
CFGSkipController
¶
Wraps a static (num_blocks, num_heads) bool skip schedule.
True at (b, h) means: for block b, head h skips the
unconditional branch and reuses the conditional output. False means
the head computes uncond normally.
from_profiler
classmethod
¶
from_profiler(profiler: CFGSimilarityProfiler, threshold: float) -> 'CFGSkipController'
Build a controller from a profiler by thresholding its scores.
should_skip_uncond
¶
(num_heads,) bool mask — True heads skip uncond for this block.
apply
¶
Apply the cached schedule to splice cond/uncond outputs (wraps :func:splice_heads).
DDIMScheduler
¶
DDIMScheduler(num_train_timesteps: int = 1000, beta_start: float = 0.00085, beta_end: float = 0.012, beta_schedule: BetaSchedule = 'scaled_linear', prediction_type: PredictionType = 'v_prediction', rescale_betas_zero_snr: bool = True, timestep_spacing: TimestepSpacing = 'trailing', set_alpha_to_one: bool = True, clip_sample: bool = False, num_inference_steps: int = 50)
Deterministic DDIM scheduler (eta=0).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_train_timesteps
|
int
|
Number of diffusion steps used during
training. Defaults to |
1000
|
beta_start
|
float
|
First beta value. Defaults to |
0.00085
|
beta_end
|
float
|
Last beta value. Defaults to |
0.012
|
beta_schedule
|
BetaSchedule
|
|
'scaled_linear'
|
prediction_type
|
PredictionType
|
|
'v_prediction'
|
rescale_betas_zero_snr
|
bool
|
Whether to apply zero-terminal-SNR rescaling to the beta schedule. |
True
|
timestep_spacing
|
TimestepSpacing
|
|
'trailing'
|
set_alpha_to_one
|
bool
|
Whether the final |
True
|
clip_sample
|
bool
|
Whether to clip the predicted |
False
|
num_inference_steps
|
int
|
Default schedule length. Call
:meth: |
50
|
set_timesteps
¶
Recompute the timestep schedule for num_inference_steps steps.
step
¶
Run one deterministic DDIM step.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
model_output
|
array
|
The model's prediction at |
required |
timestep
|
int | array
|
Current step's index (int or 0-d |
required |
sample
|
array
|
Current noisy sample. |
required |
Returns:
| Type | Description |
|---|---|
array
|
The denoised sample for the previous timestep. |
add_noise
¶
Forward-diffuse original to noise level timestep.
HeadMaskCache
¶
HeadMaskCache(delta: float, *, relative: bool = False, layer_gate: tuple[float, float] | None = None)
Per-head sparse-mask cache refreshed on pooled Q/K drift (HEART TMR).
Two-phase protocol, once per denoising step and per layer (and per CFG branch — use one cache each):
refresh = cache.should_refresh(q, k) # (B, H) bool
new_mask = predict(q, k) # caller's predictor
mask = cache.update(new_mask, refresh) # merged per head
The first call refreshes every head. Afterwards a head refreshes when the
drift of its pooled Q/K from its anchor (the step of its last refresh)
exceeds delta; refreshed heads take the new mask and move their
anchor, the others keep their cached mask. The predictor may compute the
new mask for every head (simplest, same cost in dense MLX) or only for
refreshed ones — values for reused heads are ignored.
layer_gate=(low, high) applies HEART's layer-level override per batch
row: if the fraction r of heads marked for refresh is below low
no head refreshes, above high every head does (the paper uses
(0.4, 0.8) to amortize per-layer mask-construction overhead).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
delta
|
float
|
Drift threshold |
required |
relative
|
bool
|
Use :func: |
False
|
layer_gate
|
tuple[float, float] | None
|
Optional |
None
|
should_refresh
¶
Decide which heads rebuild their mask at this step.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
mask. Must be followed by :meth: |
update
¶
Merge freshly predicted masks into the cache and move anchors.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
new_mask
|
array
|
Mask with leading |
required |
refresh
|
array
|
The |
required |
Returns:
| Type | Description |
|---|---|
array
|
The merged mask, same shape and dtype as |
If this raises ValueError (bad refresh or mask shape), the
step stays pending: call update again with valid arguments, or
:meth:reset.
TokenStats
¶
FlowMatchEulerDiscreteScheduler
¶
Stateful flow-matching Euler scheduler.
Mirrors the diffusers FlowMatchEulerDiscreteScheduler API: set_timesteps
then step through timesteps calling step with the model's velocity output.
add_noise forward-interpolates a clean sample toward pure noise.
Convention (diffusers parity): sigmas descend from 1 (noise) to ~0
(clean), with a terminal 0.0 appended so the final step has a valid
sigma_next. step therefore moves the sample toward the clean
manifold when given the flow-matching velocity noise - x0.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_train_timesteps
|
int
|
Number of training timesteps. |
1000
|
shift
|
float
|
Shift factor applied to the sigma schedule. |
1.0
|
set_timesteps
¶
Configure the scheduler for num_inference_steps denoising steps.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_inference_steps
|
int
|
Number of denoising steps. |
required |
sigmas
|
ndarray | list[float] | None
|
Optional custom sigmas (descending, |
None
|
step
¶
Advance the sample by one Euler step using a velocity prediction.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
model_output
|
array
|
Predicted velocity (same shape as |
required |
timestep
|
array
|
Current timestep. Accepted for diffusers API parity; unused (the scheduler tracks its own step index). |
required |
sample
|
array
|
Current noisy sample. |
required |
Returns:
| Type | Description |
|---|---|
array
|
Updated sample after one Euler step. |
add_noise
¶
Flow-matching interpolation: sigma * noise + (1 - sigma) * original.
TeaCacheController
¶
TeaCacheController(num_steps: int, rel_l1_thresh: float, coefficients: Sequence[float] | None = None, *, max_consecutive_skips: int | None = None)
Stateful controller deciding when to skip a transformer forward.
Usage per denoising step::
if controller.should_compute(step_index, modulated_input):
x_in = x
x = transformer_blocks(x)
controller.cache_residual(x - x_in)
else:
x = x + controller.previous_residual
Boundary steps (step_index == 0 and step_index == num_steps - 1)
always compute and reset the accumulator.
With max_consecutive_skips=n, a step that would be the n + 1-th
skip in a row computes instead (and resets the accumulator), bounding how
stale the reused residual can get. The count restarts after every
computed step, whatever triggered it. diffusers' SeaCache uses n = 2.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_steps
|
int
|
Total number of denoising steps in the generation. |
required |
rel_l1_thresh
|
float
|
Skip threshold on the accumulated rescaled L1 distance. |
required |
coefficients
|
Sequence[float] | None
|
Polynomial coefficients in |
None
|
max_consecutive_skips
|
int | None
|
Optional cap on back-to-back skips ( |
None
|
previous_residual
property
¶
Last cached payload. Raises before the first cache_residual call.
should_compute
¶
Decide whether to run the transformer at step_index.
Side-effects: advances the stored previous_modulated_input and the
internal accumulator. Must be called once per step in order.
cache_residual
¶
Store the residual from the just-computed step for reuse on skip.
residual is whatever the caller wants to retrieve later via
previous_residual. Single-tensor models pass an mx.array;
multi-stream models (e.g. LTX-2) pass a tuple or dict. The controller
does not inspect or copy the value.
TimestepEmbedding
¶
Bases: Module
MLP that projects sinusoidal timestep embeddings into a feature space.
Weight keys: linear1.{weight,bias}, linear2.{weight,bias}.
StableConfidentStopping
¶
Stop decoding a canvas once its argmax is stable and confident.
Per batch row, stop when the argmax canvas equals each of the previous
stability_threshold argmax canvases (0 makes every call stable)
and the mean per-position entropy is below confidence_threshold.
Before stability_threshold canvases have been seen, no row is stable.
This is Hugging Face's StableAndConfidentStoppingCriteria
(DiffusionGemma defaults: 1 and 0.005); pass the entropy of the
temperature-scaled logits, as the reference does.
Call once per denoising step; :meth:reset before each new canvas.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
stability_threshold
|
int
|
Number of previous identical canvases, |
1
|
confidence_threshold
|
float
|
Mean-entropy bound (nats), |
0.005
|
VerifiedFeatureCache
¶
VerifiedFeatureCache(num_steps: int, tau_0: float, beta: float, *, order: int = 2, epsilon: float = 1e-08)
Lagrange-extrapolate a tracked feature and verify before accepting.
Usage per step::
cache = VerifiedFeatureCache(num_steps=50, tau_0=0.1, beta=0.5)
for step in range(num_steps):
if cache.can_predict(step):
predicted = cache.extrapolate(step)
actual = run_verifier_layer(state) # caller-supplied
if cache.accept(step, predicted, actual):
# use predicted for downstream layers / skip full forward
cache.record(step, predicted)
continue
feature = run_full_forward(state) # fallback
cache.record(step, feature)
Boundary policy: can_predict returns False at step == 0 and
step == num_steps - 1, forcing a full compute at both ends. This
matches the convention used by other step-aware controllers in
mlx_arsenal.diffusion.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_steps
|
int
|
Total number of iterative steps in the generation. |
required |
tau_0
|
float
|
Base relative-L2² threshold for accepting a draft. |
required |
beta
|
float
|
Geometric schedule base. |
required |
order
|
int
|
Lagrange polynomial order. Requires |
2
|
epsilon
|
float
|
Numerical floor in the relative-L2 denominator. |
1e-08
|
threshold
¶
Adaptive threshold at step_index (relative-L2² units).
can_predict
¶
True if a draft is available for step_index.
Returns False at boundary steps (0 and num_steps - 1) and
whenever fewer than order + 1 anchors have been recorded.
extrapolate
¶
Lagrange-extrapolate the tracked feature at step_index.
Raises RuntimeError if called when can_predict would
return False — the caller is expected to gate on it first.
accept
¶
Compare draft to ground truth on the verification layer.
Uses squared relative-L2 to match the SpeCa formulation:
e = ‖predicted − actual‖²₂ / (‖actual‖²₂ + ε).
Returns True if e <= threshold(step_index). A draft whose
shape differs from actual (the feature changed shape since the
anchors were recorded) is always rejected.
record
¶
Append feature as a new anchor at step_index.
Anchors are stored in a fixed-capacity FIFO of size order + 1.
step_index must be strictly greater than the last recorded
step (anchors are monotonic by construction). A feature whose
shape differs from the previous anchors (e.g. a sequence that grew)
drops them first: forecasting restarts from this single anchor.
WindowResidualController
¶
Step-aware controller for the WA-RS residual cache.
Construct via :meth:fixed, :meth:scheduled, or :meth:adaptive —
the bare __init__ is intentionally not part of the public API.
previous_residual
property
¶
Last cached residual. Raises before the first cache_residual call.
fixed
classmethod
¶
fixed(num_steps: int, *, refresh_every: int) -> WindowResidualController
Refresh on step 0, num_steps - 1, and every refresh_every step.
scheduled
classmethod
¶
scheduled(num_steps: int, *, refresh_steps: Sequence[int]) -> WindowResidualController
Refresh on step 0, num_steps - 1, and every step in refresh_steps.
adaptive
classmethod
¶
adaptive(num_steps: int, *, rel_l1_thresh: float) -> WindowResidualController
Refresh when relative-L1 input delta crosses rel_l1_thresh.
Mirrors :class:TeaCacheController semantics: at non-boundary
step i with previous input p, refresh iff
mean(|input - p|) / mean(|p|) >= rel_l1_thresh. A zero-norm
previous input also forces a refresh.
should_refresh
¶
Decide whether to recompute full attention at step_index.
attn_input is required in adaptive mode and ignored
otherwise. Boundary steps (0 and num_steps - 1) always
return True.
cache_residual
¶
Store the full - window residual from the just-refreshed step for reuse.
splice_heads
¶
Per-head select-between-tensors.
Where recompute_mask[h] is True, take new_output[:, h];
where False, take cached_output[:, h]. The head axis is assumed
to be axis 1.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
new_output
|
array
|
|
required |
cached_output
|
array
|
|
required |
recompute_mask
|
array
|
1D bool array of length |
required |
Returns:
| Type | Description |
|---|---|
array
|
Spliced tensor with the same shape as |
cfg_head_similarity
¶
Per-head similarity between conditional and unconditional outputs.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
cond
|
array
|
|
required |
uncond
|
array
|
|
required |
metric
|
Metric
|
|
'cosine'
|
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
feature axes. |
cfg_skip_mask
¶
Convert similarity scores to a skip-uncond mask.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
scores
|
array
|
|
required |
threshold
|
float
|
Cut-off. With |
required |
metric
|
Metric
|
Which inversion to apply. |
'cosine'
|
Returns:
| Type | Description |
|---|---|
array
|
Bool array with the same shape as |
array
|
head can skip the uncond branch". |
posterior_mean_capped_guidance
¶
posterior_mean_capped_guidance(cond: array, uncond: array, scale: float, *, x: array, sigma: float | array | Any, cap: float) -> array
Classifier-free guidance with PMC-CFG's per-sample posterior-mean cap.
diffusers convention: σ = 1 is noise, the velocity is
v = noise − x0 and the posterior mean (the x0 estimate) is
m = x − σ·v. With Δ = m_c − m_u the guidance increment is::
β = max{β ∈ [0, λ − 1] : ‖m_c + β·Δ‖ ≤ Γ·‖m_c‖}
solved in closed form per sample (norms over all non-batch axes), and the
result is cond + β·(cond − uncond). When the cap does not bind this is
exactly classifier_free_guidance(cond, uncond, scale); σ = 0 (no
gap) gives nominal guidance.
For the paper's convention (t = 1 is data, u = x0 − noise) pass
-u velocities and σ = 1 − t, and negate the result.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
cond
|
array
|
Conditional velocity |
required |
uncond
|
array
|
Unconditional velocity, same shape. |
required |
scale
|
float
|
Nominal guidance scale |
required |
x
|
array
|
Current sample, same shape. |
required |
sigma
|
float | array | Any
|
Current noise level in |
required |
cap
|
float
|
|
required |
Returns:
| Type | Description |
|---|---|
array
|
The guided velocity, dtype of |
pooled_qk
¶
Token-mean of queries and keys per head (HEART's drift summary).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
Returns:
| Type | Description |
|---|---|
tuple[array, array]
|
|
qk_drift
¶
qk_drift(qbar_a: array, kbar_a: array, qbar_b: array, kbar_b: array, *, relative: bool = False) -> array
Per-head drift between two pooled Q/K summaries.
‖q̄_a − q̄_b‖₁ + ‖k̄_a − k̄_b‖₁ over the feature axis (HEART Eq. 4).
This raw distance scales with the magnitude of Q and K, so HEART's
threshold is model-specific (8 or 30 in the paper). relative=True
divides by ‖q̄_a‖₁ + ‖k̄_a‖₁ (the reference summary), which makes the
drift scale-free — an extension, not part of the paper.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
qbar_a
|
array
|
|
required |
kbar_a
|
array
|
|
required |
qbar_b
|
array
|
|
required |
kbar_b
|
array
|
|
required |
relative
|
bool
|
Normalize by the reference summary's L1 norm. |
False
|
Returns:
| Type | Description |
|---|---|
array
|
|
block_ranges
¶
block_ranges(prompt_len: int, gen_len: int, block_len: int, *, align: bool = False) -> list[tuple[int, int]]
Half-open (start, end) blocks covering the generation span.
The span is [prompt_len, prompt_len + gen_len). With align=False
blocks start at prompt_len (LLaDA 1.x, Dream, Fast-dLLM). With
align=True block boundaries sit on absolute multiples of block_len,
so the first block may be partial (dInfer start_block_align); use this
for block-causal models, whose blocks match
:func:~mlx_arsenal.attention.block_causal_mask. The last block is
truncated when the span is not a multiple of block_len.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
prompt_len
|
int
|
Prompt length, |
required |
gen_len
|
int
|
Number of tokens to generate, |
required |
block_len
|
int
|
Block size, |
required |
align
|
bool
|
Align block boundaries on absolute positions. |
False
|
Returns:
| Type | Description |
|---|---|
list[tuple[int, int]]
|
List of |
edit_transfer
¶
edit_transfer(x0: array, prob: array, tokens: array, editable: array, threshold: float, *, strict: bool = True) -> array
Revise committed tokens the model now confidently predicts differently.
Token-to-token (T2T) editing, as in LLaDA2.1/2.2 (editing_threshold),
Nemotron-Labs Diffusion and SGLang's JointThreshold: an editable position
is overwritten with x0 when x0 differs from the current token and
its confidence exceeds threshold. No edit is forced — a step may edit
nothing. Apply it alongside a mask-commit rule:
tokens = mx.where(commit | edit, x0, tokens).
Which positions are editable (committed and outside the prompt in the
LLaDA2.1 reference; mlx-vlm additionally stops at the first EOS when
eos_early_stop is on) and when to stop iterating are caller-side; see
the LLaDA2.1 loop in the dLLM research note. x0 may be the mask token if
it is not suppressed in :func:token_stats.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x0
|
array
|
|
required |
prob
|
array
|
|
required |
tokens
|
array
|
|
required |
editable
|
array
|
|
required |
threshold
|
float
|
Minimum confidence, in |
required |
strict
|
bool
|
Compare with |
True
|
Returns:
| Type | Description |
|---|---|
array
|
|
entropy_bound_transfer
¶
Commit low-entropy candidates within a total entropy budget (EB-Sampler).
Candidates are sorted by entropy (ascending) and the longest prefix is
committed whose entropy, minus its largest term, stays within bound:
sum_{j < i} H_(j) <= bound for every committed i. The lowest-entropy
candidate is always committed. This is the EB-Sampler rule, also used by
DiffusionGemma's EntropyBoundSampler to accept canvas tokens.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
entropy
|
array
|
|
required |
candidates
|
array
|
|
required |
bound
|
float
|
Entropy budget |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
non-empty row. |
factor_transfer
¶
Commit a confidence-dependent number of candidates (Fast-dLLM factor rule).
Candidates are sorted by confidence (descending) and the longest prefix
of length n is committed such that every i <= n satisfies
c_(i) >= 1 - factor / (i + 1), i.e. (i + 1)(1 - c_(i)) <= factor.
The most confident candidate is always committed. This is Fast-dLLM's
get_transfer_index_dynamic, except for its off-by-one: when only the
last sorted candidate fails, the reference commits every candidate; this
function does not.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
confidence
|
array
|
|
required |
candidates
|
array
|
|
required |
factor
|
float
|
Parallelism factor, |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
non-empty row. |
threshold_transfer
¶
threshold_transfer(confidence: array, candidates: array, threshold: float, *, force_one: bool = True, strict: bool = False) -> array
Commit every candidate whose confidence reaches threshold.
With force_one, a row that has candidates but none above the threshold
commits exactly its most confident candidate (first position on ties),
so decoding always progresses. This is Fast-dLLM's get_transfer_index
with threshold (and SGLang's LowConfidence); dInfer instead lowers the
threshold to max - 1e-5, which can commit several near-tied tokens.
The comparison is confidence >= threshold (Fast-dLLM, Nemotron-Labs
Diffusion); pass strict=True for confidence > threshold, which is what
the LLaDA2.x reference generate uses.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
confidence
|
array
|
|
required |
candidates
|
array
|
|
required |
threshold
|
float
|
Minimum confidence, in |
required |
force_one
|
bool
|
Guarantee one commit per non-empty row. |
True
|
strict
|
bool
|
Use |
False
|
Returns:
| Type | Description |
|---|---|
array
|
|
token_stats
¶
token_stats(logits: array, *, temperature: float = 0.0, key: array | None = None, suppress_ids: Sequence[int] = ()) -> TokenStats
Propose a token per position with its confidence and entropy.
x0 is the argmax of the logits when temperature == 0, otherwise a
Gumbel-max sample of softmax(logits / temperature), which is what the
Fast-dLLM / LLaDA add_gumbel_noise + argmax computes. prob — the usual
"low_confidence" score — and entropy always use the un-noised,
temperature-1 softmax, as in the reference implementations.
Everything is computed in float32 (the reference code uses float64,
which the MLX GPU lacks; near-ties may resolve differently). Entropy is
NaN-safe for -inf logits.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
logits
|
array
|
|
required |
temperature
|
float
|
Sampling temperature, |
0.0
|
key
|
array | None
|
PRNG key, required when |
None
|
suppress_ids
|
Sequence[int]
|
Token ids never proposed (e.g. the mask token, or EOS
before the end of the canvas). Their logits are set to |
()
|
Returns:
| Type | Description |
|---|---|
TokenStats
|
class: |
topk_transfer
¶
Commit the k most confident candidates of each row.
Pair with :func:transfer_schedule for LLaDA-style fixed quotas
(k = schedule[:, step]). Each row uses its own k; a k above the
row's candidate count commits all of them.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
confidence
|
array
|
|
required |
candidates
|
array
|
|
required |
k
|
int | array
|
Non-negative int, or |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
transfer_schedule
¶
Per-step commit quota spreading each row's masked count over steps.
Row b commits n_b // steps tokens per step, plus one on its first
n_b % steps steps, so the quotas sum to n_b. This is LLaDA /
Fast-dLLM's get_num_transfer_tokens (the linear-schedule expectation),
computed per row. Use with :func:topk_transfer:
topk_transfer(conf, candidates, schedule[:, step]).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_masked
|
array
|
|
required |
steps
|
int
|
Number of denoising steps, |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
classifier_free_guidance
¶
Apply classifier-free guidance.
Returns uncond + scale * (cond - uncond). A scale of 1.0
yields the conditioned prediction, 0.0 the unconditioned one, and
values greater than 1.0 amplify the conditioning signal.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
cond
|
array
|
Conditioned prediction. |
required |
uncond
|
array
|
Unconditioned prediction (same shape as |
required |
scale
|
float
|
Guidance scale. |
required |
Returns:
| Type | Description |
|---|---|
array
|
Guided prediction. |
euler_step
¶
Single stateless Euler step on an x0-prediction model.
Implements x_{t-1} = x + (sigma_next - sigma) * (x - x0) / sigma.
When sigma == 0 the current sample is already clean, so x0 is
returned directly.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
array
|
Current noisy sample. |
required |
x0
|
array
|
Predicted clean sample. |
required |
sigma
|
float
|
Current noise level. |
required |
sigma_next
|
float
|
Next noise level. |
required |
Returns:
| Type | Description |
|---|---|
array
|
Updated sample at |
dynamic_shift_schedule
¶
dynamic_shift_schedule(num_steps: int, num_tokens: int, base_shift: float = 0.95, max_shift: float = 2.05, base_tokens: int = 1024, max_tokens: int = 4096, stretch: bool = True, terminal: float = 0.1) -> list[float]
Sigma schedule with token-count-dependent shift (LTX-style).
Interpolates a shift factor linearly between base_shift at
base_tokens and max_shift at max_tokens, then applies it to
a descending linspace. Optional terminal stretching matches the last
non-zero sigma to 1 - terminal.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_steps
|
int
|
Number of denoising steps. |
required |
num_tokens
|
int
|
Number of latent tokens (drives the shift). |
required |
base_shift
|
float
|
Shift at |
0.95
|
max_shift
|
float
|
Shift at |
2.05
|
base_tokens
|
int
|
Anchor for |
1024
|
max_tokens
|
int
|
Anchor for |
4096
|
stretch
|
bool
|
Rescale so the last non-zero sigma equals |
True
|
terminal
|
float
|
Target terminal for stretching. Must be in |
0.1
|
Returns:
| Type | Description |
|---|---|
list[float]
|
List of |
Raises:
| Type | Description |
|---|---|
ValueError
|
if |
get_sampling_sigmas
¶
Flow-matching sigma schedule: linspace(1, 0, steps+1) with optional shift.
The schedule includes the terminal 0.0 so pairs are formed via
zip(sigmas[:-1], sigmas[1:]).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_steps
|
int
|
Number of denoising steps. |
required |
shift
|
float
|
Shift factor applied as |
1.0
|
Returns:
| Type | Description |
|---|---|
list[float]
|
List of |
sea_filter
¶
sea_filter(x: array, signal_scale: float, noise_scale: float, *, axes: Sequence[int] | None = None, power_exp: float = 2.0, eps: float = 1e-16) -> array
Apply SeaCache's SEA filter to x over its grid axes.
Each axis gets a 1-D Wiener gain from a power-law clean spectrum
S(f) = 1 / (|f|^p + eps)::
g(f) = a·S(f) / (a²·S(f) + b² + eps)
with f in cycles per sample. The N-D gain is the product of the
per-axis gains, normalised to unit mean over the full frequency grid
(SeaCache Eq. 7), and applied in the Fourier domain in float32.
Coefficients: flow matching a = 1 − σ, b = σ; VP / DDPM
a = √ᾱ, b = √(1 − ᾱ). The references clamp σ to
[1e-6, 1 − 1e-6]. b = 0 leaves x unchanged; a = 0 returns
zeros.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
array
|
Channels-last token grid, e.g. |
required |
signal_scale
|
float
|
|
required |
noise_scale
|
float
|
|
required |
axes
|
Sequence[int] | None
|
Axes to filter. Default: every axis except the first (batch) and the last (channels). |
None
|
power_exp
|
float
|
Exponent |
2.0
|
eps
|
float
|
Regulariser of the spectrum and of the gain ( |
1e-16
|
Returns:
| Type | Description |
|---|---|
array
|
The filtered tensor, same shape and dtype as |
get_timestep_embedding
¶
get_timestep_embedding(timesteps: array, embedding_dim: int, flip_sin_to_cos: bool = True, downscale_freq_shift: float = 0.0, scale: float = 1.0, max_period: float = 10000.0) -> array
Create sinusoidal timestep embeddings.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
timesteps
|
array
|
1D array of timestep values. |
required |
embedding_dim
|
int
|
Dimension of the output embeddings. |
required |
flip_sin_to_cos
|
bool
|
If True, output |
True
|
downscale_freq_shift
|
float
|
Controls delta between frequencies. |
0.0
|
scale
|
float
|
Scaling factor applied before sin/cos. |
1.0
|
max_period
|
float
|
Controls the minimum frequency. |
10000.0
|
Returns:
| Type | Description |
|---|---|
array
|
Array of shape |
linear_temperature
¶
linear_temperature(remaining: int, num_steps: int, *, t_min: float = 0.4, t_max: float = 0.8) -> float
Temperature of a linear schedule indexed by the number of remaining steps.
t = t_min + (t_max - t_min) * remaining / num_steps: the first step
(remaining == num_steps) uses t_max and the schedule decays
towards t_min. Diffusion loops count steps down, so remaining
runs num_steps, ..., 1 (Hugging Face's
LinearTemperatureScheduleLogitsProcessor; DiffusionGemma uses
t_min = 0.4, t_max = 0.8, 48 steps).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
remaining
|
int
|
Steps remaining, in |
required |
num_steps
|
int
|
Maximum number of denoising steps, |
required |
t_min
|
float
|
Final temperature, |
0.4
|
t_max
|
float
|
Initial temperature, |
0.8
|
Returns:
| Type | Description |
|---|---|
float
|
The temperature to divide the logits by. |
renoise
¶
Keep accepted tokens and replace every other position with noise.
where(accepted, canvas, noise), with noise typically a fresh
:func:uniform_canvas. Passing the noise in keeps this function pure and
lets a caller reproduce a reference's random draws.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
canvas
|
array
|
Integer tokens (the accepted canvas). |
required |
accepted
|
array
|
Bool mask, same shape, True where the token is kept. |
required |
noise
|
array
|
Integer replacement tokens, same shape. |
required |
Returns:
| Type | Description |
|---|---|
array
|
Tokens with |
uniform_canvas
¶
Uniformly random tokens, for the initial canvas and for re-noising.
With key=None the global MLX PRNG is used, exactly like
mx.random.randint(0, vocab_size, shape), so a loop that seeds the
global PRNG reproduces a reference implementation's draws call for call.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
shape
|
Sequence[int]
|
Canvas shape, e.g. |
required |
vocab_size
|
int
|
Number of token ids, |
required |
key
|
array | None
|
Optional PRNG key. |
None
|
Returns:
| Type | Description |
|---|---|
array
|
int32 tokens in |
geometric_threshold
¶
SpeCa-style geometric threshold schedule.
τ_t = τ₀ · β^((T - 1 - t) / max(T - 1, 1))
With beta < 1: threshold grows from τ₀·β (strict early) to
τ₀ (tolerant late). With beta > 1: the opposite. beta == 1
yields a constant threshold τ₀. The "right" regime is empirical
and depends on the schedule and model — see the research note at
docs/research/verified-feature-caching.md.
adaptive_steps
¶
CAT-Flow: curvature-adaptive step sizes for flow-matching Euler sampling.
CAT-Flow (arXiv 2609.01746) picks each Euler step size from the velocity the model just returned, with no extra function evaluations: small steps where the trajectory bends, large ones where it is straight. Two rules:
- OT ("over time"):
dt = λ / ‖(u_k − u_{k−1}) / dt_{k−1}‖₂, a finite difference of the velocity, i.e. the trajectory's acceleration. - OV ("over values"): Adam-like moments of the scaled velocity
(1 − t)·u,dt = λ / sqrt(‖m2 − m1²‖₂).
The caller owns the sampling loop; :class:CurvatureAdaptiveStepper only
turns velocities into step sizes. See docs/research/adaptive-flow-steps.md.
References
https://arxiv.org/abs/2609.01746
CurvatureAdaptiveStepper
¶
CurvatureAdaptiveStepper(scale: float, *, mode: Literal['ot', 'ov'] = 'ov', beta: float = 0.3, dt_min: float = 0.01, dt_max: float | None = None, warmup_steps: int = 0, t_start: float = 0.0)
Zero-NFE adaptive step-size controller for flow-matching Euler sampling.
Time follows the paper: t runs from 0 (noise) to 1 (data) and the
Euler update is x ← x + dt·u. Both rules are invariant to the sign
of the velocity, so a diffusers-convention output v (σ = 1 − t,
x ← x + (σ_next − σ)·v) can be passed as is. Usage::
stepper = CurvatureAdaptiveStepper(1.75, mode="ov")
while not stepper.done:
sigma = 1.0 - stepper.t
v = model(x, sigma)
dt = stepper.step(v)
x = x - dt * v # σ decreases by dt
Each step size is clipped to [dt_min, min(dt_max, 1 − t)]; the upper
bound wins, so the last step lands exactly on t = 1. A zero norm (no
curvature signal) gives the upper bound. As in the paper, a step that
stops just short of 1 leaves a tiny last step, which still costs one
model evaluation.
Batches share one timeline: norms are taken per sample over all non-batch axes and the smallest step wins (at batch 1 this is the paper's rule).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
scale
|
float
|
|
required |
mode
|
Literal['ot', 'ov']
|
|
'ov'
|
beta
|
float
|
EMA factor of the OV moments (paper: |
0.3
|
dt_min
|
float
|
Smallest step (paper: |
0.01
|
dt_max
|
float | None
|
Optional largest step. |
None
|
warmup_steps
|
int
|
Leading steps forced to |
0
|
t_start
|
float
|
Initial time, e.g. for image-to-image starting mid-way.
The OV bias correction applies to the first step whatever
|
0.0
|
step
¶
Step size to take now from velocity (evaluated at :attr:t).
Advances :attr:t by the returned value. velocity is
(B, ...); its shape must not change during a trajectory. A NaN
velocity raises ValueError (OT: from the step after it).
attention_cache
¶
Attention output cache (AST-style) for diffusion transformers.
Caches attention sub-layer output across denoising steps and reuses it on the next step when the input has barely changed. Two granularities:
- :class:
PerLayerAttentionCache— scalar similarity, one decision per layer per step. Simpler, mirrors :class:mlx_arsenal.diffusion.TeaCacheControllerbut at the attention sub-layer instead of a whole transformer block. - :class:
PerHeadAttentionCache— per-head similarity, one decision per head per step; skipped heads reuse their slice of the cached output via :func:splice_heads.
Decision rule, at step i:
- If
i == 0ori == num_steps - 1→ recompute (boundary). - If
mean(abs(prev_input))is zero → recompute (degenerate). - If
mean(abs(input - prev_input)) / mean(abs(prev_input)) >= rel_l1_thresh→ recompute.
References
DiTFastAttn — Attention Sharing across Timesteps (AST).
PerLayerAttentionCache
¶
Stateful per-layer attention output cache.
previous_output
property
¶
Last cached attention output. Raises before the first cache_output call.
should_compute
¶
Decide whether to recompute attention at step_index.
Side-effects: advances the stored previous input and summary. Must be called once per step in order.
should_compute_from_summary
¶
Decide using a caller-supplied scalar summary instead of a tensor.
summary is the analogue of
mean(abs(input - prev_input)) / mean(abs(prev_input)) — the caller
has already done the math.
Do not interleave with :meth:should_compute within a single
denoising run: the two methods write semantically different values
into the internal previous-summary slot. Pick one mode per run.
cache_output
¶
Store the attention output from the just-computed step for reuse on skip.
PerHeadAttentionCache
¶
Stateful per-head attention output cache.
Returns a (num_heads,) bool decision per step. Inputs are assumed to
have the head axis at position 1 — i.e. shape (B, num_heads, ...).
previous_output
property
¶
Last cached attention output. Raises before the first cache_output call.
should_compute
¶
Per-head decide whether to recompute attention at step_index.
Side-effects: advances the stored previous input and per-head summary. Must be called once per step in order.
should_compute_from_summary
¶
Decide per head using a caller-supplied (num_heads,) summary.
summary[h] is the analogue of the per-head delta ratio. Do not
interleave with :meth:should_compute within a single denoising
run: the two methods write semantically different values into the
internal previous-summary slot. Pick one mode per run.
cache_output
¶
Store the full (B, num_heads, ...) attention output for per-head splicing on skip.
splice_heads
¶
Per-head select-between-tensors.
Where recompute_mask[h] is True, take new_output[:, h];
where False, take cached_output[:, h]. The head axis is assumed
to be axis 1.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
new_output
|
array
|
|
required |
cached_output
|
array
|
|
required |
recompute_mask
|
array
|
1D bool array of length |
required |
Returns:
| Type | Description |
|---|---|
array
|
Spliced tensor with the same shape as |
cfg_skip
¶
CFG-skip (Attention Sharing across CFG, DiTFastAttn ASC).
Profiles per-head similarity between the conditional and unconditional CFG branches, builds a static schedule, and applies that schedule at runtime to skip the unconditional branch on heads where cond and uncond outputs are near-identical.
Two metrics:
cosine— dot/norm of flattened(B, S, D)per head, in[-1, 1]. Skip when score is at least the threshold. Default; matches the DiTFastAttn ASC literature.relative_l1—mean(|c - u|) / mean(|c|)per head, ≥ 0. Skip when score is at most the threshold. Same family as :class:~mlx_arsenal.diffusion.TeaCacheControllerand the attention output caches, so existing thresholds carry over.
The runtime apply step (:meth:CFGSkipController.apply) bundles a
:func:mlx_arsenal.diffusion.splice_heads call so callers do not have to
import the splice helper directly.
CFGSimilarityProfiler
¶
Per-(block, head) running-mean similarity accumulator.
Records the per-head similarity of cond and uncond attention outputs
during a warmup pass and produces a static (num_blocks, num_heads)
skip schedule.
CFGSkipController
¶
Wraps a static (num_blocks, num_heads) bool skip schedule.
True at (b, h) means: for block b, head h skips the
unconditional branch and reuses the conditional output. False means
the head computes uncond normally.
from_profiler
classmethod
¶
from_profiler(profiler: CFGSimilarityProfiler, threshold: float) -> 'CFGSkipController'
Build a controller from a profiler by thresholding its scores.
should_skip_uncond
¶
(num_heads,) bool mask — True heads skip uncond for this block.
apply
¶
Apply the cached schedule to splice cond/uncond outputs (wraps :func:splice_heads).
cfg_head_similarity
¶
Per-head similarity between conditional and unconditional outputs.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
cond
|
array
|
|
required |
uncond
|
array
|
|
required |
metric
|
Metric
|
|
'cosine'
|
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
feature axes. |
cfg_skip_mask
¶
Convert similarity scores to a skip-uncond mask.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
scores
|
array
|
|
required |
threshold
|
float
|
Cut-off. With |
required |
metric
|
Metric
|
Which inversion to apply. |
'cosine'
|
Returns:
| Type | Description |
|---|---|
array
|
Bool array with the same shape as |
array
|
head can skip the uncond branch". |
ddim
¶
DDIM (Denoising Diffusion Implicit Models) scheduler — MLX port.
Matches the diffusers DDIMScheduler / CogVideoXDDIMScheduler
behaviour for the deterministic case (eta=0). Supports the two
common prediction types (epsilon and v_prediction), the two
common spacing strategies (leading, trailing), and optional
zero-SNR rescaling of the beta schedule.
Use this for ports of DDPM-trained models (CogVideoX, Stable Diffusion
1.x / 2.x, ...). For flow-matching models (LTX, Hunyuan-DiT,
ERNIE-Image), see :class:FlowMatchEulerDiscreteScheduler instead.
DDIMScheduler
¶
DDIMScheduler(num_train_timesteps: int = 1000, beta_start: float = 0.00085, beta_end: float = 0.012, beta_schedule: BetaSchedule = 'scaled_linear', prediction_type: PredictionType = 'v_prediction', rescale_betas_zero_snr: bool = True, timestep_spacing: TimestepSpacing = 'trailing', set_alpha_to_one: bool = True, clip_sample: bool = False, num_inference_steps: int = 50)
Deterministic DDIM scheduler (eta=0).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_train_timesteps
|
int
|
Number of diffusion steps used during
training. Defaults to |
1000
|
beta_start
|
float
|
First beta value. Defaults to |
0.00085
|
beta_end
|
float
|
Last beta value. Defaults to |
0.012
|
beta_schedule
|
BetaSchedule
|
|
'scaled_linear'
|
prediction_type
|
PredictionType
|
|
'v_prediction'
|
rescale_betas_zero_snr
|
bool
|
Whether to apply zero-terminal-SNR rescaling to the beta schedule. |
True
|
timestep_spacing
|
TimestepSpacing
|
|
'trailing'
|
set_alpha_to_one
|
bool
|
Whether the final |
True
|
clip_sample
|
bool
|
Whether to clip the predicted |
False
|
num_inference_steps
|
int
|
Default schedule length. Call
:meth: |
50
|
set_timesteps
¶
Recompute the timestep schedule for num_inference_steps steps.
step
¶
Run one deterministic DDIM step.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
model_output
|
array
|
The model's prediction at |
required |
timestep
|
int | array
|
Current step's index (int or 0-d |
required |
sample
|
array
|
Current noisy sample. |
required |
Returns:
| Type | Description |
|---|---|
array
|
The denoised sample for the previous timestep. |
add_noise
¶
Forward-diffuse original to noise level timestep.
guidance
¶
PMC-CFG: posterior-mean-capped classifier-free guidance for flow matching.
Plain CFG extrapolates v_u + λ(v_c − v_u). With large scales the implied
clean-sample estimate (the posterior mean) overshoots, which shows up as
saturation. PMC-CFG (arXiv 2609.24287) keeps the guidance direction but picks,
per sample and per step, the largest increment β ∈ [0, λ − 1] for which the
guided posterior mean stays within Γ times the norm of the conditional
one. The cap releases by itself late in sampling, when the conditional and
unconditional estimates agree. No extra model evaluations.
References
https://arxiv.org/abs/2609.24287
posterior_mean_capped_guidance
¶
posterior_mean_capped_guidance(cond: array, uncond: array, scale: float, *, x: array, sigma: float | array | Any, cap: float) -> array
Classifier-free guidance with PMC-CFG's per-sample posterior-mean cap.
diffusers convention: σ = 1 is noise, the velocity is
v = noise − x0 and the posterior mean (the x0 estimate) is
m = x − σ·v. With Δ = m_c − m_u the guidance increment is::
β = max{β ∈ [0, λ − 1] : ‖m_c + β·Δ‖ ≤ Γ·‖m_c‖}
solved in closed form per sample (norms over all non-batch axes), and the
result is cond + β·(cond − uncond). When the cap does not bind this is
exactly classifier_free_guidance(cond, uncond, scale); σ = 0 (no
gap) gives nominal guidance.
For the paper's convention (t = 1 is data, u = x0 − noise) pass
-u velocities and σ = 1 − t, and negate the result.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
cond
|
array
|
Conditional velocity |
required |
uncond
|
array
|
Unconditional velocity, same shape. |
required |
scale
|
float
|
Nominal guidance scale |
required |
x
|
array
|
Current sample, same shape. |
required |
sigma
|
float | array | Any
|
Current noise level in |
required |
cap
|
float
|
|
required |
Returns:
| Type | Description |
|---|---|
array
|
The guided velocity, dtype of |
mask_reuse
¶
Head-wise temporal reuse of sparse attention masks (HEART).
Content-dependent sparse masks (e.g.
:func:~mlx_arsenal.attention.top_p_block_mask over
:func:~mlx_arsenal.attention.antidiagonal_block_scores) change slowly
across denoising steps, and some heads change much less than others. HEART
(Temporal Mask Reuse, arXiv 2605.14513) keeps, per head, an anchor: the
token-mean of Q and K at the step where the head's mask was last rebuilt.
At every step the head's drift is the L1 distance between the current
pooled Q/K and its anchor; heads whose drift exceeds δ rebuild their
mask and move their anchor, the others reuse the anchored mask. Because the
anchor only moves on refresh, small drifts accumulate until they cross
δ.
Caller-side, as in the paper: the dense warm-up steps (simply do not use the
cache during them), one cache per layer and per CFG branch, the choice of
δ, and the mask predictor itself.
HeadMaskCache
¶
HeadMaskCache(delta: float, *, relative: bool = False, layer_gate: tuple[float, float] | None = None)
Per-head sparse-mask cache refreshed on pooled Q/K drift (HEART TMR).
Two-phase protocol, once per denoising step and per layer (and per CFG branch — use one cache each):
refresh = cache.should_refresh(q, k) # (B, H) bool
new_mask = predict(q, k) # caller's predictor
mask = cache.update(new_mask, refresh) # merged per head
The first call refreshes every head. Afterwards a head refreshes when the
drift of its pooled Q/K from its anchor (the step of its last refresh)
exceeds delta; refreshed heads take the new mask and move their
anchor, the others keep their cached mask. The predictor may compute the
new mask for every head (simplest, same cost in dense MLX) or only for
refreshed ones — values for reused heads are ignored.
layer_gate=(low, high) applies HEART's layer-level override per batch
row: if the fraction r of heads marked for refresh is below low
no head refreshes, above high every head does (the paper uses
(0.4, 0.8) to amortize per-layer mask-construction overhead).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
delta
|
float
|
Drift threshold |
required |
relative
|
bool
|
Use :func: |
False
|
layer_gate
|
tuple[float, float] | None
|
Optional |
None
|
should_refresh
¶
Decide which heads rebuild their mask at this step.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
mask. Must be followed by :meth: |
update
¶
Merge freshly predicted masks into the cache and move anchors.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
new_mask
|
array
|
Mask with leading |
required |
refresh
|
array
|
The |
required |
Returns:
| Type | Description |
|---|---|
array
|
The merged mask, same shape and dtype as |
If this raises ValueError (bad refresh or mask shape), the
step stays pending: call update again with valid arguments, or
:meth:reset.
pooled_qk
¶
Token-mean of queries and keys per head (HEART's drift summary).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
array
|
|
required |
k
|
array
|
|
required |
Returns:
| Type | Description |
|---|---|
tuple[array, array]
|
|
qk_drift
¶
qk_drift(qbar_a: array, kbar_a: array, qbar_b: array, kbar_b: array, *, relative: bool = False) -> array
Per-head drift between two pooled Q/K summaries.
‖q̄_a − q̄_b‖₁ + ‖k̄_a − k̄_b‖₁ over the feature axis (HEART Eq. 4).
This raw distance scales with the magnitude of Q and K, so HEART's
threshold is model-specific (8 or 30 in the paper). relative=True
divides by ‖q̄_a‖₁ + ‖k̄_a‖₁ (the reference summary), which makes the
drift scale-free — an extension, not part of the paper.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
qbar_a
|
array
|
|
required |
kbar_a
|
array
|
|
required |
qbar_b
|
array
|
|
required |
kbar_b
|
array
|
|
required |
relative
|
bool
|
Normalize by the reference summary's L1 norm. |
False
|
Returns:
| Type | Description |
|---|---|
array
|
|
masked_decode
¶
Commit-step primitives for absorbing-mask diffusion LLMs (dLLMs).
Masked diffusion LLMs (LLaDA, Dream, LLaDA2.x, SDAR, Nemotron-Labs Diffusion…) decode by repeatedly predicting every masked position and committing ("transferring") a subset of the predictions. Every framework (Fast-dLLM, dInfer, SGLang, LMDeploy) splits that step the same way:
- per-position statistics of the logits (:func:
token_stats); - a commit rule choosing which masked positions to fill
(:func:
threshold_transfer, :func:topk_transfer, :func:factor_transfer, :func:entropy_bound_transfer); - a per-step quota and a block schedule (:func:
transfer_schedule, :func:block_ranges).
All functions are pure and batched per row (no batch averaging), and
compute in float32. The caller owns the model forward, the KV cache, the
mask / special token ids and the stopping logic. See
:func:mlx_arsenal.attention.block_causal_mask for block-causal models.
TokenStats
¶
token_stats
¶
token_stats(logits: array, *, temperature: float = 0.0, key: array | None = None, suppress_ids: Sequence[int] = ()) -> TokenStats
Propose a token per position with its confidence and entropy.
x0 is the argmax of the logits when temperature == 0, otherwise a
Gumbel-max sample of softmax(logits / temperature), which is what the
Fast-dLLM / LLaDA add_gumbel_noise + argmax computes. prob — the usual
"low_confidence" score — and entropy always use the un-noised,
temperature-1 softmax, as in the reference implementations.
Everything is computed in float32 (the reference code uses float64,
which the MLX GPU lacks; near-ties may resolve differently). Entropy is
NaN-safe for -inf logits.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
logits
|
array
|
|
required |
temperature
|
float
|
Sampling temperature, |
0.0
|
key
|
array | None
|
PRNG key, required when |
None
|
suppress_ids
|
Sequence[int]
|
Token ids never proposed (e.g. the mask token, or EOS
before the end of the canvas). Their logits are set to |
()
|
Returns:
| Type | Description |
|---|---|
TokenStats
|
class: |
threshold_transfer
¶
threshold_transfer(confidence: array, candidates: array, threshold: float, *, force_one: bool = True, strict: bool = False) -> array
Commit every candidate whose confidence reaches threshold.
With force_one, a row that has candidates but none above the threshold
commits exactly its most confident candidate (first position on ties),
so decoding always progresses. This is Fast-dLLM's get_transfer_index
with threshold (and SGLang's LowConfidence); dInfer instead lowers the
threshold to max - 1e-5, which can commit several near-tied tokens.
The comparison is confidence >= threshold (Fast-dLLM, Nemotron-Labs
Diffusion); pass strict=True for confidence > threshold, which is what
the LLaDA2.x reference generate uses.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
confidence
|
array
|
|
required |
candidates
|
array
|
|
required |
threshold
|
float
|
Minimum confidence, in |
required |
force_one
|
bool
|
Guarantee one commit per non-empty row. |
True
|
strict
|
bool
|
Use |
False
|
Returns:
| Type | Description |
|---|---|
array
|
|
topk_transfer
¶
Commit the k most confident candidates of each row.
Pair with :func:transfer_schedule for LLaDA-style fixed quotas
(k = schedule[:, step]). Each row uses its own k; a k above the
row's candidate count commits all of them.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
confidence
|
array
|
|
required |
candidates
|
array
|
|
required |
k
|
int | array
|
Non-negative int, or |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
factor_transfer
¶
Commit a confidence-dependent number of candidates (Fast-dLLM factor rule).
Candidates are sorted by confidence (descending) and the longest prefix
of length n is committed such that every i <= n satisfies
c_(i) >= 1 - factor / (i + 1), i.e. (i + 1)(1 - c_(i)) <= factor.
The most confident candidate is always committed. This is Fast-dLLM's
get_transfer_index_dynamic, except for its off-by-one: when only the
last sorted candidate fails, the reference commits every candidate; this
function does not.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
confidence
|
array
|
|
required |
candidates
|
array
|
|
required |
factor
|
float
|
Parallelism factor, |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
non-empty row. |
entropy_bound_transfer
¶
Commit low-entropy candidates within a total entropy budget (EB-Sampler).
Candidates are sorted by entropy (ascending) and the longest prefix is
committed whose entropy, minus its largest term, stays within bound:
sum_{j < i} H_(j) <= bound for every committed i. The lowest-entropy
candidate is always committed. This is the EB-Sampler rule, also used by
DiffusionGemma's EntropyBoundSampler to accept canvas tokens.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
entropy
|
array
|
|
required |
candidates
|
array
|
|
required |
bound
|
float
|
Entropy budget |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
array
|
non-empty row. |
transfer_schedule
¶
Per-step commit quota spreading each row's masked count over steps.
Row b commits n_b // steps tokens per step, plus one on its first
n_b % steps steps, so the quotas sum to n_b. This is LLaDA /
Fast-dLLM's get_num_transfer_tokens (the linear-schedule expectation),
computed per row. Use with :func:topk_transfer:
topk_transfer(conf, candidates, schedule[:, step]).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_masked
|
array
|
|
required |
steps
|
int
|
Number of denoising steps, |
required |
Returns:
| Type | Description |
|---|---|
array
|
|
block_ranges
¶
block_ranges(prompt_len: int, gen_len: int, block_len: int, *, align: bool = False) -> list[tuple[int, int]]
Half-open (start, end) blocks covering the generation span.
The span is [prompt_len, prompt_len + gen_len). With align=False
blocks start at prompt_len (LLaDA 1.x, Dream, Fast-dLLM). With
align=True block boundaries sit on absolute multiples of block_len,
so the first block may be partial (dInfer start_block_align); use this
for block-causal models, whose blocks match
:func:~mlx_arsenal.attention.block_causal_mask. The last block is
truncated when the span is not a multiple of block_len.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
prompt_len
|
int
|
Prompt length, |
required |
gen_len
|
int
|
Number of tokens to generate, |
required |
block_len
|
int
|
Block size, |
required |
align
|
bool
|
Align block boundaries on absolute positions. |
False
|
Returns:
| Type | Description |
|---|---|
list[tuple[int, int]]
|
List of |
edit_transfer
¶
edit_transfer(x0: array, prob: array, tokens: array, editable: array, threshold: float, *, strict: bool = True) -> array
Revise committed tokens the model now confidently predicts differently.
Token-to-token (T2T) editing, as in LLaDA2.1/2.2 (editing_threshold),
Nemotron-Labs Diffusion and SGLang's JointThreshold: an editable position
is overwritten with x0 when x0 differs from the current token and
its confidence exceeds threshold. No edit is forced — a step may edit
nothing. Apply it alongside a mask-commit rule:
tokens = mx.where(commit | edit, x0, tokens).
Which positions are editable (committed and outside the prompt in the
LLaDA2.1 reference; mlx-vlm additionally stops at the first EOS when
eos_early_stop is on) and when to stop iterating are caller-side; see
the LLaDA2.1 loop in the dLLM research note. x0 may be the mask token if
it is not suppressed in :func:token_stats.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x0
|
array
|
|
required |
prob
|
array
|
|
required |
tokens
|
array
|
|
required |
editable
|
array
|
|
required |
threshold
|
float
|
Minimum confidence, in |
required |
strict
|
bool
|
Compare with |
True
|
Returns:
| Type | Description |
|---|---|
array
|
|
samplers
¶
Stateless samplers and guidance primitives for diffusion denoising.
euler_step
¶
Single stateless Euler step on an x0-prediction model.
Implements x_{t-1} = x + (sigma_next - sigma) * (x - x0) / sigma.
When sigma == 0 the current sample is already clean, so x0 is
returned directly.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
array
|
Current noisy sample. |
required |
x0
|
array
|
Predicted clean sample. |
required |
sigma
|
float
|
Current noise level. |
required |
sigma_next
|
float
|
Next noise level. |
required |
Returns:
| Type | Description |
|---|---|
array
|
Updated sample at |
classifier_free_guidance
¶
Apply classifier-free guidance.
Returns uncond + scale * (cond - uncond). A scale of 1.0
yields the conditioned prediction, 0.0 the unconditioned one, and
values greater than 1.0 amplify the conditioning signal.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
cond
|
array
|
Conditioned prediction. |
required |
uncond
|
array
|
Unconditioned prediction (same shape as |
required |
scale
|
float
|
Guidance scale. |
required |
Returns:
| Type | Description |
|---|---|
array
|
Guided prediction. |
schedulers
¶
Sigma schedules and noise schedulers for flow-matching diffusion.
FlowMatchEulerDiscreteScheduler
¶
Stateful flow-matching Euler scheduler.
Mirrors the diffusers FlowMatchEulerDiscreteScheduler API: set_timesteps
then step through timesteps calling step with the model's velocity output.
add_noise forward-interpolates a clean sample toward pure noise.
Convention (diffusers parity): sigmas descend from 1 (noise) to ~0
(clean), with a terminal 0.0 appended so the final step has a valid
sigma_next. step therefore moves the sample toward the clean
manifold when given the flow-matching velocity noise - x0.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_train_timesteps
|
int
|
Number of training timesteps. |
1000
|
shift
|
float
|
Shift factor applied to the sigma schedule. |
1.0
|
set_timesteps
¶
Configure the scheduler for num_inference_steps denoising steps.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_inference_steps
|
int
|
Number of denoising steps. |
required |
sigmas
|
ndarray | list[float] | None
|
Optional custom sigmas (descending, |
None
|
step
¶
Advance the sample by one Euler step using a velocity prediction.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
model_output
|
array
|
Predicted velocity (same shape as |
required |
timestep
|
array
|
Current timestep. Accepted for diffusers API parity; unused (the scheduler tracks its own step index). |
required |
sample
|
array
|
Current noisy sample. |
required |
Returns:
| Type | Description |
|---|---|
array
|
Updated sample after one Euler step. |
add_noise
¶
Flow-matching interpolation: sigma * noise + (1 - sigma) * original.
get_sampling_sigmas
¶
Flow-matching sigma schedule: linspace(1, 0, steps+1) with optional shift.
The schedule includes the terminal 0.0 so pairs are formed via
zip(sigmas[:-1], sigmas[1:]).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_steps
|
int
|
Number of denoising steps. |
required |
shift
|
float
|
Shift factor applied as |
1.0
|
Returns:
| Type | Description |
|---|---|
list[float]
|
List of |
dynamic_shift_schedule
¶
dynamic_shift_schedule(num_steps: int, num_tokens: int, base_shift: float = 0.95, max_shift: float = 2.05, base_tokens: int = 1024, max_tokens: int = 4096, stretch: bool = True, terminal: float = 0.1) -> list[float]
Sigma schedule with token-count-dependent shift (LTX-style).
Interpolates a shift factor linearly between base_shift at
base_tokens and max_shift at max_tokens, then applies it to
a descending linspace. Optional terminal stretching matches the last
non-zero sigma to 1 - terminal.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_steps
|
int
|
Number of denoising steps. |
required |
num_tokens
|
int
|
Number of latent tokens (drives the shift). |
required |
base_shift
|
float
|
Shift at |
0.95
|
max_shift
|
float
|
Shift at |
2.05
|
base_tokens
|
int
|
Anchor for |
1024
|
max_tokens
|
int
|
Anchor for |
4096
|
stretch
|
bool
|
Rescale so the last non-zero sigma equals |
True
|
terminal
|
float
|
Target terminal for stretching. Must be in |
0.1
|
Returns:
| Type | Description |
|---|---|
list[float]
|
List of |
Raises:
| Type | Description |
|---|---|
ValueError
|
if |
seacache
¶
SeaCache: spectral-evolution-aware distance for TeaCache-style gating.
SeaCache (Chung et al., CVPR 2026) replaces TeaCache's per-model polynomial rescale with a calibration-free spectral filter. The first-block modulated input is passed through a Wiener-like gain built from the noise schedule and a power-law clean-signal spectrum, and the usual accumulated relative-L1 distance is taken on the filtered tensor. Early (noisy) steps are low-passed, so high-frequency noise no longer inflates the distance.
Combine :func:sea_filter with :class:~mlx_arsenal.diffusion.TeaCacheController
built without coefficients. See docs/research/seacache.md.
References
https://arxiv.org/abs/2602.18993
sea_filter
¶
sea_filter(x: array, signal_scale: float, noise_scale: float, *, axes: Sequence[int] | None = None, power_exp: float = 2.0, eps: float = 1e-16) -> array
Apply SeaCache's SEA filter to x over its grid axes.
Each axis gets a 1-D Wiener gain from a power-law clean spectrum
S(f) = 1 / (|f|^p + eps)::
g(f) = a·S(f) / (a²·S(f) + b² + eps)
with f in cycles per sample. The N-D gain is the product of the
per-axis gains, normalised to unit mean over the full frequency grid
(SeaCache Eq. 7), and applied in the Fourier domain in float32.
Coefficients: flow matching a = 1 − σ, b = σ; VP / DDPM
a = √ᾱ, b = √(1 − ᾱ). The references clamp σ to
[1e-6, 1 − 1e-6]. b = 0 leaves x unchanged; a = 0 returns
zeros.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
array
|
Channels-last token grid, e.g. |
required |
signal_scale
|
float
|
|
required |
noise_scale
|
float
|
|
required |
axes
|
Sequence[int] | None
|
Axes to filter. Default: every axis except the first (batch) and the last (channels). |
None
|
power_exp
|
float
|
Exponent |
2.0
|
eps
|
float
|
Regulariser of the spectrum and of the gain ( |
1e-16
|
Returns:
| Type | Description |
|---|---|
array
|
The filtered tensor, same shape and dtype as |
teacache
¶
TeaCache: timestep-aware residual caching for diffusion transformers.
TeaCache (Liu et al., "Timestep Embedding Aware Cache") accelerates diffusion
inference by reusing a cached residual across timesteps when the modulated
input doesn't move much. The trade-off is governed by rel_l1_thresh:
higher = more skipping = faster but lossier.
Architecture-agnostic mechanism — coefficients are model-specific and live
with each model's port (e.g. inside the LTX-2 / Hunyuan / Flux MLX
implementation). mlx-arsenal ships only the engine.
References
https://github.com/ali-vilab/TeaCache
TeaCacheController
¶
TeaCacheController(num_steps: int, rel_l1_thresh: float, coefficients: Sequence[float] | None = None, *, max_consecutive_skips: int | None = None)
Stateful controller deciding when to skip a transformer forward.
Usage per denoising step::
if controller.should_compute(step_index, modulated_input):
x_in = x
x = transformer_blocks(x)
controller.cache_residual(x - x_in)
else:
x = x + controller.previous_residual
Boundary steps (step_index == 0 and step_index == num_steps - 1)
always compute and reset the accumulator.
With max_consecutive_skips=n, a step that would be the n + 1-th
skip in a row computes instead (and resets the accumulator), bounding how
stale the reused residual can get. The count restarts after every
computed step, whatever triggered it. diffusers' SeaCache uses n = 2.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_steps
|
int
|
Total number of denoising steps in the generation. |
required |
rel_l1_thresh
|
float
|
Skip threshold on the accumulated rescaled L1 distance. |
required |
coefficients
|
Sequence[float] | None
|
Polynomial coefficients in |
None
|
max_consecutive_skips
|
int | None
|
Optional cap on back-to-back skips ( |
None
|
previous_residual
property
¶
Last cached payload. Raises before the first cache_residual call.
should_compute
¶
Decide whether to run the transformer at step_index.
Side-effects: advances the stored previous_modulated_input and the
internal accumulator. Must be called once per step in order.
cache_residual
¶
Store the residual from the just-computed step for reuse on skip.
residual is whatever the caller wants to retrieve later via
previous_residual. Single-tensor models pass an mx.array;
multi-stream models (e.g. LTX-2) pass a tuple or dict. The controller
does not inspect or copy the value.
timestep
¶
Timestep embeddings for diffusion models.
TimestepEmbedding
¶
Bases: Module
MLP that projects sinusoidal timestep embeddings into a feature space.
Weight keys: linear1.{weight,bias}, linear2.{weight,bias}.
get_timestep_embedding
¶
get_timestep_embedding(timesteps: array, embedding_dim: int, flip_sin_to_cos: bool = True, downscale_freq_shift: float = 0.0, scale: float = 1.0, max_period: float = 10000.0) -> array
Create sinusoidal timestep embeddings.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
timesteps
|
array
|
1D array of timestep values. |
required |
embedding_dim
|
int
|
Dimension of the output embeddings. |
required |
flip_sin_to_cos
|
bool
|
If True, output |
True
|
downscale_freq_shift
|
float
|
Controls delta between frequencies. |
0.0
|
scale
|
float
|
Scaling factor applied before sin/cos. |
1.0
|
max_period
|
float
|
Controls the minimum frequency. |
10000.0
|
Returns:
| Type | Description |
|---|---|
array
|
Array of shape |
uniform_decode
¶
Uniform-noise diffusion LLM decoding helpers (DiffusionGemma).
Unlike absorbing-mask dLLMs (see :mod:mlx_arsenal.diffusion.masked_decode),
uniform-noise dLLMs such as DiffusionGemma start from a canvas of random
tokens and re-decide the whole canvas at every step: the lowest-entropy
predictions are accepted (the EB-Sampler rule,
:func:~mlx_arsenal.diffusion.entropy_bound_transfer over all positions)
and every other position is re-noised with fresh uniform tokens. The final
output is the argmax canvas of the last step. Temperature follows a linear
schedule, and decoding stops early once the argmax canvas is stable and
confident.
This module ships the pieces that are not already in masked_decode:
:func:linear_temperature, :func:uniform_canvas, :func:renoise and
:class:StableConfidentStopping. Sampling from softmax(logits / t) is a
single mx.random.categorical call on the caller side; self-conditioning
and the encoder cache are model-specific. Semantics follow Hugging Face's
generation_diffusion_gemma.py.
StableConfidentStopping
¶
Stop decoding a canvas once its argmax is stable and confident.
Per batch row, stop when the argmax canvas equals each of the previous
stability_threshold argmax canvases (0 makes every call stable)
and the mean per-position entropy is below confidence_threshold.
Before stability_threshold canvases have been seen, no row is stable.
This is Hugging Face's StableAndConfidentStoppingCriteria
(DiffusionGemma defaults: 1 and 0.005); pass the entropy of the
temperature-scaled logits, as the reference does.
Call once per denoising step; :meth:reset before each new canvas.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
stability_threshold
|
int
|
Number of previous identical canvases, |
1
|
confidence_threshold
|
float
|
Mean-entropy bound (nats), |
0.005
|
linear_temperature
¶
linear_temperature(remaining: int, num_steps: int, *, t_min: float = 0.4, t_max: float = 0.8) -> float
Temperature of a linear schedule indexed by the number of remaining steps.
t = t_min + (t_max - t_min) * remaining / num_steps: the first step
(remaining == num_steps) uses t_max and the schedule decays
towards t_min. Diffusion loops count steps down, so remaining
runs num_steps, ..., 1 (Hugging Face's
LinearTemperatureScheduleLogitsProcessor; DiffusionGemma uses
t_min = 0.4, t_max = 0.8, 48 steps).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
remaining
|
int
|
Steps remaining, in |
required |
num_steps
|
int
|
Maximum number of denoising steps, |
required |
t_min
|
float
|
Final temperature, |
0.4
|
t_max
|
float
|
Initial temperature, |
0.8
|
Returns:
| Type | Description |
|---|---|
float
|
The temperature to divide the logits by. |
uniform_canvas
¶
Uniformly random tokens, for the initial canvas and for re-noising.
With key=None the global MLX PRNG is used, exactly like
mx.random.randint(0, vocab_size, shape), so a loop that seeds the
global PRNG reproduces a reference implementation's draws call for call.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
shape
|
Sequence[int]
|
Canvas shape, e.g. |
required |
vocab_size
|
int
|
Number of token ids, |
required |
key
|
array | None
|
Optional PRNG key. |
None
|
Returns:
| Type | Description |
|---|---|
array
|
int32 tokens in |
renoise
¶
Keep accepted tokens and replace every other position with noise.
where(accepted, canvas, noise), with noise typically a fresh
:func:uniform_canvas. Passing the noise in keeps this function pure and
lets a caller reproduce a reference's random draws.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
canvas
|
array
|
Integer tokens (the accepted canvas). |
required |
accepted
|
array
|
Bool mask, same shape, True where the token is kept. |
required |
noise
|
array
|
Integer replacement tokens, same shape. |
required |
Returns:
| Type | Description |
|---|---|
array
|
Tokens with |
verified_cache
¶
Verified feature caching (SpeCa-style) for iterative generative models.
Forecast-then-verify cache: at each iterative step (denoising timestep, autoregressive step, etc.) the controller can extrapolate a tracked feature from previously observed anchors via Lagrange polynomial extrapolation, then accept or reject the prediction by comparing it to a freshly computed ground-truth feature (typically the output of a single inner layer — the "verification layer").
The pattern is lossy but bounded — see the SpeCa convergence analysis (Zou et al., 2025, arxiv:2509.11628). For bit-exact reproducibility, do not use.
Architecture-agnostic mechanism. The caller decides:
- Which feature to track (the verification-layer output is a common pick).
- How to map the controller's abstract
step_indexto its own iterative axis (denoising step, diffusion timestep, etc.). - The Lagrange
order(number of anchors − 1) and the threshold schedule parameterstau_0andbeta.
References
SpeCa — https://arxiv.org/abs/2509.11628 TaylorSeer — https://arxiv.org/abs/2503.06923
VerifiedFeatureCache
¶
VerifiedFeatureCache(num_steps: int, tau_0: float, beta: float, *, order: int = 2, epsilon: float = 1e-08)
Lagrange-extrapolate a tracked feature and verify before accepting.
Usage per step::
cache = VerifiedFeatureCache(num_steps=50, tau_0=0.1, beta=0.5)
for step in range(num_steps):
if cache.can_predict(step):
predicted = cache.extrapolate(step)
actual = run_verifier_layer(state) # caller-supplied
if cache.accept(step, predicted, actual):
# use predicted for downstream layers / skip full forward
cache.record(step, predicted)
continue
feature = run_full_forward(state) # fallback
cache.record(step, feature)
Boundary policy: can_predict returns False at step == 0 and
step == num_steps - 1, forcing a full compute at both ends. This
matches the convention used by other step-aware controllers in
mlx_arsenal.diffusion.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_steps
|
int
|
Total number of iterative steps in the generation. |
required |
tau_0
|
float
|
Base relative-L2² threshold for accepting a draft. |
required |
beta
|
float
|
Geometric schedule base. |
required |
order
|
int
|
Lagrange polynomial order. Requires |
2
|
epsilon
|
float
|
Numerical floor in the relative-L2 denominator. |
1e-08
|
threshold
¶
Adaptive threshold at step_index (relative-L2² units).
can_predict
¶
True if a draft is available for step_index.
Returns False at boundary steps (0 and num_steps - 1) and
whenever fewer than order + 1 anchors have been recorded.
extrapolate
¶
Lagrange-extrapolate the tracked feature at step_index.
Raises RuntimeError if called when can_predict would
return False — the caller is expected to gate on it first.
accept
¶
Compare draft to ground truth on the verification layer.
Uses squared relative-L2 to match the SpeCa formulation:
e = ‖predicted − actual‖²₂ / (‖actual‖²₂ + ε).
Returns True if e <= threshold(step_index). A draft whose
shape differs from actual (the feature changed shape since the
anchors were recorded) is always rejected.
record
¶
Append feature as a new anchor at step_index.
Anchors are stored in a fixed-capacity FIFO of size order + 1.
step_index must be strictly greater than the last recorded
step (anchors are monotonic by construction). A feature whose
shape differs from the previous anchors (e.g. a sequence that grew)
drops them first: forecasting restarts from this single anchor.
geometric_threshold
¶
SpeCa-style geometric threshold schedule.
τ_t = τ₀ · β^((T - 1 - t) / max(T - 1, 1))
With beta < 1: threshold grows from τ₀·β (strict early) to
τ₀ (tolerant late). With beta > 1: the opposite. beta == 1
yields a constant threshold τ₀. The "right" regime is empirical
and depends on the schedule and model — see the research note at
docs/research/verified-feature-caching.md.
window_residual
¶
WA-RS controller (DiTFastAttn Window Attention + Residual Sharing).
Decides when to refresh the cached full - window attention residual:
- :meth:
WindowResidualController.fixed— refresh everyKsteps. - :meth:
WindowResidualController.scheduled— refresh on an explicit list. - :meth:
WindowResidualController.adaptive— refresh when the attention input has moved more thanrel_l1_threshsince the previous step (same relative-L1 metric as :class:mlx_arsenal.diffusion.TeaCacheController).
Boundary steps (0 and num_steps - 1) always refresh regardless of
mode. The cached residual itself is opaque — the controller does not look
inside mx.array shapes.
References
DiTFastAttn — Window Attention with Residual Sharing (WA-RS).
WindowResidualController
¶
Step-aware controller for the WA-RS residual cache.
Construct via :meth:fixed, :meth:scheduled, or :meth:adaptive —
the bare __init__ is intentionally not part of the public API.
previous_residual
property
¶
Last cached residual. Raises before the first cache_residual call.
fixed
classmethod
¶
fixed(num_steps: int, *, refresh_every: int) -> WindowResidualController
Refresh on step 0, num_steps - 1, and every refresh_every step.
scheduled
classmethod
¶
scheduled(num_steps: int, *, refresh_steps: Sequence[int]) -> WindowResidualController
Refresh on step 0, num_steps - 1, and every step in refresh_steps.
adaptive
classmethod
¶
adaptive(num_steps: int, *, rel_l1_thresh: float) -> WindowResidualController
Refresh when relative-L1 input delta crosses rel_l1_thresh.
Mirrors :class:TeaCacheController semantics: at non-boundary
step i with previous input p, refresh iff
mean(|input - p|) / mean(|p|) >= rel_l1_thresh. A zero-norm
previous input also forces a refresh.
should_refresh
¶
Decide whether to recompute full attention at step_index.
attn_input is required in adaptive mode and ignored
otherwise. Boundary steps (0 and num_steps - 1) always
return True.
cache_residual
¶
Store the full - window residual from the just-refreshed step for reuse.