nemo_automodel.recipes.llm.train_jetspec
nemo_automodel.recipes.llm.train_jetspec
JetSpec draft-model training recipe (Qwen3-style targets).
JetSpec (arXiv:2606.18394) reuses the DFlash parallel draft backbone but trains it
as a causal parallel tree drafter: in-block attention is causal (so each branch
is conditioned on its own prefix) and the draft is distilled against the target’s
per-position soft distribution with a temperature-scaled forward-KL loss. See
nemo_automodel.components.speculative.dflash.jetspec_core.
This recipe reuses every piece of the DFlash recipe — online target hidden-state capture, anchor sampling, the block attention mask machinery, gradient accumulation, and checkpointing — and only (a) enables target-logit capture so the teacher distribution is available, and (b) swaps in the JetSpec trainer wrapper (causal mask + forward-KL).
IMPORTANT: as with DFlash, regenerate the training responses with the target model first — training teacher-forces ground-truth tokens while inference is autoregressive, and the distribution mismatch hurts acceptance length otherwise.
Module Contents
Classes
Functions
Data
API
Bases: TrainDFlashRecipe
Recipe for JetSpec draft-model training: DFlash backbone + causal mask + forward-KL.
Stamp causal=true: JetSpec drafts with causal in-block attention, and a serving
engine (vLLM reads dflash_config.causal) must match it at inference.
Capture the target’s full-vocab logits too — JetSpec distills against them.
Build the JetSpec trainer wrapper (causal parallel drafting + forward-KL).
Log the JetSpec acceptance-length proxy (tau) for the most recent step.
Forward through the JetSpec wrapper, passing the captured teacher logits.
Entrypoint for TrainJetSpecRecipe.