nemo_automodel.components.speculative.dflash

View as Markdown

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

Package Contents

Classes

NameDescription
CandidateSelectorPairwise path selector over each draft position’s top-k candidates.
DFlash2StepMetricsPer-step training outputs for the DFlash 2 draft.
DFlash2TrainerModuleDFlash 2 online training wrapper: DFlash block CE + candidate-selection CE.
DFlashDraftSpecHow to build a DFlash draft model for a particular target architecture.
DFlashStepMetricsPer-step training outputs for the DFlash draft.
DFlashTargetBatchTarget-model context features needed by the DFlash trainer.
DFlashTrainerModuleDFlash online training wrapper with block-wise CE loss.
DominoStepMetricsPer-step training outputs for the Domino draft.
DominoTrainerModuleDomino online training wrapper: DFlash backbone + causal correction head.
GroupedDynamicCausalConvTwo-tap dynamic depthwise convolution wrapped around one draft sublayer.
HFDFlashTargetModelCapture a set of decoder-layer hidden states from a frozen HF causal LM.
KimiK3DFlashDraftModelDFlash draft model: a small dense non-causal K3 MLA stack over [context | noise].
NoValidAnchorsErrorRaised when a batch has no sample long enough to form a DFlash block.
Qwen3DFlash2DraftModelDFlash 2 draft model: the DFlash stack plus in-block convs and a path selector.
Qwen3DFlashDraftModelDFlash draft model: a small non-causal Qwen3 stack over [context | noise].

Functions

NameDescription
build_target_layer_idsPick num_draft_layers target layers spread across the target’s depth.
compute_accept_lenCount consecutive accepted predictions for every draft block.
create_dflash_block_maskBuild a sparse FlexAttention BlockMask for DFlash training.
create_dflash_sdpa_maskBuild a dense additive attention mask for the SDPA backend.
extract_context_featureConcatenate the selected target layers’ hidden states along the feature dim.
get_lambda_baseBase-anchor curriculum weight at global_step.
resolve_dflash_draft_specReturn the first registered DFlash draft spec matching any architecture in the list.

Data

DFLASH_DRAFT_REGISTRY

API

class nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2.CandidateSelector(
vocab_size: int,
hidden_size: int,
rank: int,
top_k: int
)

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:

vocab_size
int

Size of the (shared with the target) token vocabulary.

hidden_size
int

Channel count of the draft hidden states.

rank
int

Codebook / gate width (selector_rank; 256 in DFlash 2).

top_k
int

Candidates kept per position (selector_top_k; 16 in DFlash 2).

hidden_projection
= nn.Linear(hidden_size, rank, bias=False)
predecessor_codebook
= nn.Parameter(torch.empty(vocab_size, rank))
successor_codebook
= nn.Parameter(torch.empty(vocab_size, rank))
nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2.CandidateSelector.pair_scores(
hidden: torch.Tensor,
unary: torch.Tensor,
candidate_ids: torch.Tensor,
predecessor_ids: torch.Tensor
) -> torch.Tensor

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:

hidden
torch.Tensor

Tensor of shape […, hidden]; the draft hidden state at each scored position, with arbitrary leading dimensions.

unary
torch.Tensor

Tensor of shape […, candidates]; U_t(b), the draft logit of each candidate, with the same leading dimensions as hidden.

candidate_ids
torch.Tensor

Long tensor of shape […, candidates]; the candidate token ids, with the same leading dimensions as hidden.

predecessor_ids
torch.Tensor

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

nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2.CandidateSelector.reset_parameters() -> None

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.

nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2.CandidateSelector.walk(
hidden: torch.Tensor,
logits: torch.Tensor,
anchor_ids: torch.Tensor,
temperature: float
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]

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:

hidden
torch.Tensor

Tensor of shape [batch, draft, hidden]; the draft hidden states of the block’s predicted positions.

logits
torch.Tensor

Tensor of shape [batch, draft, vocab]; the draft logits at those positions.

anchor_ids
torch.Tensor

Long tensor of shape [batch]; the last verified token, i.e. the predecessor of draft position 0.

temperature
float

Sampling temperature; below GREEDY_TEMPERATURE_EPS selects greedily.

Returns: torch.Tensor

Tuple (path, candidate_ids, draft_probs): path is a Long tensor

class nemo_automodel.components.speculative.dflash.dflash2_core.DFlash2StepMetrics(
loss: torch.Tensor,
loss_weight: torch.Tensor,
accuracy: torch.Tensor,
valid_tokens: torch.Tensor,
correct_tokens: torch.Tensor,
accept_len: torch.Tensor,
accept_len_sum: torch.Tensor,
valid_blocks: torch.Tensor,
base_loss: torch.Tensor,
selector_loss: torch.Tensor,
base_accuracy: torch.Tensor,
base_correct_tokens: torch.Tensor,
base_accept_len: torch.Tensor,
base_accept_len_sum: torch.Tensor,
candidate_recall: torch.Tensor
)
Dataclass

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.

accept_len
Tensor

Scalar tensor containing selector-path mean acceptance length.

accept_len_sum
Tensor

Scalar tensor containing selector-path additive acceptance length.

accuracy
Tensor

Scalar tensor containing selector-path greedy token accuracy.

base_accept_len
Tensor

Scalar tensor containing backbone mean acceptance length.

base_accept_len_sum
Tensor

Scalar tensor containing backbone additive acceptance length.

base_accuracy
Tensor

Scalar tensor containing backbone top-1 token accuracy.

base_correct_tokens
Tensor

Scalar tensor containing backbone correct-token count.

base_loss
Tensor

Scalar tensor containing the backbone block-CE term.

candidate_recall
Tensor

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.

correct_tokens
Tensor

Scalar tensor containing selector-path correct-token count.

loss
Tensor

Scalar tensor containing the differentiable training loss.

loss_weight
Tensor

Scalar tensor containing the effective loss denominator.

selector_loss
Tensor

Scalar tensor containing the candidate-selection CE term.

valid_blocks
Tensor

Scalar tensor containing the number of evaluated draft blocks.

valid_tokens
Tensor

Scalar tensor containing the supervised-token count.

class nemo_automodel.components.speculative.dflash.dflash2_core.DFlash2TrainerModule(
target_lm_head: torch.nn.Module,
target_embed_tokens: torch.nn.Module,
mask_token_id: int,
block_size: int = 16,
attention_backend: str = 'flex_attention',
num_anchors: int = 512,
loss_decay_gamma: float | None = None,
selector_loss_weight: float = 1.0,
sliding_window: int | None = None
)

Bases: DFlashTrainerModule

DFlash 2 online training wrapper: DFlash block CE + candidate-selection CE.

selector_loss_weight
= float(selector_loss_weight)
nemo_automodel.components.speculative.dflash.dflash2_core.DFlash2TrainerModule._depth_weights(
mask: torch.Tensor
) -> torch.Tensor

Block-position decay weights for the block_size - 1 predicted positions.

Parameters:

mask
torch.Tensor

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

nemo_automodel.components.speculative.dflash.dflash2_core.DFlash2TrainerModule._selector_scores(
hidden: torch.Tensor,
logits: torch.Tensor,
target_ids: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]

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:

hidden
torch.Tensor

Tensor of shape [batch, blocks, depth, hidden]; the draft hidden states of the predicted (non-anchor) block positions.

logits
torch.Tensor

Tensor of shape [batch, blocks, depth, vocab]; the backbone logits at those positions.

target_ids
torch.Tensor

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

nemo_automodel.components.speculative.dflash.dflash2_core.DFlash2TrainerModule.forward(
input_ids: torch.Tensor,
hidden_states: torch.Tensor,
loss_mask: torch.Tensor,
position_ids: torch.Tensor | None = None,
seq_lens: torch.Tensor | None = None,
doc_remaining: torch.Tensor | None = None

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:

input_ids
torch.Tensor

Long tensor of shape [batch, sequence]; the context tokens.

hidden_states
torch.Tensor

Tensor of shape [batch, sequence, layers * hidden]; the captured target-model context features.

loss_mask
torch.Tensor

Tensor of shape [batch, sequence]; the supervised-token mask.

position_ids
torch.Tensor | NoneDefaults to None

Long tensor of shape [batch, sequence] with per-document reset positions under packing, or None.

seq_lens
torch.Tensor | NoneDefaults to None

Long tensor of shape [batch, max_docs] with packed document lengths, or None when unpacked.

doc_remaining
torch.Tensor | NoneDefaults to None

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.

class nemo_automodel.components.speculative.dflash.registry.DFlashDraftSpec(
draft_cls: type[torch.nn.Module],
build_draft_config: typing.Callable[..., transformers.PretrainedConfig] | None = None,
draft2_cls: type[torch.nn.Module] | None = None,
build_target_kwargs: typing.Callable[[Any], dict] = _no_target_kwargs,
attention_backends: tuple[str, ...] = ('flex_attention', 'sdpa', ...,
supports_context_parallel: bool = True
)
Dataclass

How to build a DFlash draft model for a particular target architecture.

attention_backends
tuple[str, ...] = ('flex_attention', 'sdpa', 'eager')

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

build_draft_config
Callable[..., PretrainedConfig] | None = None

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.

build_target_kwargs
Callable[[Any], dict] = _no_target_kwargs

Extra from_pretrained keyword arguments for the frozen target, derived from the recipe’s recipe_args.

draft2_cls
type[Module] | None = None

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.

draft_cls
type[Module]

The draft model class.

supports_context_parallel
bool = True

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.

class nemo_automodel.components.speculative.dflash.core.DFlashStepMetrics(
loss: torch.Tensor,
loss_weight: torch.Tensor,
accuracy: torch.Tensor,
valid_tokens: torch.Tensor,
correct_tokens: torch.Tensor,
accept_len: torch.Tensor,
accept_len_sum: torch.Tensor,
valid_blocks: torch.Tensor
)
Dataclass

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.

accept_len
Tensor

Scalar tensor containing mean bonus-token-inclusive acceptance length.

accept_len_sum
Tensor

Scalar tensor containing the additive acceptance-length sum.

accuracy
Tensor

Scalar tensor containing greedy token accuracy.

correct_tokens
Tensor

Scalar tensor containing the greedy-correct token count.

loss
Tensor

Scalar tensor containing the differentiable training loss.

loss_weight
Tensor

Scalar tensor containing the effective loss denominator.

valid_blocks
Tensor

Scalar tensor containing the number of evaluated draft blocks.

valid_tokens
Tensor

Scalar tensor containing the supervised-token count.

class nemo_automodel.components.speculative.dflash.target.DFlashTargetBatch(
hidden_states: torch.Tensor,
input_ids: torch.Tensor,
attention_mask: torch.Tensor,
loss_mask: torch.Tensor,
logits: torch.Tensor | None = None,
position_ids: torch.Tensor | None = None,
seq_lens: torch.Tensor | None = None,
doc_remaining: torch.Tensor | None = None
)
Dataclass

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.

attention_mask
Tensor
doc_remaining
Tensor | None = None
hidden_states
Tensor
input_ids
Tensor
logits
Tensor | None = None
loss_mask
Tensor
position_ids
Tensor | None = None
seq_lens
Tensor | None = None
class nemo_automodel.components.speculative.dflash.core.DFlashTrainerModule(
target_lm_head: torch.nn.Module,
target_embed_tokens: torch.nn.Module,
mask_token_id: int,
block_size: int = 16,
attention_backend: str = 'flex_attention',
num_anchors: int = 512,
loss_decay_gamma: float | None = None,
loss_type: str = 'dflash',
prefix_weight_base: float = 0.9,
sliding_window: int | None = None
)

Bases: Module

DFlash online training wrapper with block-wise CE loss.

_min_prefix
= min(2, block_size - 1)
loss_fn
prefix_weight_base
= float(prefix_weight_base)
nemo_automodel.components.speculative.dflash.core.DFlashTrainerModule._build_block_targets(
input_ids: torch.Tensor,
loss_mask: torch.Tensor,
anchor_positions: torch.Tensor,
block_keep_mask: torch.Tensor,
seq_len: int,
label_start: int = 0,
doc_remaining: torch.Tensor | None = None
) -> typing.Tuple[torch.Tensor, torch.Tensor, torch.Tensor]

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.

nemo_automodel.components.speculative.dflash.core.DFlashTrainerModule._create_noise_embed(
input_ids,
anchor_positions,
block_keep_mask
)

Embed each block as [anchor_token, MASK, MASK, ...] (invalid blocks all MASK).

nemo_automodel.components.speculative.dflash.core.DFlashTrainerModule._create_position_ids(
anchor_positions: torch.Tensor,
context_position_ids: torch.Tensor | None = None
) -> torch.Tensor

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.

nemo_automodel.components.speculative.dflash.core.DFlashTrainerModule._create_vp_noise_embed(
input_ids,
anchor_positions,
block_keep_mask,
prefix_lengths
)

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:

input_ids

Long tensor of shape [batch, sequence].

anchor_positions

Long tensor of shape [batch, blocks]; each block’s start position in the sequence.

block_keep_mask

Bool tensor of shape [batch, blocks]; invalid (padding) blocks are embedded as all MASK.

prefix_lengths

Long tensor of shape [batch, blocks]; this block’s visible-prefix length.

Returns:

Tensor of shape [batch, blocks * block_size, hidden].

nemo_automodel.components.speculative.dflash.core.DFlashTrainerModule._embed_noise_ids(
noise_ids: torch.Tensor
) -> torch.Tensor

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.

nemo_automodel.components.speculative.dflash.core.DFlashTrainerModule._prepare_block_inputs(
input_ids: torch.Tensor,
loss_mask: torch.Tensor,
position_ids: torch.Tensor | None = None,
seq_lens: torch.Tensor | None = None,
doc_remaining: torch.Tensor | None = None,
causal: bool = False
) -> typing.Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, 'torch.Tensor | BlockMask', torch.Tensor | None]

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:

input_ids
torch.Tensor

[B, S] context token ids (long).

loss_mask
torch.Tensor

[B, S] supervised-token mask.

position_ids
torch.Tensor | NoneDefaults to None

[B, S] per-document reset positions (packing), or None.

seq_lens
torch.Tensor | NoneDefaults to None

[B, max_docs] packed document lengths, or None (unpacked).

doc_remaining
torch.Tensor | NoneDefaults to None

[B, S] remaining real tokens of each position’s document.

causal
boolDefaults to False

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

nemo_automodel.components.speculative.dflash.core.DFlashTrainerModule._sample_anchor_positions(
seq_len: int,
loss_mask: torch.Tensor,
device: torch.device,
doc_remaining: torch.Tensor | None = None
) -> typing.Tuple[torch.Tensor, torch.Tensor]

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.

nemo_automodel.components.speculative.dflash.core.DFlashTrainerModule._sample_prefix_lengths(
bsz: int,
n_blocks: int,
device: torch.device
) -> torch.Tensor

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

nemo_automodel.components.speculative.dflash.core.DFlashTrainerModule._variable_prefix_loss(
logits: torch.Tensor,
target_ids: torch.Tensor,
block_mask: torch.Tensor,
prefix_lengths: torch.Tensor

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:

logits
torch.Tensor

Tensor of shape [batch, blocks, block_size, vocab].

target_ids
torch.Tensor

Long tensor of shape [batch, blocks, block_size]; the ground-truth token at anchor + k for block position k.

block_mask
torch.Tensor

Tensor of shape [batch, blocks, block_size]; 0/1 product of block validity, in-bounds, and the loss mask.

prefix_lengths
torch.Tensor

Long tensor of shape [batch, blocks]; this block’s visible-prefix length.

Returns: DFlashStepMetrics

DFlashStepMetrics with the weighted-mean loss, the argmax accuracy

nemo_automodel.components.speculative.dflash.core.DFlashTrainerModule.forward(
input_ids: torch.Tensor,
hidden_states: torch.Tensor,
loss_mask: torch.Tensor,
position_ids: torch.Tensor | None = None,
seq_lens: torch.Tensor | None = None,
doc_remaining: torch.Tensor | None = None

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.

class nemo_automodel.components.speculative.dflash.domino_core.DominoStepMetrics(
loss: torch.Tensor,
loss_weight: torch.Tensor,
accuracy: torch.Tensor,
valid_tokens: torch.Tensor,
correct_tokens: torch.Tensor,
accept_len: torch.Tensor,
accept_len_sum: torch.Tensor,
valid_blocks: torch.Tensor,
final_loss: torch.Tensor,
base_loss: torch.Tensor,
base_accuracy: torch.Tensor,
base_correct_tokens: torch.Tensor,
base_accept_len: torch.Tensor,
base_accept_len_sum: torch.Tensor,
lambda_base: torch.Tensor
)
Dataclass

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.

accept_len
Tensor

Scalar tensor containing final-head mean acceptance length.

accept_len_sum
Tensor

Scalar tensor containing final-head additive acceptance length.

accuracy
Tensor

Scalar tensor containing final-head greedy token accuracy.

base_accept_len
Tensor

Scalar tensor containing base-head mean acceptance length.

base_accept_len_sum
Tensor

Scalar tensor containing base-head additive acceptance length.

base_accuracy
Tensor

Scalar tensor containing base-head greedy token accuracy.

base_correct_tokens
Tensor

Scalar tensor containing base-head correct-token count.

base_loss
Tensor

Scalar tensor containing base-head validation loss.

correct_tokens
Tensor

Scalar tensor containing final-head correct-token count.

final_loss
Tensor

Scalar tensor containing final-head validation loss.

lambda_base
Tensor

Scalar tensor containing the current base-loss curriculum weight.

loss
Tensor

Scalar tensor containing the mixed differentiable training loss.

loss_weight
Tensor

Scalar tensor containing the effective loss denominator.

valid_blocks
Tensor

Scalar tensor containing the number of evaluated draft blocks.

valid_tokens
Tensor

Scalar tensor containing the supervised-token count.

class nemo_automodel.components.speculative.dflash.domino_core.DominoTrainerModule(
target_lm_head: torch.nn.Module,
target_embed_tokens: torch.nn.Module,
mask_token_id: int,
block_size: int = 16,
attention_backend: str = 'flex_attention',
num_anchors: int = 512,
loss_decay_gamma: float | None = None,
shift_label: bool = False
)

Bases: DFlashTrainerModule

Domino online training wrapper: DFlash backbone + causal correction head.

_suffix_start
int

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.

nemo_automodel.components.speculative.dflash.domino_core.DominoTrainerModule._apply_domino_head(
base_logits4d: torch.Tensor,
hidden4d: torch.Tensor,
prev_ids: torch.Tensor,
target_ids: torch.Tensor
) -> torch.Tensor

Add the GRU-conditioned low-rank correction to the suffix base logits.

nemo_automodel.components.speculative.dflash.domino_core.DominoTrainerModule._build_domino_head_inputs(
input_ids: torch.Tensor,
anchor_positions: torch.Tensor,
target_ids: torch.Tensor,
output_hidden: torch.Tensor
) -> typing.Tuple[torch.Tensor, torch.Tensor]

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.

nemo_automodel.components.speculative.dflash.domino_core.DominoTrainerModule._compute_extra_metrics(
pred_ids: torch.Tensor,
flat_base_logits: torch.Tensor,
flat_targets: torch.Tensor,
binary_eval_mask: torch.Tensor,
target_ids: torch.Tensor,
eval_weight_mask: torch.Tensor,
final_loss: torch.Tensor,
base_loss: torch.Tensor,
lambda_base: float
) -> typing.Dict[str, torch.Tensor]

Diagnostics for both heads (acceptance length, base accuracy). No gradient.

nemo_automodel.components.speculative.dflash.domino_core.DominoTrainerModule._compute_weighted_losses(
final_logits: torch.Tensor,
base_logits: torch.Tensor,
target_ids: torch.Tensor,
weight_mask: torch.Tensor,
lambda_base: float
) -> typing.Tuple[torch.Tensor, torch.Tensor, torch.Tensor]

Decay-weighted CE on both logits, mixed by the curriculum weight.

nemo_automodel.components.speculative.dflash.domino_core.DominoTrainerModule.forward(
input_ids: torch.Tensor,
hidden_states: torch.Tensor,
loss_mask: torch.Tensor,
lambda_base: float = 0.0,
position_ids: torch.Tensor | None = None,
seq_lens: torch.Tensor | None = None,
doc_remaining: torch.Tensor | None = None

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.

class nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2.GroupedDynamicCausalConv(
hidden_size: int,
kernel_size: int,
group_size: int
)

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:

hidden_size
int

Channel count of the draft hidden states.

kernel_size
int

Number of taps; DFlash 2 uses 2 (self and predecessor).

group_size
int

Channels sharing one dynamic correction (16 in DFlash 2).

base_kernel
kernel_projection
nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2.GroupedDynamicCausalConv.finish(
hidden: torch.Tensor,
dynamic: torch.Tensor,
block_size: int
) -> torch.Tensor

Apply the post-sublayer convolution using the coefficients from prepare.

Parameters:

hidden
torch.Tensor

Tensor of shape [batch, sequence, hidden]; the sublayer output.

dynamic
torch.Tensor

Tensor of shape [batch, sequence, kernel, groups]; the second half of prepare’s coefficients.

block_size
int

Draft-block length; taps never cross a block boundary.

Returns: torch.Tensor

Tensor of shape [batch, sequence, hidden].

nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2.GroupedDynamicCausalConv.prepare(
hidden: torch.Tensor,
block_size: int
) -> tuple[torch.Tensor, torch.Tensor]

Apply the pre-sublayer convolution and emit the post-sublayer coefficients.

Parameters:

hidden
torch.Tensor

Tensor of shape [batch, sequence, hidden]; the sublayer input.

block_size
int

Draft-block length; taps never cross a block boundary.

Returns: torch.Tensor

Tuple of (convolved, dynamic): convolved is a Tensor of shape

nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2.GroupedDynamicCausalConv.reset_parameters() -> None

Reset the convolution to the identity: unit self-tap, no correction.

class nemo_automodel.components.speculative.dflash.target.HFDFlashTargetModel(
model: torch.nn.Module,
target_layer_ids: typing.Sequence[int],
capture_logits: bool = False,
cp_mesh = None
)

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

_cp_size
= cp_mesh.size() if cp_mesh is not None else 1
capture_logits
= bool(capture_logits)
model
= model.eval()
target_layer_ids
= self._validate_layer_ids(target_layer_ids)
nemo_automodel.components.speculative.dflash.target.HFDFlashTargetModel._check_captured(
captured: dict[int, torch.Tensor]
) -> None
nemo_automodel.components.speculative.dflash.target.HFDFlashTargetModel._get_transformer_layers() -> list[torch.nn.Module]

Return decoder layers as an ordered, integer-indexable list.

nemo_automodel.components.speculative.dflash.target.HFDFlashTargetModel._validate_layer_ids(
target_layer_ids: typing.Sequence[int]
) -> list[int]
nemo_automodel.components.speculative.dflash.target.HFDFlashTargetModel.generate_batch(
input_ids: torch.Tensor,
attention_mask: torch.Tensor,
loss_mask: torch.Tensor,
position_ids: torch.Tensor | None = None,
seq_lens: torch.Tensor | None = None,
doc_remaining: torch.Tensor | None = None

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.

nemo_automodel.components.speculative.dflash.target.HFDFlashTargetModel.get_input_embeddings() -> torch.nn.Embedding

Return the target model input embeddings.

class nemo_automodel.components.speculative.dflash.draft_kimi_k3.KimiK3DFlashDraftModel(
)

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.

_no_split_modules
= ['KimiK3DFlashDecoderLayer']
fc
final_logit_softcapping
= None if softcap is None else float(softcap)
hidden_norm
input_embedding_scale
layers
norm
target_layer_ids
nemo_automodel.components.speculative.dflash.draft_kimi_k3.KimiK3DFlashDraftModel.compute_logits(
hidden: torch.Tensor,
output_head: torch.nn.Module
) -> torch.Tensor

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:

hidden
torch.Tensor

Tensor of shape […, hidden]; draft hidden states with arbitrary leading dimensions.

output_head
nn.Module

The frozen target’s output projection.

Returns: torch.Tensor

Tensor of shape […, vocab].

nemo_automodel.components.speculative.dflash.draft_kimi_k3.KimiK3DFlashDraftModel.forward(
position_ids: torch.LongTensor | None = None,
attention_mask: torch.Tensor | None = None,
noise_embedding: torch.Tensor | None = None,
target_hidden: torch.Tensor | None = None,
kwargs: typing.Any = {}
) -> torch.Tensor

Predict the draft blocks’ hidden states.

Parameters:

position_ids
torch.LongTensor | NoneDefaults to None

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.

attention_mask
torch.Tensor | NoneDefaults to None

Additive DFlash mask of shape [batch, 1, blocks * block_size, sequence + blocks * block_size].

noise_embedding
torch.Tensor | NoneDefaults to None

Tensor of shape [batch, blocks * block_size, hidden].

target_hidden
torch.Tensor | NoneDefaults to None

Tensor of shape [batch, sequence, len(target_layer_ids) * hidden].

**kwargs
AnyDefaults to {}

Ignored.

Returns: torch.Tensor

Tensor of shape [batch, blocks * block_size, hidden].

class nemo_automodel.components.speculative.dflash.core.NoValidAnchorsError()

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.

class nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2.Qwen3DFlash2DraftModel(
config
)

Bases: Qwen3DFlashDraftModel

DFlash 2 draft model: the DFlash stack plus in-block convs and a path selector.

_no_split_modules
= ['Qwen3DFlash2DecoderLayer']
candidate_selector
nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2.Qwen3DFlash2DraftModel.forward(
position_ids: torch.LongTensor,
attention_mask: torch.Tensor | None = None,
noise_embedding: torch.Tensor | None = None,
target_hidden: torch.Tensor | None = None,
past_key_values: transformers.cache_utils.Cache | None = None,
use_cache: bool = False,
conv_block_size: int | None = None,
kwargs = {}
) -> torch.Tensor

Run the DFlash 2 draft stack over [context | noise-block].

Parameters:

position_ids
torch.LongTensor

Long tensor of shape [batch, context + draft].

attention_mask
torch.Tensor | NoneDefaults to None

Attention mask over [batch, 1, draft, context + draft], a flex BlockMask, or None.

noise_embedding
torch.Tensor | NoneDefaults to None

Tensor of shape [batch, draft, hidden]; the embedded [anchor, MASK, ...] blocks laid end to end.

target_hidden
torch.Tensor | NoneDefaults to None

Tensor of shape [batch, context, layers * hidden]; the concatenated target-model context features.

past_key_values
Cache | NoneDefaults to None

Draft KV cache, or None.

use_cache
boolDefaults to False

Whether to write into past_key_values.

conv_block_size
int | NoneDefaults to None

Draft-block length for the in-block convolutions; see resolve_conv_block_size for how None is resolved.

**kwargs
Defaults to {}

Forwarded to the attention implementation.

Returns: torch.Tensor

Tensor of shape [batch, draft, hidden]; the normalised draft hidden

nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2.Qwen3DFlash2DraftModel.resolve_conv_block_size(
query_len: int,
conv_block_size: int | None
) -> int

Resolve the block length the in-block convolutions must not reach across.

Parameters:

query_len
int

Number of draft (noise-block) query positions in this call.

conv_block_size
int | None

Explicit block length, or None to infer it.

Returns: int

The block length to convolve within. None resolves to

Raises:

  • ValueError: If conv_block_size does not divide query_len.
nemo_automodel.components.speculative.dflash.draft_qwen3_dflash2.Qwen3DFlash2DraftModel.spec_generate(
target: torch.nn.Module,
input_ids: torch.LongTensor,
max_new_tokens: int,
stop_token_ids: list[int] | None,
temperature: float,
top_p: float = 1.0,
top_k: int = 0
) -> torch.LongTensor

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:

target
nn.Module

The frozen verifier; must expose model.embed_tokens, lm_head, and an HF-style forward with output_hidden_states.

input_ids
torch.LongTensor

Long tensor of shape [1, prompt].

max_new_tokens
int

Maximum number of tokens to generate.

stop_token_ids
list[int] | None

Token ids that end generation, or None.

temperature
float

Sampling temperature; 0 decodes greedily.

top_p
floatDefaults to 1.0

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.

top_k
intDefaults to 0

Candidates to keep, 0 for the whole vocabulary.

Returns: torch.LongTensor

Long tensor of shape [1, prompt + generated] containing the prompt

class nemo_automodel.components.speculative.dflash.draft_qwen3.Qwen3DFlashDraftModel(
config
)

Bases: Qwen3PreTrainedModel

DFlash draft model: a small non-causal Qwen3 stack over [context | noise].

_no_split_modules
= ['Qwen3DFlashDecoderLayer']
block_size
decoder_layer_cls
type[Qwen3DFlashDecoderLayer] = Qwen3DFlashDecoderLayer
emb_dim
= dflash_config['emb_dim']
embed_proj
fc
final_logit_softcapping
= None if softcap is None else float(softcap)
gru_hidden_dim
= dflash_config['gru_hidden_dim']
hidden_norm
input_embedding_scale
layers
mask_token_id
= dflash_config.get('mask_token_id', None)
norm
prefix_gru
projector_type
= dflash_config.get('projector_type', None)
pure_draft_prefix_len
= dflash_config.get('pure_draft_prefix_len', 0)
rotary_emb
= Qwen3RotaryEmbedding(config)
shift_label
= dflash_config.get('shift_label', False)
target_layer_ids
nemo_automodel.components.speculative.dflash.draft_qwen3.Qwen3DFlashDraftModel._apply(
fn,
recurse = True
)

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.

nemo_automodel.components.speculative.dflash.draft_qwen3.Qwen3DFlashDraftModel.compute_logits(
hidden: torch.Tensor,
output_head: torch.nn.Module
) -> torch.Tensor

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:

hidden
torch.Tensor

Tensor of shape […, hidden]; draft hidden states with arbitrary leading dimensions.

output_head
nn.Module

The frozen target’s output projection.

Returns: torch.Tensor

Tensor of shape […, vocab].

nemo_automodel.components.speculative.dflash.draft_qwen3.Qwen3DFlashDraftModel.embed_noise_block(
target: torch.nn.Module,
block_ids: torch.LongTensor
) -> torch.Tensor

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:

target
nn.Module

The frozen verifier.

block_ids
torch.LongTensor

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

nemo_automodel.components.speculative.dflash.draft_qwen3.Qwen3DFlashDraftModel.forward(
position_ids: torch.LongTensor,
attention_mask: torch.Tensor | None = None,
noise_embedding: torch.Tensor | None = None,
target_hidden: torch.Tensor | None = None,
past_key_values: transformers.cache_utils.Cache | None = None,
use_cache: bool = False,
kwargs = {}
) -> torch.Tensor
nemo_automodel.components.speculative.dflash.draft_qwen3.Qwen3DFlashDraftModel.spec_generate(
target: torch.nn.Module,
input_ids: torch.LongTensor,
max_new_tokens: int,
stop_token_ids: list[int] | None,
temperature: float,
top_p: float = 1.0,
top_k: int = 0
) -> torch.LongTensor

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:

target
nn.Module

The frozen verifier.

input_ids
torch.LongTensor

Long tensor of shape [1, prompt].

max_new_tokens
int

Maximum number of tokens to generate.

stop_token_ids
list[int] | None

Token ids that end generation, or None.

temperature
float

Sampling temperature; 0 decodes greedily.

top_p
floatDefaults to 1.0

Nucleus mass to keep, in (0, 1].

top_k
intDefaults to 0

Candidates to keep, 0 for the whole vocabulary.

Returns: torch.LongTensor

Long tensor of shape [1, prompt + generated].

nemo_automodel.components.speculative.dflash.draft_qwen3.build_target_layer_ids(
num_target_layers: int,
num_draft_layers: int
) -> list[int]

Pick num_draft_layers target layers spread across the target’s depth.

nemo_automodel.components.speculative.dflash.core.compute_accept_len(
pred_ids_4d: torch.Tensor,
target_ids_4d: torch.Tensor,
valid_mask_4d: torch.Tensor
) -> torch.Tensor

Count consecutive accepted predictions for every draft block.

Parameters:

pred_ids_4d
torch.Tensor

Long tensor of shape [batch, blocks, depth].

target_ids_4d
torch.Tensor

Long tensor of shape [batch, blocks, depth].

valid_mask_4d
torch.Tensor

Bool tensor of shape [batch, blocks, depth].

Returns: torch.Tensor

Float tensor of shape [batch, blocks] containing the accepted draft

nemo_automodel.components.attention.dflash_mask.create_dflash_block_mask(
anchor_positions: torch.Tensor,
block_keep_mask: torch.Tensor,
ctx_len: int,
block_size: int,
device: torch.device,
use_compile: bool = True,
causal: bool = False,
ctx_doc_id: torch.Tensor | None = None,
anchor_doc_id: torch.Tensor | None = None,
sliding_window: int | None = None
) -> 'BlockMask'

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:

anchor_positions
torch.Tensor

[B, N] anchor positions (long).

block_keep_mask
torch.Tensor

[B, N] valid-anchor mask (bool).

ctx_len
int

context length.

block_size
int

block size.

device
torch.device

torch device.

use_compile
boolDefaults to True

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.

causal
boolDefaults to False

When True, make in-block attention causal (JetSpec); otherwise it is bidirectional (DFlash).

ctx_doc_id
torch.Tensor | NoneDefaults to None

[B, S] long per-context-token document id (sequence packing), or None to disable the per-document constraint.

anchor_doc_id
torch.Tensor | NoneDefaults to None

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

sliding_window
int | NoneDefaults to None

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.

nemo_automodel.components.attention.dflash_mask.create_dflash_sdpa_mask(
anchor_positions: torch.Tensor,
block_keep_mask: torch.Tensor,
ctx_len: int,
block_size: int,
device: torch.device,
dtype: torch.dtype,
causal: bool = False,
ctx_doc_id: torch.Tensor | None = None,
anchor_doc_id: torch.Tensor | None = None,
sliding_window: int | None = None
) -> torch.Tensor

Build a dense additive attention mask for the SDPA backend.

Parameters:

anchor_positions
torch.Tensor

[B, N] anchor positions per sample (long).

block_keep_mask
torch.Tensor

[B, N] per-sample valid-anchor mask (bool).

ctx_len
int

context length S.

block_size
int

block size.

device
torch.device

torch device.

dtype
torch.dtype

dtype for the additive mask (typically the model dtype).

causal
boolDefaults to False

When True, make in-block attention causal (JetSpec); otherwise it is bidirectional (DFlash).

ctx_doc_id
torch.Tensor | NoneDefaults to None

[B, S] long per-context-token document id (sequence packing), or None to disable the per-document constraint.

anchor_doc_id
torch.Tensor | NoneDefaults to None

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

sliding_window
int | NoneDefaults to None

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

nemo_automodel.components.speculative.dflash.draft_qwen3.extract_context_feature(
hidden_states: list[torch.Tensor],
layer_ids: list[int]
) -> torch.Tensor

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.

nemo_automodel.components.speculative.dflash.domino_core.get_lambda_base(
global_step: int,
total_steps: int,
lambda_start: float = 1.0,
decay_ratio: float = 0.5
) -> float

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

nemo_automodel.components.speculative.dflash.registry.resolve_dflash_draft_spec(
architectures: list[str]

Return the first registered DFlash draft spec matching any architecture in the list.

nemo_automodel.components.speculative.dflash.registry.DFLASH_DRAFT_REGISTRY: dict[str, DFlashDraftSpec] = {None: {arch: (DFlashDraftSpec(draft_cls=Qwen3DFlashDraftModel, draft2_cls=Qwen3...