> This page is for version Nightly (default).
> For other versions, use one of these documentation indexes:
> - Nightly (default): https://docs.nvidia.com/nemo/automodel/nightly/llms.txt
> - Latest: https://docs.nvidia.com/nemo/automodel/latest/llms.txt
> - 0.5.0 · 26.06: https://docs.nvidia.com/nemo/automodel/v0.5/llms.txt
> - 0.4.0 · 26.04: https://docs.nvidia.com/nemo/automodel/v0.4/llms.txt

> For clean Markdown of any page, append .md to the page URL.
> For a complete documentation index, see https://docs.nvidia.com/nemo/automodel/llms.txt.
> For AI client integration (Claude Code, Cursor, etc.), connect to the MCP server at https://docs.nvidia.com/nemo/automodel/_mcp/server.

# 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

| Name                                                                                 | Description                                                                          |
| ------------------------------------------------------------------------------------ | ------------------------------------------------------------------------------------ |
| [`TrainJetSpecRecipe`](#nemo_automodel-recipes-llm-train_jetspec-TrainJetSpecRecipe) | Recipe for JetSpec draft-model training: DFlash backbone + causal mask + forward-KL. |

### Functions

| Name                                                     | Description                          |
| -------------------------------------------------------- | ------------------------------------ |
| [`main`](#nemo_automodel-recipes-llm-train_jetspec-main) | Entrypoint for `TrainJetSpecRecipe`. |

### Data

[`logger`](#nemo_automodel-recipes-llm-train_jetspec-logger)

### API

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

**Bases:** [TrainDFlashRecipe](/nemo-automodel/nemo_automodel/recipes/llm/train_dflash#nemo_automodel-recipes-llm-train_dflash-TrainDFlashRecipe)

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

```python
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.

```python
nemo_automodel.recipes.llm.train_jetspec.TrainJetSpecRecipe._build_target_wrapper(
    target_layer_ids: list[int]
) -> nemo_automodel.components.speculative.dflash.target.HFDFlashTargetModel
```

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

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

```python
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.

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

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

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

Entrypoint for `TrainJetSpecRecipe`.

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