nemo_automodel.components.speculative.dspark

View as Markdown

DSpark speculative-decoding draft model and training objective.

A semi-autoregressive parallel drafter: a parallel backbone produces every position of a block in one pass, a lightweight serial Markov head injects intra-block token dependency, and a confidence head predicts per-position acceptance probability for scheduled verification.

Submodules

Package Contents

Classes

NameDescription
DSparkForwardOutputOutputs for one DSpark training forward.
Qwen3DSparkModel-

Functions

API

class nemo_automodel.components.speculative.dspark.common.DSparkForwardOutput(
draft_logits: torch.Tensor,
target_ids: torch.Tensor,
eval_mask: torch.Tensor,
block_keep_mask: torch.Tensor,
confidence_pred: torch.Tensor | None = None,
aligned_target_logits: torch.Tensor | None = None
)
Dataclass

Outputs for one DSpark training forward.

Shape symbols: batch_size: number of samples in the batch seq_len: source sequence length num_anchors: sampled anchor blocks per sample block_size: number of draft positions per anchor vocab_size: vocabulary size

The sampler keeps anchors whose first draft target is enabled by loss_mask. Later slots are supervised only while they remain inside seq_len and form a contiguous enabled prefix. Dummy anchors can still appear when a sample has too few valid anchors; they are masked out by block_keep_mask and eval_mask.

aligned_target_logits
Tensor | None = None
block_keep_mask
Tensor
confidence_pred
Tensor | None = None
draft_logits
Tensor
eval_mask
Tensor
target_ids
Tensor
class nemo_automodel.components.speculative.dspark.draft_qwen3.Qwen3DSparkModel(
config
)

Bases: Qwen3PreTrainedModel

_no_split_modules
= ['Qwen3DSparkDecoderLayer']
block_size
= int(config.block_size)
embed_tokens
enable_confidence_head
= bool(config.enable_confidence_head)
fc
hidden_norm
layers
lm_head
markov_head
= build_markov_head(config)
mask_token_id
= config.mask_token_id
norm
num_anchors
= int(config.num_anchors)
rotary_emb
= Qwen3RotaryEmbedding(config)
target_layer_ids
= config.target_layer_ids
nemo_automodel.components.speculative.dspark.draft_qwen3.Qwen3DSparkModel._apply(
fn,
recurse: bool = True
)

Keep the RoPE inv_freq buffer in fp32 across dtype casts.

model.to(bfloat16) (the training build path) would otherwise round inv_freq to bf16 and dephase RoPE with absolute position, eroding draft acceptance (see pin_rope_inv_freq_fp32).

nemo_automodel.components.speculative.dspark.draft_qwen3.Qwen3DSparkModel._forward_backbone(
position_ids: torch.LongTensor,
attention_mask: torch.Tensor | None = None,
noise_embedding: torch.Tensor | None = None,
target_hidden_states: torch.Tensor | None = None,
past_key_values: transformers.cache_utils.Cache | None = None,
use_cache: bool = False,
kwargs = {}
) -> torch.Tensor
nemo_automodel.components.speculative.dspark.draft_qwen3.Qwen3DSparkModel.compute_logits(
hidden_states: torch.Tensor
) -> torch.Tensor
nemo_automodel.components.speculative.dspark.draft_qwen3.Qwen3DSparkModel.forward(
input_ids: torch.Tensor,
target_hidden_states: torch.Tensor,
loss_mask: torch.Tensor,
target_last_hidden_states: torch.Tensor | None = None,
position_ids: torch.Tensor | None = None,
seq_lens: torch.Tensor | None = None,
doc_remaining: torch.Tensor | None = None

Run one DSpark training forward.

Sequence packing (position_ids [B, S] per-document reset positions, seq_lens [B, max_docs], doc_remaining [B, S]) keeps every block inside its anchor’s document: the anchor’s first target must be in-document, the block’s context prefix and supervision are restricted to that document, and the draft’s RoPE uses the per-document positions.

nemo_automodel.components.speculative.dspark.draft_qwen3.Qwen3DSparkModel.initialize_embeddings_and_head(
embed_tokens: torch.nn.Module,
lm_head: torch.nn.Module,
freeze: bool = True
)
nemo_automodel.components.speculative.dspark.draft_qwen3.Qwen3DSparkModel.predict_confidence_step(
hidden_states: torch.Tensor,
prev_token_ids: torch.Tensor | None = None
) -> torch.Tensor | None
nemo_automodel.components.speculative.dspark.draft_qwen3.Qwen3DSparkModel.sample_draft_token_step(
base_logits: torch.Tensor,
prev_token_ids: torch.Tensor,
temperature: float = 0.0,
hidden_states: torch.Tensor | None = None
) -> tuple[torch.Tensor, torch.Tensor]
nemo_automodel.components.speculative.dspark.draft_qwen3.Qwen3DSparkModel.sample_draft_tokens(
base_logits: torch.Tensor,
first_prev_token_ids: torch.Tensor,
temperature: float = 0.0,
hidden_states: torch.Tensor | None = None
) -> tuple[torch.Tensor, torch.Tensor]
nemo_automodel.components.speculative.dspark.draft_qwen3.Qwen3DSparkModel.set_embedding_head_trainable(
trainable: bool
)
nemo_automodel.components.speculative.dspark.config.build_draft_config(
target_config,
model_args
)
nemo_automodel.components.speculative.dspark.loss.compute_dspark_loss(
loss_decay_gamma: float | None,
ce_loss_alpha: float,
l1_loss_alpha: float,
confidence_head_alpha: float,
return_terms: bool = False
)