Recipes & E2E ExamplesTrain a DSpark Drafter for Speculative Decoding

Train a DSpark Drafter for Speculative Decoding

View as Markdown

This guide shows you how to train a DSpark speculative-decoding drafter to accelerate LLM inference with NeMo AutoModel.

What is DSpark?

DSpark is a semi-autoregressive parallel drafter. A parallel backbone proposes every position of a block in a single forward pass, a lightweight serial Markov head injects intra-block token dependency (mitigating the acceptance decay of purely parallel drafters), and a confidence head predicts per-position acceptance probability for scheduled verification. The draft shares and freezes the target’s embed_tokens and lm_head, training only the backbone, the feature projection, the Markov head, and the confidence head.

It follows the same scaffolding as the EAGLE and DFlash recipes: online target hidden-state capture, gradient accumulation, and consolidated-safetensors checkpointing.

Objective

The draft is trained with a three-term, position-decay-weighted objective:

TermMeaning
L_ce (ce_loss_alpha, default 0.1)cross-entropy against the next target token
L_tv (l1_loss_alpha, default 0.9)total-variation distance to the target distribution (a direct acceptance proxy)
L_conf (confidence_head_alpha, default 1.0)BCE training the confidence head against measured acceptance

Positions are weighted by exp(-(k-1)/loss_decay_gamma).

Data

Use a chat dataset of OpenAI-format messages rows. As in the DSpark paper, use Open-PerfectBlend prompts with the responses regenerated by the target model (training is teacher-forced; regenerate before training to avoid a train/inference distribution mismatch). Point recipe_args.train_data_path at the regenerated JSONL or a Hugging Face dataset id.

Run It

Example configs live under examples/speculative/dspark/ (qwen3_0.6b_dspark.yaml, gemma4_12b_dspark.yaml, and deepseek_v41_flash_dspark.yaml). Multi-GPU defaults to FSDP2 (distributed.strategy: fsdp2); set it to ddp for simple data parallelism.

torchrun --standalone --nproc_per_node=2 \
-m nemo_automodel.recipes.llm.train_dspark \
-c examples/speculative/dspark/qwen3_0.6b_dspark.yaml

Per-step metrics (loss, ce_loss, l1_loss, confidence_loss, lr, mem) are reduced across data-parallel ranks and written to <output_dir>/dspark_train_metrics.jsonl.

Key Config Fields

FieldMeaning
target_model_name_or_pathfrozen target (e.g. Qwen/Qwen3-4B)
draft_num_hidden_layersdraft backbone depth (paper: 5)
block_sizetokens drafted per block (paper: 7)
num_anchorsblocks sampled per sequence per step
target_layer_idstarget feature layers fed to the draft (defaults to an even spread)
mask_token_idreserved token id filling non-anchor block positions
markov_rank / markov_head_typeserial head size and variant (vanilla / gated / rnn)
confidence_head_alpha / confidence_head_with_markovconfidence-head weight and conditioning
confidence_head_stop_gradientDeepSeek V4.1 only: train the confidence head on detached inputs so its loss cannot back-propagate into the draft backbone (default false)

DeepSeek V4.1 Flash

DeepSeek V4.1 uses the native DSpark architecture stored under the checkpoint’s mtp.* namespace rather than the configurable dense draft used for DeepSeek V4. The implementation reuses NeMo AutoModel’s V4.1 mHC, MLA, RoPE, Mixture-of-Experts (MoE) routing, normalization, and distributed-training components. The target remains frozen; only the three draft stages, target-feature projection, Markov head, and confidence head receive DSpark gradients.

NeMo AutoModel feeds the final RMSNorm output into the V4.1 confidence head to reduce its sensitivity to residual magnitude. The example keeps confidence_head_alpha: 1.0 and confidence_head_stop_gradient: false so the confidence loss still trains the draft backbone. This differs from the released confidence head, which reads the raw residual. Serving checkpoints trained with this recipe requires the same normalized confidence input; loading released confidence weights does not preserve their original predictions. Full-scale training stability with this normalization still requires validation.

The recipe starts those trainable DSpark weights from a fresh initialization and copies only the frozen embedding and LM head from the target. Use checkpoint.restore_from to resume a DSpark checkpoint produced by NeMo AutoModel. The V4.1 technical report does not name the dataset used for the dedicated DSpark stage, so the example data path is illustrative rather than a claim about the released model’s training data.

The following settings are fixed by the released checkpoint and the recipe rejects mismatches:

FieldDeepSeek V4.1 value
draft_num_hidden_layers3
block_size5
target_layer_ids[37, 38, 39]
mask_token_id128799
markov_rank256
draft attentionbidirectional five-position block plus a 128-token target window
draft MoE128 routed experts, top-3

Use the production-scale deepseek_v41_flash_dspark.yaml as the starting point. It uses FSDP2 for the trainable draft and the model’s distributed EP/FSDP path for the frozen target. Pipeline parallelism, context parallelism, sequence packing, image inputs, and the serving-time dynamic verification scheduler are not part of this training path.

Supported targets include Qwen3 (dense and MoE), Gemma4, DeepSeek V4 and V4.1, GLM-5.2, Kimi K3, and MiniMax M3.