nemo_automodel.components.speculative.dflash
nemo_automodel.components.speculative.dflash
DFlash speculative-decoding training components.
DFlash drafts a whole block of tokens in parallel via MASK-token “denoising”
conditioned on the target model’s hidden states, in contrast to EAGLE’s
autoregressive single-step drafting. See
nemo_automodel.components.speculative.dflash.core for the training wrapper.
Submodules
nemo_automodel.components.speculative.dflash.corenemo_automodel.components.speculative.dflash.dflash2_corenemo_automodel.components.speculative.dflash.domino_corenemo_automodel.components.speculative.dflash.draft_kimi_k3nemo_automodel.components.speculative.dflash.draft_qwen3nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2nemo_automodel.components.speculative.dflash.jetspec_corenemo_automodel.components.speculative.dflash.registrynemo_automodel.components.speculative.dflash.target
Package Contents
Classes
Functions
Data
API
Bases: Module
Pairwise path selector over each draft position’s top-k candidates.
Scores S_t(a, b) = U_t(b) + <A(a) * H(h_t), B(b)> for predecessor token
a and candidate b: DFlash’s own logit plus a low-rank bilinear match
between the two tokens’ codebook embeddings, gated by the draft hidden state.
The codebooks are plain [vocab, rank] parameters rather than
nn.Embedding modules so the saved keys are candidate_selector.*_codebook
— the names the published DFlash 2 drafters and the serving runtimes use —
instead of candidate_selector.*_codebook.weight.
Parameters:
Size of the (shared with the target) token vocabulary.
Channel count of the draft hidden states.
Codebook / gate width (selector_rank; 256 in DFlash 2).
Candidates kept per position (selector_top_k; 16 in DFlash 2).
Score every (predecessor, candidate) pair for a batch of draft positions.
Fully parallel: no position depends on another’s score. Training calls this
once with the ground-truth predecessors; the decode-time walk in
walk calls it one position at a time with the token it just picked.
Parameters:
Tensor of shape […, hidden]; the draft hidden state at each scored position, with arbitrary leading dimensions.
Tensor of shape […, candidates]; U_t(b), the draft logit of
each candidate, with the same leading dimensions as hidden.
Long tensor of shape […, candidates]; the candidate
token ids, with the same leading dimensions as hidden.
Long tensor of shape […]; the token preceding each
scored position, with the same leading dimensions as hidden.
Returns: torch.Tensor
Tensor of shape […, candidates] holding S_t(a, b).
Reset the selector to a no-op: every score collapses to the draft logit.
S_t is bilinear in the two codebooks, so collapsing it needs one factor
at zero. Zeroing the successor codebook does that; it still receives
gradient immediately, while the predecessor codebook and the context gate
multiply it and therefore only start training on the second step.
Trace one coherent path through the per-position candidate lists.
Below GREEDY_TEMPERATURE_EPS (matching sample’s own greedy
threshold — dividing by a positive-but-tiny temperature blows the scores
up enough that the softmax can turn to NaN) this follows the best
successor at each step; otherwise the step is sampled from the softmax
over the candidate scores, and the returned per-step distribution is the
draft proposal q that dflash2_rejection_sample needs to stay
lossless.
Parameters:
Tensor of shape [batch, draft, hidden]; the draft hidden states of the block’s predicted positions.
Tensor of shape [batch, draft, vocab]; the draft logits at those positions.
Long tensor of shape [batch]; the last verified token, i.e. the predecessor of draft position 0.
Sampling temperature; below GREEDY_TEMPERATURE_EPS
selects greedily.
Returns: torch.Tensor
Tuple (path, candidate_ids, draft_probs): path is a Long tensor
Per-step training outputs for the DFlash 2 draft.
loss / accuracy / valid_tokens mirror DFlashStepMetrics so the
shared DFlash training loop consumes them unchanged. The primary accuracy and
acceptance-length fields describe the selector path — what the draft
actually emits at decode time — and the base_* fields the backbone’s own
top-1 picks, so the two are directly comparable on the same denominator.
Scalar tensor containing selector-path mean acceptance length.
Scalar tensor containing selector-path additive acceptance length.
Scalar tensor containing selector-path greedy token accuracy.
Scalar tensor containing backbone mean acceptance length.
Scalar tensor containing backbone additive acceptance length.
Scalar tensor containing backbone top-1 token accuracy.
Scalar tensor containing backbone correct-token count.
Scalar tensor containing the backbone block-CE term.
Scalar tensor containing the fraction of supervised positions whose true token is in the backbone’s top-k candidates — the ceiling the selector can reach.
Scalar tensor containing selector-path correct-token count.
Scalar tensor containing the differentiable training loss.
Scalar tensor containing the effective loss denominator.
Scalar tensor containing the candidate-selection CE term.
Scalar tensor containing the number of evaluated draft blocks.
Scalar tensor containing the supervised-token count.
Bases: DFlashTrainerModule
DFlash 2 online training wrapper: DFlash block CE + candidate-selection CE.
Block-position decay weights for the block_size - 1 predicted positions.
Parameters:
Tensor of shape [batch, blocks, depth]; the supervised-position mask the weights are multiplied into.
Returns: torch.Tensor
Tensor of shape [batch, blocks, depth] equal to mask scaled by
Score each supervised position’s top-k candidates against its true predecessor.
Teacher-forces the predecessor: position k is scored after the token at
anchor + k - 1, which is what the decode-time walk would have committed
if every earlier position had been accepted. Position 1’s predecessor is the
anchor token itself, exactly as at decode time.
Parameters:
Tensor of shape [batch, blocks, depth, hidden]; the draft hidden states of the predicted (non-anchor) block positions.
Tensor of shape [batch, blocks, depth, vocab]; the backbone logits at those positions.
Long tensor of shape [batch, blocks, block_size]; the
ground-truth token at anchor + k for block position k, so
depth == block_size - 1.
Returns: torch.Tensor
Tuple (scores, candidate_ids, target_index, has_target): scores
Parallel block-wise training forward with the DFlash 2 path selector.
Sequence packing (position_ids [B, S] per-document reset positions,
seq_lens [B, max_docs] document lengths, doc_remaining [B, S])
is handled by the shared DFlash prologue, which keeps every block inside one
document.
Parameters:
Long tensor of shape [batch, sequence]; the context tokens.
Tensor of shape [batch, sequence, layers * hidden]; the captured target-model context features.
Tensor of shape [batch, sequence]; the supervised-token mask.
Long tensor of shape [batch, sequence] with per-document
reset positions under packing, or None.
Long tensor of shape [batch, max_docs] with packed document
lengths, or None when unpacked.
Long tensor of shape [batch, sequence]; real tokens left
in each position’s document, or None when unpacked.
Returns: DFlash2StepMetrics
DFlash2StepMetrics for this micro-batch.
How to build a DFlash draft model for a particular target architecture.
The attention backends this draft can run. The
trainer’s mask format is driven by the same knob and must agree (a
flex BlockMask only works with the flex attention function, a
dense additive mask only with sdpa / eager).
Builds the draft config from the target’s text
config; called with the keyword arguments num_draft_layers,
num_target_layers, block_size, dflash_config, and
attention_backend.
Extra from_pretrained keyword arguments for the
frozen target, derived from the recipe’s recipe_args.
The DFlash 2 draft class for this target: the same backbone
plus the in-block convolutions and the candidate selector. None
when the family has no DFlash 2 draft, which makes the DFlash 2
recipe reject the target rather than silently training the plain
DFlash architecture under a DFlash 2 config.
The draft model class.
Whether the frozen target can be sharded by
the DFlash context-parallel path, which installs a key/value-gather
hook on the target’s SDPA call. False for a target that shards the
sequence itself (it declares _owns_cp_attention and never routes
through that hook). This is declared here rather than read off the
loaded target because the recipe’s other CP gates — which force the
target onto HuggingFace SDPA — run before the target exists, and
would otherwise report a misleading error first.
Per-step training outputs for the DFlash draft.
The count and sum fields retain the additive statistics needed for
token-weighted, distributed validation. Averaging accuracy or
accept_len per micro-batch would bias batches with fewer valid tokens or blocks.
Scalar tensor containing mean bonus-token-inclusive acceptance length.
Scalar tensor containing the additive acceptance-length sum.
Scalar tensor containing greedy token accuracy.
Scalar tensor containing the greedy-correct token count.
Scalar tensor containing the differentiable training loss.
Scalar tensor containing the effective loss denominator.
Scalar tensor containing the number of evaluated draft blocks.
Scalar tensor containing the supervised-token count.
Target-model context features needed by the DFlash trainer.
position_ids / seq_lens / doc_remaining are None off the
packing path and carry the (unshifted) packing metadata to the trainer on it.
Bases: Module
DFlash online training wrapper with block-wise CE loss.
Per-block ground-truth tokens and the supervised-position mask.
Returns (label_indices, target_ids, block_mask) each of shape
[B, N, block_size]. label_indices[..., k] is the sequence position
block position k predicts (anchor + label_start + k); block_mask
is the product of block validity, in-bounds, and the gathered loss mask.
Shared by the block-wise trainers (DFlash, JetSpec, and Domino, which passes
label_start=1 for shift_label) so the label/mask gathering lives in
one place. Under packing, doc_remaining [B, S] truncates labels at
the anchor’s document boundary: anchor sampling only keeps offsets up to
block_size - 1 inside the anchor’s document, and a shifted label window
(label_start > 0) reaches one past that guarantee.
Embed each block as [anchor_token, MASK, MASK, ...] (invalid blocks all MASK).
Position ids for the parallel draft blocks (anchor position + offset).
Without packing the anchor’s position equals its row index, so the block
positions are anchor + offset. Under packing context_position_ids
[B, S] holds per-document reset positions, so the block’s base position
is gathered from it at the anchor (context_position_ids[anchor] + offset)
to keep the draft’s RoPE phase document-local.
Embed the draft blocks with a visible prefix of real tokens, then MASK.
The variable-prefix analogue of _create_noise_embed: block
positions < prefix_lengths hold the real sequence tokens, the rest
(and every position of an invalid block) hold MASK.
Parameters:
Long tensor of shape [batch, sequence].
Long tensor of shape [batch, blocks]; each block’s
start position in the sequence.
Bool tensor of shape [batch, blocks]; invalid
(padding) blocks are embedded as all MASK.
Long tensor of shape [batch, blocks]; this block’s
visible-prefix length.
Returns:
Tensor of shape [batch, blocks * block_size, hidden].
Embed noise-block token ids and apply the target’s input-embedding scale.
Mirrors Qwen3DFlashDraftModel.embed_noise_block, which spec_generate
uses on the decode side: a target whose dflash_config sets
input_embedding_scale must see the same scaled embeddings during
training, or the draft learns one distribution and is served another —
the same train/serve mismatch compute_logits closes on the output side.
Shared block-drafting prologue: anchors, noise embedding, positions, mask.
Centralises the sequence-packing handling for every block-wise trainer (DFlash and its Domino / JetSpec subclasses): anchors are sampled so each block stays inside one document, the block’s context prefix is restricted to its anchor’s document, and the draft RoPE uses per-document positions.
Under loss_type="variable_prefix" the block embedding shows a sampled
real-token prefix (_create_vp_noise_embed) instead of the single anchor
token, and the sampled prefix_lengths are returned for the loss; every
other loss_type returns prefix_lengths=None.
Parameters:
[B, S] context token ids (long).
[B, S] supervised-token mask.
[B, S] per-document reset positions (packing), or None.
[B, max_docs] packed document lengths, or None (unpacked).
[B, S] remaining real tokens of each position’s document.
When True, build the in-block-causal (JetSpec) mask instead of the bidirectional (DFlash / Domino) one.
Returns: torch.Tensor
“(anchor_positions [B, N], block_keep_mask [B, N], noise_embedding
Randomly sample anchor positions per sample; returns (anchors, keep_mask).
doc_remaining [B, S] (sequence packing) restricts anchors so the
whole block stays inside one document (doc_remaining >= block_size - 1),
the per-document analogue of the anchor <= seq_len - block_size bound.
This is required for correctness — _build_block_targets gathers labels
by absolute offset and does not encode document boundaries, so a block that
crossed one would be supervised on the next document’s tokens. A side effect
is that a packed document shorter than block_size yields no anchors (the
unpacked path still supervises such a short sequence’s partial block); pack
with documents at least block_size long to avoid dropping their signal.
Sample a visible-prefix length per block for variable-prefix training.
A prefix length l means block positions [0, l) are visible real
tokens and positions [l, block_size) are masked prediction targets.
Lengths follow D2SD’s truncated geometric prior Pr(l) ~ base ** l over
[min(2, block_size - 1), block_size - 1]; a base below 1 biases toward
short prefixes. The lower bound skips the degenerate fixed-anchor DFlash
case, and the upper bound keeps at least one masked target per block.
Returns: torch.Tensor
Long tensor of shape [batch, blocks].
Decay-weighted CE over each block’s masked suffix (D2SD Eq. for L_VP).
Only positions at or past the sampled visible prefix are supervised, and
the exponential decay restarts at the prefix boundary:
w_k = exp(-(k - l) / gamma) for block position k >= l
(loss_decay_gamma=None disables decay). The loss is the weighted mean
sum(nll * w) / sum(w), mirroring the fixed-anchor path’s
normalize="mean". Assumes every prefix_lengths entry is at least
min(2, block_size - 1) (what _sample_prefix_lengths produces),
so the leading always-visible positions can be sliced off before the CE.
Parameters:
Tensor of shape [batch, blocks, block_size, vocab].
Long tensor of shape [batch, blocks, block_size]; the
ground-truth token at anchor + k for block position k.
Tensor of shape [batch, blocks, block_size]; 0/1
product of block validity, in-bounds, and the loss mask.
Long tensor of shape [batch, blocks]; this block’s
visible-prefix length.
Returns: DFlashStepMetrics
DFlashStepMetrics with the weighted-mean loss, the argmax accuracy
Parallel block-wise training forward pass.
Sequence packing (position_ids [B, S] per-document reset positions,
seq_lens [B, max_docs] document lengths, doc_remaining [B, S])
keeps every block inside one document: anchors are constrained so the block
does not cross a boundary, the block’s context prefix attends only within the
anchor’s document, and the draft’s RoPE uses the per-document positions.
Per-step training outputs for the Domino draft.
loss/accuracy/valid_tokens mirror DFlashStepMetrics so the
DFlash training loop consumes them unchanged. The remaining fields are
diagnostics for the two supervised logits and the curriculum weight.
Scalar tensor containing final-head mean acceptance length.
Scalar tensor containing final-head additive acceptance length.
Scalar tensor containing final-head greedy token accuracy.
Scalar tensor containing base-head mean acceptance length.
Scalar tensor containing base-head additive acceptance length.
Scalar tensor containing base-head greedy token accuracy.
Scalar tensor containing base-head correct-token count.
Scalar tensor containing base-head validation loss.
Scalar tensor containing final-head correct-token count.
Scalar tensor containing final-head validation loss.
Scalar tensor containing the current base-loss curriculum weight.
Scalar tensor containing the mixed differentiable training loss.
Scalar tensor containing the effective loss denominator.
Scalar tensor containing the number of evaluated draft blocks.
Scalar tensor containing the supervised-token count.
Bases: DFlashTrainerModule
Domino online training wrapper: DFlash backbone + causal correction head.
First block position that receives the Domino correction.
With shift_label the block predicts anchor+1+k, so position 0 is
already a real next-token prediction; otherwise position 0 is the clean
anchor and is excluded. pure_draft_prefix_len reserves additional
leading positions for the backbone-only (uncorrected) base logits.
Add the GRU-conditioned low-rank correction to the suffix base logits.
Reshape backbone hidden states and gather the block’s previous tokens.
prev_ids[..., k] is the token that precedes block-position k’s
target — the input token at anchor+k when shift_label (the GRU
consumes ground-truth context), else the target sequence itself.
Diagnostics for both heads (acceptance length, base accuracy). No gradient.
Decay-weighted CE on both logits, mixed by the curriculum weight.
Parallel block-wise training forward with the Domino correction head.
Sequence packing (position_ids [B, S] per-document reset positions,
seq_lens [B, max_docs] document lengths, doc_remaining [B, S])
is handled by the shared DFlash prologue; shift_label labels reaching
one past the block (anchor + block_size) are truncated at the anchor’s
document boundary by _build_block_targets.
Bases: Module
Two-tap dynamic depthwise convolution wrapped around one draft sublayer.
One instance covers the convolution before a sublayer (prepare) and
the one after it (finish). Both sets of dynamic coefficients are
predicted from the sublayer’s input, so prepare returns the
coefficients finish needs and the projection runs once per sublayer.
Parameters:
Channel count of the draft hidden states.
Number of taps; DFlash 2 uses 2 (self and predecessor).
Channels sharing one dynamic correction (16 in DFlash 2).
Apply the post-sublayer convolution using the coefficients from prepare.
Parameters:
Tensor of shape [batch, sequence, hidden]; the sublayer output.
Tensor of shape [batch, sequence, kernel, groups]; the second
half of prepare’s coefficients.
Draft-block length; taps never cross a block boundary.
Returns: torch.Tensor
Tensor of shape [batch, sequence, hidden].
Apply the pre-sublayer convolution and emit the post-sublayer coefficients.
Parameters:
Tensor of shape [batch, sequence, hidden]; the sublayer input.
Draft-block length; taps never cross a block boundary.
Returns: torch.Tensor
Tuple of (convolved, dynamic): convolved is a Tensor of shape
Reset the convolution to the identity: unit self-tap, no correction.
Capture a set of decoder-layer hidden states from a frozen HF causal LM.
A forward hook on decoder layer i captures that layer’s output, which in
HuggingFace’s output_hidden_states convention is hidden_states[i + 1]
— matching SpecForge’s extract_context_feature (offset 1).
Return decoder layers as an ordered, integer-indexable list.
Run the target model and capture the selected layers’ hidden states as context.
With seq_lens ([B, max_docs], per-document lengths summing to S)
the target runs with a document-level block-causal mask and per-document
position_ids so the captured context hidden states do not leak across
document boundaries (SDPA/eager consume the [B, 1, S, S] block-causal
additive mask; FlashAttention infers boundaries from position_ids at batch
size 1). The packing metadata is carried through to the trainer unchanged.
Return the target model input embeddings.
Bases: Module
DFlash draft model: a small dense non-causal K3 MLA stack over [context | noise].
The draft owns no embedding table and no LM head: the DFlash trainer embeds
the [anchor, MASK, ...] blocks with the frozen target’s embed_tokens
and decodes this model’s output with the frozen target’s lm_head.
Project draft hidden states to logits, applying the target’s output transform.
See Qwen3DFlashDraftModel.compute_logits — same contract, identity
unless the target’s dflash_config sets output_multiplier /
final_logit_softcapping.
Parameters:
Tensor of shape […, hidden]; draft hidden states with arbitrary leading dimensions.
The frozen target’s output projection.
Returns: torch.Tensor
Tensor of shape […, vocab].
Predict the draft blocks’ hidden states.
Parameters:
Unused. K3’s MLA is NoPE, so the draft has no rotary embedding; the argument is accepted to keep the trainer’s call signature identical across DFlash drafts.
Additive DFlash mask of shape
[batch, 1, blocks * block_size, sequence + blocks * block_size].
Tensor of shape [batch, blocks * block_size, hidden].
Tensor of shape
[batch, sequence, len(target_layer_ids) * hidden].
Ignored.
Returns: torch.Tensor
Tensor of shape [batch, blocks * block_size, hidden].
Bases: ValueError
Raised when a batch has no sample long enough to form a DFlash block.
A DFlash anchor is a supervised position (loss_mask=1) that still has a
full block ahead of it, i.e. one in [0, seq_len - block_size] (and, when
packing, with at least block_size - 1 further real tokens in its own
document). Datasets always contain some short conversations; the training
loop catches this and skips the offending micro-batch rather than aborting
the run.
Bases: Qwen3DFlashDraftModel
DFlash 2 draft model: the DFlash stack plus in-block convs and a path selector.
Run the DFlash 2 draft stack over [context | noise-block].
Parameters:
Long tensor of shape [batch, context + draft].
Attention mask over [batch, 1, draft, context + draft],
a flex BlockMask, or None.
Tensor of shape [batch, draft, hidden]; the embedded
[anchor, MASK, ...] blocks laid end to end.
Tensor of shape [batch, context, layers * hidden]; the concatenated target-model context features.
Draft KV cache, or None.
Whether to write into past_key_values.
Draft-block length for the in-block convolutions; see
resolve_conv_block_size for how None is resolved.
Forwarded to the attention implementation.
Returns: torch.Tensor
Tensor of shape [batch, draft, hidden]; the normalised draft hidden
Resolve the block length the in-block convolutions must not reach across.
Parameters:
Number of draft (noise-block) query positions in this call.
Explicit block length, or None to infer it.
Returns: int
The block length to convolve within. None resolves to
Raises:
ValueError: Ifconv_block_sizedoes not dividequery_len.
Block-parallel speculative decoding with pairwise path selection.
Each cycle drafts one block in a single draft forward, walks the selector
over the per-position candidates to pick a coherent path, and verifies the
whole block with one target forward. Below GREEDY_TEMPERATURE_EPS this
accepts the longest exact-match prefix, matching sample’s own greedy
threshold; above it, it accepts via rejection sampling, so the emitted
tokens follow the target’s own distribution.
Parameters:
The frozen verifier; must expose model.embed_tokens,
lm_head, and an HF-style forward with output_hidden_states.
Long tensor of shape [1, prompt].
Maximum number of tokens to generate.
Token ids that end generation, or None.
Sampling temperature; 0 decodes greedily.
Nucleus mass to keep, in (0, 1]; truncates the target’s
distribution, which is what the emitted tokens must follow. The
draft proposes from plain temperature, as the reference does.
Candidates to keep, 0 for the whole vocabulary.
Returns: torch.LongTensor
Long tensor of shape [1, prompt + generated] containing the prompt
Bases: Qwen3PreTrainedModel
DFlash draft model: a small non-causal Qwen3 stack over [context | noise].
Keep the RoPE inv_freq buffer in fp32 across dtype casts.
Qwen3RotaryEmbedding computes the rotary angles in fp32 but reads the
frequencies from a stored inv_freq buffer. model.to(bfloat16) — the
training build path — rounds that buffer to bf16, whereas the serving
runtime (SGLang keeps an fp32 RoPE cache) and HF’s from_pretrained reload
keep it in fp32. The resulting train/inference RoPE mismatch grows with
absolute position (the bf16 frequencies dephase) and erodes draft
acceptance, so inv_freq must stay fp32 on both the training and reload
paths. A bf16 round-trip cannot be undone by upcasting, so when a cast
rounds the buffer we recompute fresh fp32 frequencies from the rotary
config (the same values HF derives on the fp32 paths) instead of upcasting
the corrupted ones.
Project draft hidden states to logits, applying the target’s output transform.
Some targets do not read their lm_head output raw: Muse Glimmer scales
it by output_multiplier and squashes it through
final_logit_softcapping. The draft is trained against, and verified
by, those transformed logits, so both training and decoding have to apply
them — otherwise the draft learns one distribution and is served another.
Both fields live in dflash_config; absent (every Qwen3-family target),
this is just the head.
Parameters:
Tensor of shape […, hidden]; draft hidden states with arbitrary leading dimensions.
The frozen target’s output projection.
Returns: torch.Tensor
Tensor of shape […, vocab].
Embed a draft block with the target’s (frozen, shared) input embedding.
Applies input_embedding_scale from dflash_config — 1.0, hence a
no-op, for every Qwen3-family target. Unlike the reference, this calls the
embedding module rather than indexing its weight table directly: under
tensor parallelism the target’s embed_tokens is vocab-parallel and only
the module call carries the DTensor plan.
Parameters:
The frozen verifier.
Long tensor of shape [batch, draft]; the block’s token ids,
position 0 holding the verified anchor and the rest MASK.
Returns: torch.Tensor
Tensor of shape [batch, draft, hidden].
Block-parallel speculative decoding: draft a block, verify with the target, accept the matching prefix.
top_p / top_k truncate the target’s distribution, which is what
defines the emitted tokens; the draft proposes from plain temperature.
Parameters:
The frozen verifier.
Long tensor of shape [1, prompt].
Maximum number of tokens to generate.
Token ids that end generation, or None.
Sampling temperature; 0 decodes greedily.
Nucleus mass to keep, in (0, 1].
Candidates to keep, 0 for the whole vocabulary.
Returns: torch.LongTensor
Long tensor of shape [1, prompt + generated].
Pick num_draft_layers target layers spread across the target’s depth.
Count consecutive accepted predictions for every draft block.
Parameters:
Long tensor of shape [batch, blocks, depth].
Long tensor of shape [batch, blocks, depth].
Bool tensor of shape [batch, blocks, depth].
Returns: torch.Tensor
Float tensor of shape [batch, blocks] containing the accepted draft
Build a sparse FlexAttention BlockMask for DFlash training.
See module docstring for the mask semantics. The returned BlockMask is
consumed directly by transformers’ flex_attention backend when
_attn_implementation="flex_attention" is set on the draft model — pass
it via the attention_mask kwarg.
Parameters:
[B, N] anchor positions (long).
[B, N] valid-anchor mask (bool).
context length.
block size.
torch device.
Cache and reuse a torch.compile’d create_block_mask
across calls (default True). Set to False when running on PyTorch
builds that hit Inductor errors during compile.
When True, make in-block attention causal (JetSpec); otherwise it is bidirectional (DFlash).
[B, S] long per-context-token document id (sequence
packing), or None to disable the per-document constraint.
[B, N] long document id of each anchor. Required when
ctx_doc_id is given; a block then attends only to context tokens in
its anchor’s document.
When set, a block additionally attends only to context
tokens within sliding_window positions of the query’s own sequence
position (anchor + offset), matching the published DFlash drafters’
sliding_attention layers. None (default) keeps the full prefix.
Returns: 'BlockMask'
class:torch.nn.attention.flex_attention.BlockMask.
Build a dense additive attention mask for the SDPA backend.
Parameters:
[B, N] anchor positions per sample (long).
[B, N] per-sample valid-anchor mask (bool).
context length S.
block size.
torch device.
dtype for the additive mask (typically the model dtype).
When True, make in-block attention causal (JetSpec); otherwise it is bidirectional (DFlash).
[B, S] long per-context-token document id (sequence
packing), or None to disable the per-document constraint.
[B, N] long document id of each anchor. Required when
ctx_doc_id is given; a block then attends only to context tokens in
its anchor’s document.
When set, a block additionally attends only to context
tokens within sliding_window positions of the query’s own sequence
position (anchor + offset), matching the published DFlash drafters’
sliding_attention layers. None (default) keeps the full prefix.
Returns: torch.Tensor
[B, 1, N*block_size, S + N*block_size] float tensor: 0 at
Concatenate the selected target layers’ hidden states along the feature dim.
hidden_states follows HF’s output_hidden_states convention where
index 0 is the embedding output, so layer i’s output is at index
i + 1.
Base-anchor curriculum weight at global_step.
lambda_base starts at lambda_start and decays linearly to 0 over the
first decay_ratio fraction of total_steps, then stays at 0. The result
is clamped to [0, 1].
Return the first registered DFlash draft spec matching any architecture in the list.