nemo_automodel.recipes.llm.train_jetspec

View as Markdown

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

NameDescription
TrainJetSpecRecipeRecipe for JetSpec draft-model training: DFlash backbone + causal mask + forward-KL.

Functions

NameDescription
mainEntrypoint for TrainJetSpecRecipe.

Data

logger

API

class nemo_automodel.recipes.llm.train_jetspec.TrainJetSpecRecipe()

Bases: TrainDFlashRecipe

Recipe for JetSpec draft-model training: DFlash backbone + causal mask + forward-KL.

nemo_automodel.recipes.llm.train_jetspec.TrainJetSpecRecipe._build_dflash_config(
recipe_cfg,
target_layer_ids: list[int]
) -> dict

Stamp causal=true: JetSpec drafts with causal in-block attention, and a serving engine (vLLM reads dflash_config.causal) must match it at inference.

nemo_automodel.recipes.llm.train_jetspec.TrainJetSpecRecipe._build_target_wrapper(
target_layer_ids: list[int]

Capture the target’s full-vocab logits too — JetSpec distills against them.

nemo_automodel.recipes.llm.train_jetspec.TrainJetSpecRecipe._build_trainer_module(
attention_backend: str,
recipe_cfg
)

Build the JetSpec trainer wrapper (causal parallel drafting + forward-KL).

nemo_automodel.recipes.llm.train_jetspec.TrainJetSpecRecipe._log_extra_train_metrics(
epoch_idx: int
) -> None

Log the JetSpec acceptance-length proxy (tau) for the most recent step.

nemo_automodel.recipes.llm.train_jetspec.TrainJetSpecRecipe._run_trainer_step(
target_batch
)

Forward through the JetSpec wrapper, passing the captured teacher logits.

nemo_automodel.recipes.llm.train_jetspec.main(
config_path: str | None = None
)

Entrypoint for TrainJetSpecRecipe.

nemo_automodel.recipes.llm.train_jetspec.logger = logging.getLogger(__name__)