Skip to content

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

λ — larger means larger steps and fewer of them. The paper finds 1.5-2 best (≈15 steps on FLUX.1-dev).

required
mode Literal['ot', 'ov']

"ov" (default, the paper's best) or "ot".

'ov'
beta float

EMA factor of the OV moments (paper: 0.3).

0.3
dt_min float

Smallest step (paper: 0.01).

0.01
dt_max float | None

Optional largest step. None (paper) caps at 1 − t.

None
warmup_steps int

Leading steps forced to dt_min while the state still updates; the paper uses 2 (FLUX.1-dev), 3 (SD3.5, Krea) or 0 (FLUX.1-schnell) because early velocities are unreliable.

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 t_start is: it corrects the zero-initialised moments.

0.0

t property

t: float

Current time: 0 = noise, 1 = data (diffusers σ = 1 − t).

steps property

steps: int

Number of :meth:step calls since the last reset.

done property

done: bool

True once t has reached 1.

reset

reset() -> None

Return to t_start and drop the velocity history.

step

step(velocity: array) -> float

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

PerHeadAttentionCache(num_heads: int, num_steps: int, rel_l1_thresh: float)

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

previous_output: array

Last cached attention output. Raises before the first cache_output call.

reset

reset() -> None

Clear all state. Call at the start of each new generation.

should_compute

should_compute(step_index: int, attn_input: array) -> array

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

should_compute_from_summary(step_index: int, summary: array) -> array

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

cache_output(output: array) -> None

Store the full (B, num_heads, ...) attention output for per-head splicing on skip.

PerLayerAttentionCache

PerLayerAttentionCache(num_steps: int, rel_l1_thresh: float)

Stateful per-layer attention output cache.

previous_output property

previous_output: array

Last cached attention output. Raises before the first cache_output call.

reset

reset() -> None

Clear all state. Call at the start of each new generation.

should_compute

should_compute(step_index: int, attn_input: array) -> bool

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

should_compute_from_summary(step_index: int, summary: float) -> bool

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

cache_output(output: array) -> None

Store the attention output from the just-computed step for reuse on skip.

CFGSimilarityProfiler

CFGSimilarityProfiler(num_blocks: int, num_heads: int, *, metric: Metric = 'cosine')

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.

scores property

scores: array

(num_blocks, num_heads) running mean. Zero-count blocks → 0.0.

call_counts property

call_counts: array

(num_blocks,) int32 count of :meth:record calls per block.

reset

reset() -> None

Clear accumulated scores and counts. Call before a new warmup pass.

record

record(block_idx: int, cond: array, uncond: array) -> None

Accumulate per-head cond/uncond similarity for block block_idx.

build_skip_schedule

build_skip_schedule(threshold: float) -> array

Threshold :attr:scores into a (num_blocks, num_heads) skip mask.

CFGSkipController

CFGSkipController(schedule: array)

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.

num_blocks property

num_blocks: int

Number of transformer blocks in the schedule.

num_heads property

num_heads: int

Number of attention heads in the schedule.

from_profiler classmethod

from_profiler(profiler: CFGSimilarityProfiler, threshold: float) -> 'CFGSkipController'

Build a controller from a profiler by thresholding its scores.

should_skip_uncond

should_skip_uncond(block_idx: int) -> array

(num_heads,) bool mask — True heads skip uncond for this block.

apply

apply(block_idx: int, cond_output: array, uncond_output: array) -> array

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.

1000
beta_start float

First beta value. Defaults to 0.00085 (Stable Diffusion convention).

0.00085
beta_end float

Last beta value. Defaults to 0.012.

0.012
beta_schedule BetaSchedule

"scaled_linear" (default) or "linear". "scaled_linear" matches Stable Diffusion's linspace(sqrt(b0), sqrt(bT), N)**2.

'scaled_linear'
prediction_type PredictionType

"epsilon" (vanilla) or "v_prediction" (Salimans & Ho 2022, used by CogVideoX, SD 2.x v-pred).

'v_prediction'
rescale_betas_zero_snr bool

Whether to apply zero-terminal-SNR rescaling to the beta schedule.

True
timestep_spacing TimestepSpacing

"leading" (Stable Diffusion convention) or "trailing" (CogVideoX convention).

'trailing'
set_alpha_to_one bool

Whether the final alpha_cumprod (used as prev when stepping past index 0) is forced to 1.0.

True
clip_sample bool

Whether to clip the predicted x0 to [-1, 1] before computing the step. Set False for latent-space diffusion (default).

False
num_inference_steps int

Default schedule length. Call :meth:set_timesteps to change it later.

50

timesteps property

timesteps: array

Schedule indices in iteration order (descending in t).

set_timesteps

set_timesteps(num_inference_steps: int) -> None

Recompute the timestep schedule for num_inference_steps steps.

step

step(model_output: array, timestep: int | array, sample: array) -> array

Run one deterministic DDIM step.

Parameters:

Name Type Description Default
model_output array

The model's prediction at timestep (epsilon or v depending on prediction_type).

required
timestep int | array

Current step's index (int or 0-d mx.array).

required
sample array

Current noisy sample.

required

Returns:

Type Description
array

The denoised sample for the previous timestep.

add_noise

add_noise(original: array, noise: array, timestep: int | array) -> array

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 δ >= 0; a head refreshes when its drift is strictly greater, or not finite. Raw drift is model-scale dependent (the paper uses 8 or 30); see relative.

required
relative bool

Use :func:qk_drift with relative=True (drift divided by the anchor's L1 norm), making delta scale-free.

False
layer_gate tuple[float, float] | None

Optional (low, high) with 0 <= low <= high <= 1.

None

mask property

mask: array

The current merged mask (after the last :meth:update).

reset

reset() -> None

Clear anchors and masks. Call at the start of each new generation.

should_refresh

should_refresh(q: array, k: array) -> array

Decide which heads rebuild their mask at this step.

Parameters:

Name Type Description Default
q array

(B, H, Nq, D) queries of this step (the tensors the mask predictor sees).

required
k array

(B, H, Nk, D) keys.

required

Returns:

Type Description
array

(B, H) bool, True where the head must use a freshly predicted

array

mask. Must be followed by :meth:update before the next call.

update

update(new_mask: array, refresh: array) -> array

Merge freshly predicted masks into the cache and move anchors.

Parameters:

Name Type Description Default
new_mask array

Mask with leading (B, H) axes (any trailing shape, e.g. a (B, H, Cq, Ck) block mask). Only refreshed heads' entries are used.

required
refresh array

The (B, H) bool returned by :meth:should_refresh.

required

Returns:

Type Description
array

The merged mask, same shape and dtype as new_mask.

If this raises ValueError (bad refresh or mask shape), the step stays pending: call update again with valid arguments, or :meth:reset.

TokenStats

Bases: NamedTuple

Per-position statistics of dLLM logits, as returned by :func:token_stats.

x0 instance-attribute

x0: array

(B, L) int32 proposed token: argmax, or a Gumbel-max sample.

prob instance-attribute

prob: array

(B, L) float32 probability of x0 under the temperature-1 softmax.

entropy instance-attribute

entropy: array

(B, L) float32 entropy (nats) of the temperature-1 softmax.

FlowMatchEulerDiscreteScheduler

FlowMatchEulerDiscreteScheduler(num_train_timesteps: int = 1000, shift: float = 1.0)

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

set_timesteps(num_inference_steps: int, sigmas: ndarray | list[float] | None = None) -> None

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, [0,1]). If omitted, a linspace schedule is generated.

None

step

step(model_output: array, timestep: array, sample: array) -> array

Advance the sample by one Euler step using a velocity prediction.

Parameters:

Name Type Description Default
model_output array

Predicted velocity (same shape as sample).

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

add_noise(original: array, noise: array, sigma: array) -> array

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 numpy.poly1d order (highest degree first), calibrated to map raw L1 distances to a quality budget. None (default) uses the raw distance, as SeaCache does on :func:~mlx_arsenal.diffusion.sea_filter outputs.

None
max_consecutive_skips int | None

Optional cap on back-to-back skips (>= 1). None (default) never forces a compute.

None

previous_residual property

previous_residual: Any

Last cached payload. Raises before the first cache_residual call.

reset

reset() -> None

Clear all state. Call at the start of each new generation.

should_compute

should_compute(step_index: int, modulated_input: array) -> bool

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

cache_residual(residual: Any) -> None

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

TimestepEmbedding(in_channels: int, time_embed_dim: int)

Bases: Module

MLP that projects sinusoidal timestep embeddings into a feature space.

Weight keys: linear1.{weight,bias}, linear2.{weight,bias}.

StableConfidentStopping

StableConfidentStopping(stability_threshold: int = 1, confidence_threshold: float = 0.005)

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, >= 0.

1
confidence_threshold float

Mean-entropy bound (nats), > 0.

0.005

reset

reset() -> None

Forget previous canvases. Call before decoding a new canvas.

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. beta < 1 → strict early / tolerant late (SpeCa default regime).

required
order int

Lagrange polynomial order. Requires order + 1 anchors before predictions become available. Typical values 1-3.

2
epsilon float

Numerical floor in the relative-L2 denominator.

1e-08

reset

reset() -> None

Drop all anchors. Call at the start of each new generation.

threshold

threshold(step_index: int) -> float

Adaptive threshold at step_index (relative-L2² units).

can_predict

can_predict(step_index: int) -> bool

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

extrapolate(step_index: int) -> array

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

accept(step_index: int, predicted: array, actual: array) -> bool

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

record(step_index: int, feature: array) -> None

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

WindowResidualController(num_steps: int)

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

previous_residual: array

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.

reset

reset() -> None

Clear all state. Call at the start of each new generation.

should_refresh

should_refresh(step_index: int, attn_input: array | None = None) -> bool

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

cache_residual(residual: array) -> None

Store the full - window residual from the just-refreshed step for reuse.

splice_heads

splice_heads(new_output: array, cached_output: array, recompute_mask: array) -> array

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

(B, num_heads, ...) newly computed attention output.

required
cached_output array

(B, num_heads, ...) previously cached output, same shape as new_output.

required
recompute_mask array

1D bool array of length num_heads.

required

Returns:

Type Description
array

Spliced tensor with the same shape as new_output.

cfg_head_similarity

cfg_head_similarity(cond: array, uncond: array, *, metric: Metric = 'cosine') -> array

Per-head similarity between conditional and unconditional outputs.

Parameters:

Name Type Description Default
cond array

(B, num_heads, S, D) conditional attention output.

required
uncond array

(B, num_heads, S, D) unconditional output, same shape.

required
metric Metric

"cosine" (default, [-1, 1]) or "relative_l1" (≥ 0).

'cosine'

Returns:

Type Description
array

(num_heads,) float array. Reduced over batch, sequence, and

array

feature axes.

cfg_skip_mask

cfg_skip_mask(scores: array, threshold: float, *, metric: Metric = 'cosine') -> array

Convert similarity scores to a skip-uncond mask.

Parameters:

Name Type Description Default
scores array

(num_heads,) or (num_blocks, num_heads) scores from :func:cfg_head_similarity or :class:CFGSimilarityProfiler.scores.

required
threshold float

Cut-off. With cosine, skip when scores >= threshold; with relative_l1, skip when scores <= threshold.

required
metric Metric

Which inversion to apply.

'cosine'

Returns:

Type Description
array

Bool array with the same shape as scores. True means "this

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 (B, ...).

required
uncond array

Unconditional velocity, same shape.

required
scale float

Nominal guidance scale λ >= 1 (as in classifier_free_guidance).

required
x array

Current sample, same shape.

required
sigma float | array | Any

Current noise level in [0, 1]: a float, a 0-d array (e.g. scheduler.sigmas[i]) or a (B,) array.

required
cap float

Γ >= 1. The paper uses 1.05–1.10 (ImageNet, toy GMM); Γ < 1 would make even the unguided m_c infeasible.

required

Returns:

Type Description
array

The guided velocity, dtype of cond.

pooled_qk

pooled_qk(q: array, k: array) -> tuple[array, array]

Token-mean of queries and keys per head (HEART's drift summary).

Parameters:

Name Type Description Default
q array

(B, H, Nq, D) queries.

required
k array

(B, H, Nk, D) keys, same (B, H, ·, D).

required

Returns:

Type Description
tuple[array, array]

(qbar, kbar), each (B, H, D) float32.

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

(B, H, D) reference (anchor) query summary.

required
kbar_a array

(B, H, D) reference (anchor) key summary.

required
qbar_b array

(B, H, D) current query summary.

required
kbar_b array

(B, H, D) current key summary.

required
relative bool

Normalize by the reference summary's L1 norm.

False

Returns:

Type Description
array

(B, H) float32 drift.

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, >= 0.

required
gen_len int

Number of tokens to generate, >= 1.

required
block_len int

Block size, >= 1.

required
align bool

Align block boundaries on absolute positions.

False

Returns:

Type Description
list[tuple[int, int]]

List of (start, end) index pairs, in order.

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

(B, L) integer proposed tokens (:attr:TokenStats.x0).

required
prob array

(B, L) confidence of x0 (:attr:TokenStats.prob).

required
tokens array

(B, L) integer current tokens.

required
editable array

(B, L) bool, positions that may be revised.

required
threshold float

Minimum confidence, in [0, 1].

required
strict bool

Compare with > (LLaDA2.1 reference, default) or >=.

True

Returns:

Type Description
array

(B, L) bool, a subset of editable & (x0 != tokens).

entropy_bound_transfer

entropy_bound_transfer(entropy: array, candidates: array, bound: float) -> array

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

(B, L) non-negative per-position entropy (e.g. :attr:TokenStats.entropy), lower commits first.

required
candidates array

(B, L) bool, positions that may be committed.

required
bound float

Entropy budget >= 0 (nats).

required

Returns:

Type Description
array

(B, L) bool, a subset of candidates with one or more commits per

array

non-empty row.

factor_transfer

factor_transfer(confidence: array, candidates: array, factor: float) -> array

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

(B, L) score in [0, 1], higher commits first.

required
candidates array

(B, L) bool, positions that may be committed.

required
factor float

Parallelism factor, > 0; larger commits more per step.

required

Returns:

Type Description
array

(B, L) bool, a subset of candidates with one or more commits per

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

(B, L) score, higher commits first (e.g. :attr:TokenStats.prob).

required
candidates array

(B, L) bool, positions that may be committed (masked, and inside the current block).

required
threshold float

Minimum confidence, in [0, 1].

required
force_one bool

Guarantee one commit per non-empty row.

True
strict bool

Use > instead of >=.

False

Returns:

Type Description
array

(B, L) bool, a subset of candidates.

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

(B, L, V) logits, any float dtype.

required
temperature float

Sampling temperature, >= 0. 0 is greedy.

0.0
key array | None

PRNG key, required when temperature > 0; ignored otherwise.

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 -inf first, so prob and entropy are over the remaining vocabulary. Every position must keep at least one finite logit outside suppress_ids: for a position whose remaining logits are all -inf (e.g. a constrained vocabulary equal to the suppressed set), x0 is undefined (argmax of a row of -inf, possibly a suppressed id) and prob is NaN. This is not checked, to avoid a host sync per decoding step.

()

Returns:

Type Description
TokenStats

class:TokenStats (x0, prob, entropy), each (B, L).

topk_transfer

topk_transfer(confidence: array, candidates: array, k: int | array) -> array

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

(B, L) score, higher commits first.

required
candidates array

(B, L) bool, positions that may be committed.

required
k int | array

Non-negative int, or (B,) integer array of per-row counts. An array k is range-checked on the host (one sync per call), negligible next to the model forward it follows.

required

Returns:

Type Description
array

(B, L) bool, a subset of candidates.

transfer_schedule

transfer_schedule(num_masked: array, steps: int) -> array

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

(B,) non-negative integer count of masked positions per row (typically in the current block). Range-checked on the host (one sync per call).

required
steps int

Number of denoising steps, >= 1.

required

Returns:

Type Description
array

(B, steps) int32 quotas.

classifier_free_guidance

classifier_free_guidance(cond: array, uncond: array, scale: float) -> array

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 cond).

required
scale float

Guidance scale.

required

Returns:

Type Description
array

Guided prediction.

euler_step

euler_step(x: array, x0: array, sigma: float, sigma_next: float) -> array

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 sigma_next.

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 base_tokens.

0.95
max_shift float

Shift at max_tokens.

2.05
base_tokens int

Anchor for base_shift.

1024
max_tokens int

Anchor for max_shift.

4096
stretch bool

Rescale so the last non-zero sigma equals 1 - terminal.

True
terminal float

Target terminal for stretching. Must be in [0, 1).

0.1

Returns:

Type Description
list[float]

List of num_steps + 1 sigma values ending at 0.0.

Raises:

Type Description
ValueError

if stretch is enabled with terminal outside [0, 1) (terminal == 1.0 would divide by zero and collapse every non-zero sigma to 1.0).

get_sampling_sigmas

get_sampling_sigmas(num_steps: int, shift: float = 1.0) -> list[float]

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 shift * s / (1 + (shift - 1) * s).

1.0

Returns:

Type Description
list[float]

List of num_steps + 1 sigma values descending from 1.0 to 0.0.

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. (B, H, W, C) or (B, T, H, W, C) (the first-block modulated input reshaped to its latent grid).

required
signal_scale float

a_t, the clean-signal coefficient (>= 0).

required
noise_scale float

b_t, the noise coefficient (>= 0).

required
axes Sequence[int] | None

Axes to filter. Default: every axis except the first (batch) and the last (channels).

None
power_exp float

Exponent p of the clean spectrum. SeaCache uses 2 for images and 3 for video.

2.0
eps float

Regulariser of the spectrum and of the gain (> 0).

1e-16

Returns:

Type Description
array

The filtered tensor, same shape and dtype as x.

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 [cos, sin]; if False, [sin, cos].

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 (len(timesteps), embedding_dim).

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 [0, num_steps].

required
num_steps int

Maximum number of denoising steps, >= 1.

required
t_min float

Final temperature, > 0.

0.4
t_max float

Initial temperature, > 0 (usually >= t_min; a rising schedule is accepted, as in the reference).

0.8

Returns:

Type Description
float

The temperature to divide the logits by.

renoise

renoise(canvas: array, accepted: array, noise: array) -> array

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 canvas's shape and dtype.

uniform_canvas

uniform_canvas(shape: Sequence[int], vocab_size: int, *, key: array | None = None) -> array

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. (B, L), positive sizes.

required
vocab_size int

Number of token ids, >= 1.

required
key array | None

Optional PRNG key.

None

Returns:

Type Description
array

int32 tokens in [0, vocab_size).

geometric_threshold

geometric_threshold(step_index: int, num_steps: int, tau_0: float, beta: float) -> float

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

λ — larger means larger steps and fewer of them. The paper finds 1.5-2 best (≈15 steps on FLUX.1-dev).

required
mode Literal['ot', 'ov']

"ov" (default, the paper's best) or "ot".

'ov'
beta float

EMA factor of the OV moments (paper: 0.3).

0.3
dt_min float

Smallest step (paper: 0.01).

0.01
dt_max float | None

Optional largest step. None (paper) caps at 1 − t.

None
warmup_steps int

Leading steps forced to dt_min while the state still updates; the paper uses 2 (FLUX.1-dev), 3 (SD3.5, Krea) or 0 (FLUX.1-schnell) because early velocities are unreliable.

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 t_start is: it corrects the zero-initialised moments.

0.0
t property
t: float

Current time: 0 = noise, 1 = data (diffusers σ = 1 − t).

steps property
steps: int

Number of :meth:step calls since the last reset.

done property
done: bool

True once t has reached 1.

reset
reset() -> None

Return to t_start and drop the velocity history.

step
step(velocity: array) -> float

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.TeaCacheController but 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:

  1. If i == 0 or i == num_steps - 1 → recompute (boundary).
  2. If mean(abs(prev_input)) is zero → recompute (degenerate).
  3. If mean(abs(input - prev_input)) / mean(abs(prev_input)) >= rel_l1_thresh → recompute.
References

DiTFastAttn — Attention Sharing across Timesteps (AST).

PerLayerAttentionCache

PerLayerAttentionCache(num_steps: int, rel_l1_thresh: float)

Stateful per-layer attention output cache.

previous_output property
previous_output: array

Last cached attention output. Raises before the first cache_output call.

reset
reset() -> None

Clear all state. Call at the start of each new generation.

should_compute
should_compute(step_index: int, attn_input: array) -> bool

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
should_compute_from_summary(step_index: int, summary: float) -> bool

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
cache_output(output: array) -> None

Store the attention output from the just-computed step for reuse on skip.

PerHeadAttentionCache

PerHeadAttentionCache(num_heads: int, num_steps: int, rel_l1_thresh: float)

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
previous_output: array

Last cached attention output. Raises before the first cache_output call.

reset
reset() -> None

Clear all state. Call at the start of each new generation.

should_compute
should_compute(step_index: int, attn_input: array) -> array

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
should_compute_from_summary(step_index: int, summary: array) -> array

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
cache_output(output: array) -> None

Store the full (B, num_heads, ...) attention output for per-head splicing on skip.

splice_heads

splice_heads(new_output: array, cached_output: array, recompute_mask: array) -> array

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

(B, num_heads, ...) newly computed attention output.

required
cached_output array

(B, num_heads, ...) previously cached output, same shape as new_output.

required
recompute_mask array

1D bool array of length num_heads.

required

Returns:

Type Description
array

Spliced tensor with the same shape as new_output.

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.TeaCacheController and 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

CFGSimilarityProfiler(num_blocks: int, num_heads: int, *, metric: Metric = 'cosine')

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.

scores property
scores: array

(num_blocks, num_heads) running mean. Zero-count blocks → 0.0.

call_counts property
call_counts: array

(num_blocks,) int32 count of :meth:record calls per block.

reset
reset() -> None

Clear accumulated scores and counts. Call before a new warmup pass.

record
record(block_idx: int, cond: array, uncond: array) -> None

Accumulate per-head cond/uncond similarity for block block_idx.

build_skip_schedule
build_skip_schedule(threshold: float) -> array

Threshold :attr:scores into a (num_blocks, num_heads) skip mask.

CFGSkipController

CFGSkipController(schedule: array)

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.

num_blocks property
num_blocks: int

Number of transformer blocks in the schedule.

num_heads property
num_heads: int

Number of attention heads in the schedule.

from_profiler classmethod
from_profiler(profiler: CFGSimilarityProfiler, threshold: float) -> 'CFGSkipController'

Build a controller from a profiler by thresholding its scores.

should_skip_uncond
should_skip_uncond(block_idx: int) -> array

(num_heads,) bool mask — True heads skip uncond for this block.

apply
apply(block_idx: int, cond_output: array, uncond_output: array) -> array

Apply the cached schedule to splice cond/uncond outputs (wraps :func:splice_heads).

cfg_head_similarity

cfg_head_similarity(cond: array, uncond: array, *, metric: Metric = 'cosine') -> array

Per-head similarity between conditional and unconditional outputs.

Parameters:

Name Type Description Default
cond array

(B, num_heads, S, D) conditional attention output.

required
uncond array

(B, num_heads, S, D) unconditional output, same shape.

required
metric Metric

"cosine" (default, [-1, 1]) or "relative_l1" (≥ 0).

'cosine'

Returns:

Type Description
array

(num_heads,) float array. Reduced over batch, sequence, and

array

feature axes.

cfg_skip_mask

cfg_skip_mask(scores: array, threshold: float, *, metric: Metric = 'cosine') -> array

Convert similarity scores to a skip-uncond mask.

Parameters:

Name Type Description Default
scores array

(num_heads,) or (num_blocks, num_heads) scores from :func:cfg_head_similarity or :class:CFGSimilarityProfiler.scores.

required
threshold float

Cut-off. With cosine, skip when scores >= threshold; with relative_l1, skip when scores <= threshold.

required
metric Metric

Which inversion to apply.

'cosine'

Returns:

Type Description
array

Bool array with the same shape as scores. True means "this

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.

1000
beta_start float

First beta value. Defaults to 0.00085 (Stable Diffusion convention).

0.00085
beta_end float

Last beta value. Defaults to 0.012.

0.012
beta_schedule BetaSchedule

"scaled_linear" (default) or "linear". "scaled_linear" matches Stable Diffusion's linspace(sqrt(b0), sqrt(bT), N)**2.

'scaled_linear'
prediction_type PredictionType

"epsilon" (vanilla) or "v_prediction" (Salimans & Ho 2022, used by CogVideoX, SD 2.x v-pred).

'v_prediction'
rescale_betas_zero_snr bool

Whether to apply zero-terminal-SNR rescaling to the beta schedule.

True
timestep_spacing TimestepSpacing

"leading" (Stable Diffusion convention) or "trailing" (CogVideoX convention).

'trailing'
set_alpha_to_one bool

Whether the final alpha_cumprod (used as prev when stepping past index 0) is forced to 1.0.

True
clip_sample bool

Whether to clip the predicted x0 to [-1, 1] before computing the step. Set False for latent-space diffusion (default).

False
num_inference_steps int

Default schedule length. Call :meth:set_timesteps to change it later.

50
timesteps property
timesteps: array

Schedule indices in iteration order (descending in t).

set_timesteps
set_timesteps(num_inference_steps: int) -> None

Recompute the timestep schedule for num_inference_steps steps.

step
step(model_output: array, timestep: int | array, sample: array) -> array

Run one deterministic DDIM step.

Parameters:

Name Type Description Default
model_output array

The model's prediction at timestep (epsilon or v depending on prediction_type).

required
timestep int | array

Current step's index (int or 0-d mx.array).

required
sample array

Current noisy sample.

required

Returns:

Type Description
array

The denoised sample for the previous timestep.

add_noise
add_noise(original: array, noise: array, timestep: int | array) -> array

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 (B, ...).

required
uncond array

Unconditional velocity, same shape.

required
scale float

Nominal guidance scale λ >= 1 (as in classifier_free_guidance).

required
x array

Current sample, same shape.

required
sigma float | array | Any

Current noise level in [0, 1]: a float, a 0-d array (e.g. scheduler.sigmas[i]) or a (B,) array.

required
cap float

Γ >= 1. The paper uses 1.05–1.10 (ImageNet, toy GMM); Γ < 1 would make even the unguided m_c infeasible.

required

Returns:

Type Description
array

The guided velocity, dtype of cond.

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 δ >= 0; a head refreshes when its drift is strictly greater, or not finite. Raw drift is model-scale dependent (the paper uses 8 or 30); see relative.

required
relative bool

Use :func:qk_drift with relative=True (drift divided by the anchor's L1 norm), making delta scale-free.

False
layer_gate tuple[float, float] | None

Optional (low, high) with 0 <= low <= high <= 1.

None
mask property
mask: array

The current merged mask (after the last :meth:update).

reset
reset() -> None

Clear anchors and masks. Call at the start of each new generation.

should_refresh
should_refresh(q: array, k: array) -> array

Decide which heads rebuild their mask at this step.

Parameters:

Name Type Description Default
q array

(B, H, Nq, D) queries of this step (the tensors the mask predictor sees).

required
k array

(B, H, Nk, D) keys.

required

Returns:

Type Description
array

(B, H) bool, True where the head must use a freshly predicted

array

mask. Must be followed by :meth:update before the next call.

update
update(new_mask: array, refresh: array) -> array

Merge freshly predicted masks into the cache and move anchors.

Parameters:

Name Type Description Default
new_mask array

Mask with leading (B, H) axes (any trailing shape, e.g. a (B, H, Cq, Ck) block mask). Only refreshed heads' entries are used.

required
refresh array

The (B, H) bool returned by :meth:should_refresh.

required

Returns:

Type Description
array

The merged mask, same shape and dtype as new_mask.

If this raises ValueError (bad refresh or mask shape), the step stays pending: call update again with valid arguments, or :meth:reset.

pooled_qk

pooled_qk(q: array, k: array) -> tuple[array, array]

Token-mean of queries and keys per head (HEART's drift summary).

Parameters:

Name Type Description Default
q array

(B, H, Nq, D) queries.

required
k array

(B, H, Nk, D) keys, same (B, H, ·, D).

required

Returns:

Type Description
tuple[array, array]

(qbar, kbar), each (B, H, D) float32.

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

(B, H, D) reference (anchor) query summary.

required
kbar_a array

(B, H, D) reference (anchor) key summary.

required
qbar_b array

(B, H, D) current query summary.

required
kbar_b array

(B, H, D) current key summary.

required
relative bool

Normalize by the reference summary's L1 norm.

False

Returns:

Type Description
array

(B, H) float32 drift.

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:

  1. per-position statistics of the logits (:func:token_stats);
  2. a commit rule choosing which masked positions to fill (:func:threshold_transfer, :func:topk_transfer, :func:factor_transfer, :func:entropy_bound_transfer);
  3. 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

Bases: NamedTuple

Per-position statistics of dLLM logits, as returned by :func:token_stats.

x0 instance-attribute
x0: array

(B, L) int32 proposed token: argmax, or a Gumbel-max sample.

prob instance-attribute
prob: array

(B, L) float32 probability of x0 under the temperature-1 softmax.

entropy instance-attribute
entropy: array

(B, L) float32 entropy (nats) of the temperature-1 softmax.

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

(B, L, V) logits, any float dtype.

required
temperature float

Sampling temperature, >= 0. 0 is greedy.

0.0
key array | None

PRNG key, required when temperature > 0; ignored otherwise.

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 -inf first, so prob and entropy are over the remaining vocabulary. Every position must keep at least one finite logit outside suppress_ids: for a position whose remaining logits are all -inf (e.g. a constrained vocabulary equal to the suppressed set), x0 is undefined (argmax of a row of -inf, possibly a suppressed id) and prob is NaN. This is not checked, to avoid a host sync per decoding step.

()

Returns:

Type Description
TokenStats

class:TokenStats (x0, prob, entropy), each (B, L).

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

(B, L) score, higher commits first (e.g. :attr:TokenStats.prob).

required
candidates array

(B, L) bool, positions that may be committed (masked, and inside the current block).

required
threshold float

Minimum confidence, in [0, 1].

required
force_one bool

Guarantee one commit per non-empty row.

True
strict bool

Use > instead of >=.

False

Returns:

Type Description
array

(B, L) bool, a subset of candidates.

topk_transfer

topk_transfer(confidence: array, candidates: array, k: int | array) -> array

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

(B, L) score, higher commits first.

required
candidates array

(B, L) bool, positions that may be committed.

required
k int | array

Non-negative int, or (B,) integer array of per-row counts. An array k is range-checked on the host (one sync per call), negligible next to the model forward it follows.

required

Returns:

Type Description
array

(B, L) bool, a subset of candidates.

factor_transfer

factor_transfer(confidence: array, candidates: array, factor: float) -> array

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

(B, L) score in [0, 1], higher commits first.

required
candidates array

(B, L) bool, positions that may be committed.

required
factor float

Parallelism factor, > 0; larger commits more per step.

required

Returns:

Type Description
array

(B, L) bool, a subset of candidates with one or more commits per

array

non-empty row.

entropy_bound_transfer

entropy_bound_transfer(entropy: array, candidates: array, bound: float) -> array

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

(B, L) non-negative per-position entropy (e.g. :attr:TokenStats.entropy), lower commits first.

required
candidates array

(B, L) bool, positions that may be committed.

required
bound float

Entropy budget >= 0 (nats).

required

Returns:

Type Description
array

(B, L) bool, a subset of candidates with one or more commits per

array

non-empty row.

transfer_schedule

transfer_schedule(num_masked: array, steps: int) -> array

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

(B,) non-negative integer count of masked positions per row (typically in the current block). Range-checked on the host (one sync per call).

required
steps int

Number of denoising steps, >= 1.

required

Returns:

Type Description
array

(B, steps) int32 quotas.

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, >= 0.

required
gen_len int

Number of tokens to generate, >= 1.

required
block_len int

Block size, >= 1.

required
align bool

Align block boundaries on absolute positions.

False

Returns:

Type Description
list[tuple[int, int]]

List of (start, end) index pairs, in order.

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

(B, L) integer proposed tokens (:attr:TokenStats.x0).

required
prob array

(B, L) confidence of x0 (:attr:TokenStats.prob).

required
tokens array

(B, L) integer current tokens.

required
editable array

(B, L) bool, positions that may be revised.

required
threshold float

Minimum confidence, in [0, 1].

required
strict bool

Compare with > (LLaDA2.1 reference, default) or >=.

True

Returns:

Type Description
array

(B, L) bool, a subset of editable & (x0 != tokens).

samplers

Stateless samplers and guidance primitives for diffusion denoising.

euler_step

euler_step(x: array, x0: array, sigma: float, sigma_next: float) -> array

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 sigma_next.

classifier_free_guidance

classifier_free_guidance(cond: array, uncond: array, scale: float) -> array

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 cond).

required
scale float

Guidance scale.

required

Returns:

Type Description
array

Guided prediction.

schedulers

Sigma schedules and noise schedulers for flow-matching diffusion.

FlowMatchEulerDiscreteScheduler

FlowMatchEulerDiscreteScheduler(num_train_timesteps: int = 1000, shift: float = 1.0)

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
set_timesteps(num_inference_steps: int, sigmas: ndarray | list[float] | None = None) -> None

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, [0,1]). If omitted, a linspace schedule is generated.

None
step
step(model_output: array, timestep: array, sample: array) -> array

Advance the sample by one Euler step using a velocity prediction.

Parameters:

Name Type Description Default
model_output array

Predicted velocity (same shape as sample).

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
add_noise(original: array, noise: array, sigma: array) -> array

Flow-matching interpolation: sigma * noise + (1 - sigma) * original.

get_sampling_sigmas

get_sampling_sigmas(num_steps: int, shift: float = 1.0) -> list[float]

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 shift * s / (1 + (shift - 1) * s).

1.0

Returns:

Type Description
list[float]

List of num_steps + 1 sigma values descending from 1.0 to 0.0.

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 base_tokens.

0.95
max_shift float

Shift at max_tokens.

2.05
base_tokens int

Anchor for base_shift.

1024
max_tokens int

Anchor for max_shift.

4096
stretch bool

Rescale so the last non-zero sigma equals 1 - terminal.

True
terminal float

Target terminal for stretching. Must be in [0, 1).

0.1

Returns:

Type Description
list[float]

List of num_steps + 1 sigma values ending at 0.0.

Raises:

Type Description
ValueError

if stretch is enabled with terminal outside [0, 1) (terminal == 1.0 would divide by zero and collapse every non-zero sigma to 1.0).

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. (B, H, W, C) or (B, T, H, W, C) (the first-block modulated input reshaped to its latent grid).

required
signal_scale float

a_t, the clean-signal coefficient (>= 0).

required
noise_scale float

b_t, the noise coefficient (>= 0).

required
axes Sequence[int] | None

Axes to filter. Default: every axis except the first (batch) and the last (channels).

None
power_exp float

Exponent p of the clean spectrum. SeaCache uses 2 for images and 3 for video.

2.0
eps float

Regulariser of the spectrum and of the gain (> 0).

1e-16

Returns:

Type Description
array

The filtered tensor, same shape and dtype as x.

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 numpy.poly1d order (highest degree first), calibrated to map raw L1 distances to a quality budget. None (default) uses the raw distance, as SeaCache does on :func:~mlx_arsenal.diffusion.sea_filter outputs.

None
max_consecutive_skips int | None

Optional cap on back-to-back skips (>= 1). None (default) never forces a compute.

None
previous_residual property
previous_residual: Any

Last cached payload. Raises before the first cache_residual call.

reset
reset() -> None

Clear all state. Call at the start of each new generation.

should_compute
should_compute(step_index: int, modulated_input: array) -> bool

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
cache_residual(residual: Any) -> None

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

TimestepEmbedding(in_channels: int, time_embed_dim: int)

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 [cos, sin]; if False, [sin, cos].

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 (len(timesteps), embedding_dim).

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

StableConfidentStopping(stability_threshold: int = 1, confidence_threshold: float = 0.005)

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, >= 0.

1
confidence_threshold float

Mean-entropy bound (nats), > 0.

0.005
reset
reset() -> None

Forget previous canvases. Call before decoding a new canvas.

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 [0, num_steps].

required
num_steps int

Maximum number of denoising steps, >= 1.

required
t_min float

Final temperature, > 0.

0.4
t_max float

Initial temperature, > 0 (usually >= t_min; a rising schedule is accepted, as in the reference).

0.8

Returns:

Type Description
float

The temperature to divide the logits by.

uniform_canvas

uniform_canvas(shape: Sequence[int], vocab_size: int, *, key: array | None = None) -> array

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. (B, L), positive sizes.

required
vocab_size int

Number of token ids, >= 1.

required
key array | None

Optional PRNG key.

None

Returns:

Type Description
array

int32 tokens in [0, vocab_size).

renoise

renoise(canvas: array, accepted: array, noise: array) -> array

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 canvas's shape and dtype.

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_index to its own iterative axis (denoising step, diffusion timestep, etc.).
  • The Lagrange order (number of anchors − 1) and the threshold schedule parameters tau_0 and beta.
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. beta < 1 → strict early / tolerant late (SpeCa default regime).

required
order int

Lagrange polynomial order. Requires order + 1 anchors before predictions become available. Typical values 1-3.

2
epsilon float

Numerical floor in the relative-L2 denominator.

1e-08
reset
reset() -> None

Drop all anchors. Call at the start of each new generation.

threshold
threshold(step_index: int) -> float

Adaptive threshold at step_index (relative-L2² units).

can_predict
can_predict(step_index: int) -> bool

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
extrapolate(step_index: int) -> array

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
accept(step_index: int, predicted: array, actual: array) -> bool

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
record(step_index: int, feature: array) -> None

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

geometric_threshold(step_index: int, num_steps: int, tau_0: float, beta: float) -> float

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 every K steps.
  • :meth:WindowResidualController.scheduled — refresh on an explicit list.
  • :meth:WindowResidualController.adaptive — refresh when the attention input has moved more than rel_l1_thresh since 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

WindowResidualController(num_steps: int)

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
previous_residual: array

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.

reset
reset() -> None

Clear all state. Call at the start of each new generation.

should_refresh
should_refresh(step_index: int, attn_input: array | None = None) -> bool

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
cache_residual(residual: array) -> None

Store the full - window residual from the just-refreshed step for reuse.