Files
h-guo18andClaude Opus 5 87f7d1432f fix(speculative): hold the DFlash draft's fp32 master weights in the optimizer (#2483)
### What does this PR do?

Type of change: Bug fix

**Follow-up to #2342**, which split this out on review (commit
`c67784d9`), and a rethink of how
the flag is implemented.

`dflash_fp32_master_weights` exists because the DFlash draft is cast to
the frozen bf16 target's
dtype, so AdamW allocates its moments in bf16 — and bf16 is too coarse
to hold them. At
`beta2=0.999` a single step changes `v` by at most **0.100%**, while the
smallest change bf16 can
represent near `v` is **0.164% mean / 0.388% max** (measured): every
decrease rounds away, `v` only
grows, and the effective step size decays on its own from step 1.

#2342 fixed that by **promoting the draft model to fp32**. Everything
else followed from giving the
model a dtype the rest of it does not have — a bf16 autocast at every
entry point, two transformers
loader hints so `from_pretrained(dtype="auto")` would not round the
draft away, a post-condition
check because those hints fail silently, and a doubled DDP gradient
all-reduce.

**This PR puts the fp32 in the optimizer instead**, where Megatron-LM,
DeepSpeed and apex put it.
`MasterWeightAdamW` holds an fp32 master copy of each non-fp32 parameter
plus fp32 moments in
`self.state[p]`, steps on the master, and copies back at the parameter's
dtype. The model is never
anything but the base dtype, so every one of those follow-on pieces is
deleted, gradients stay
bf16, and the exported drafter is unchanged. What the placement costs is
that wiring the optimizer
becomes the training loop's job:
`EagleTrainerWithAccLog.create_optimizer` builds it, and
`VerifyMasterWeightsCallback` raises at the end of step 1 if the moments
are not fp32.

**The default flips to `True`** — the flag now changes optimizer memory
and optimizer arithmetic
and nothing else. Flipping it on the model-promoted implementation turns
**25 of 259** unit tests
red; flipping it here is **259 passed**. Set it to `False` to reclaim
the memory, about 12 bytes per
draft parameter instead of 4.

<details>
<summary>Three drive-by fixes, independent of the above</summary>

- `_place_draft` is folded back into `modify()` — it fused the draft's
dtype, its device and an
  eager rotary buffer behind one meta guard.
- The module docstring's claim that `DFlashModule` has an `_apply`
meta-buffer fix is removed
  (`grep "def _apply"` matches nothing, and never did).
- #2342's field description no longer lists `evaluation` as a broken
path — `forward`
short-circuits to the base model when `not self.training`, so the draft
never runs there.

</details>

### Usage

No API change. `dflash_fp32_master_weights` now means the *optimizer*
holds fp32 master weights
rather than the draft model being fp32.

### Testing

**1 · The refactor is arithmetically a no-op.** Both implementations run
AdamW on an fp32 tensor,
so given the same starting values and the same gradients the
trajectories are identical — 1000
steps, `weight_decay=0.01`:

```
old fp32 parameter  vs  new fp32 master : bitwise equal = True  (max |diff| 0.0e+00)
exp_avg / exp_avg_sq                    : bitwise equal = True
optimizer state dtypes                  : ['torch.float32']
model parameter dtype                   : torch.bfloat16
```

Initial values have to be matched at bf16 first, or the bf16 arm's
one-time rounding of the draw
shows up as a 2e-4 "difference" that is not arithmetic. With that
controlled, the two
implementations differ only in their *inputs*: gradient precision (fp32
vs bf16 — torch 2.10
requires `grad.dtype == param.dtype`) and that one-time rounding.

**2 · End to end on GPU: the effect survives the refactor.** Qwen3-1.7B
base, real corpus, one GPU
per arm, three arms — pure bf16 (flag off), the #2342 implementation,
and this one — on two
algorithms trained independently, sharing seed, data order and
initialisation within an algorithm.

<img width="2925" height="960" alt="image"
src="https://github.com/user-attachments/assets/40d37059-ab8a-4924-b049-85f76b70b156"
/>


The two fp32 arms sit on top of each other for the whole run while bf16
stays above both, and the
old-vs-new gap is 10–23× smaller than the fp32-vs-bf16 effect it has to
be compared against.

**Acceptance length says the same thing, and settles what the loss could
not.** All six drafters at
the end of those curves were exported and served under vLLM against the
same base, and measured on
MT-Bench (80 prompts, 8 categories, greedy, one request at a time,
`num_speculative_tokens` =
trained `block_size` − 1, every knob but the drafter held fixed):

| | pure bf16 | fp32 in model (#2342) | fp32 in optimizer (this PR) |
new − old | fp32 − bf16 |
|---|---|---|---|---|---|
| `dflash` | 1.3068 | 1.3708 | **1.3666** | −0.0042 &nbsp;`t=−0.90` |
+0.0619 &nbsp;`t=+13.8` |
| `lilicorr` | 1.2536 | 1.2814 | **1.2882** | +0.0068 &nbsp;`t=+1.42` |
+0.0312 &nbsp;`t=+9.2` |

Paired by prompt, n=80. On both algorithms the new-vs-old 95% CI
straddles zero (`dflash`
[−0.0134, +0.0051], `lilicorr` [−0.0028, +0.0164]) while fp32-vs-bf16
does not come close to it,
and the sign of new-vs-old **flips between the two algorithms** — what a
rounding difference looks
like, not a bias. This is also the comparison the training loss could
not give: all three arms are
**exported and served in bf16**, so the old implementation's fp32 draft
weights are rounded at
export exactly as they would be for deployment, and the "its loss was
computed on a more precise
forward" caveat below does not apply. `lilicorr` needs

[vllm-project/vllm#57934](https://github.com/vllm-project/vllm/pull/57934),
applied as an overlay so
that both algorithms are measured on one engine build.

The right panel is the mechanism, and the one signal that depends on
neither the seed nor the
choice of loss statistic: Adam's updates to the draft's RMSNorm gains
are smaller than the bf16 ULP
at 1.0 (0.0078), so in the bf16 arm every one of them rounds away and
the gains never move — not
one of `dflash`'s 14 in 30000 steps, and two of `lilicorr`'s 20 by
3e-06. Both fp32 arms move
all of them, by the same amount.

Two results behind the figure rather than in it. **fp32-vs-bf16 grows
with the horizon** while
old-vs-new does not — on `dflash` −0.129 at 1500 steps → −0.262 at 15000
→ −0.341 at 30000, and on
`lilicorr` −0.191 → −0.220 → −0.285, against an old-vs-new difference
that stays near 0.02 at every
horizon and changes sign between them (−0.026 → +0.028 on `lilicorr`).
That is what a compounding
bias and a rounding difference respectively should look like, and it is
the reason the longer runs
were worth doing. And
**across seeds**, the paired old-vs-new difference at 1500 steps is
+0.0003 (n=6) on `dflash` and
+0.0643 (n=10) on `lilicorr`, both with a 95% CI straddling zero.

<details>
<summary>Limits of the above, stated rather than smoothed over</summary>

At 5 seeds the `lilicorr` paired difference read +0.1610 ± 0.0557
(t=+2.89, 4/5 seeds in the same
direction) — nominally significant, suggesting the new implementation
was genuinely worse there.
Four further `lilicorr` seeds were run against that pre-declared
question; two came back strongly
negative and the estimate settled at +0.0643 (95% CI [−0.086, +0.214]).
The earlier reading was
small-sample noise.

At 1500 steps on `lilicorr` that CI is *not* narrower than the
fp32-vs-bf16 effect it is being
compared against, so the 1500-step sweep alone cannot certify
equivalence there — `lilicorr` is
still at loss 9.3 and deep in its early transient, and it is the long
runs that resolve it. On
`dflash` the 1500-step CI (±0.031) is already 4× tighter than the effect
(−0.129).

One asymmetry the loss comparison cannot separate: the old
implementation held the draft weights in
fp32 *at forward time*, so its training loss was computed on a more
precise forward, while both
implementations export bf16. Any residual advantage it appears to have
is therefore an upper bound.

</details>

**3 · Unit tests.** `tests/unit/torch/speculative/` — **259 passed** on
**transformers 5.0.0** and
**5.3.0**, both ends of the supported `>=5.0,<5.13` (CPU, torch 2.10).
`TestDFlashFp32MasterWeights`
is rewritten for the new mechanism; the two that would have caught the
traps in this design are
`test_resume_does_not_round_the_master_back_down`
(`Optimizer.load_state_dict` casts float state to
its parameter's dtype, so a naive subclass rounds the master and both
moments to bf16 on *every*
resume, silently, with the loss still falling) and
`test_the_callback_refuses_a_loop_that_forgot_the_optimizer`. The rest
cover the draft's dtype with
the flag either way, that no forward path needs an autocast any more,
that plain AdamW really does
leave the moments in bf16, and that an fp32 model allocates no redundant
master. A sharded FSDP2
`DTensor` keeps an fp32 master and fp32 moments through a step.

### Before your PR is "*Ready for review*"

- Is this change backward compatible?: ✅ for artifacts, with one
intentional default change.
The draft's stored dtype goes back to matching the base, as it was
before #2342; existing
checkpoints load unchanged and the exported drafter is unaffected. The
flag now defaults to
**`True`** — the measurements above are the reason, and the cost is fp32
master + fp32 moments for
the draft only. A training loop that builds its own optimizer instead of
using the shipped
`create_optimizer` gets plain AdamW and none of this;
`VerifyMasterWeightsCallback` makes that
  fail loudly at step 1 rather than skip the feature quietly.
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ — no new
dependencies.
- Did you write any new necessary tests?: ✅
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅
- Did you get Claude approval on this PR?: ❌ — draft; will run `/claude
review` before marking ready.

### Additional Information

**On the 7–14% acceptance-length gain quoted in #2342:** that was
measured on the fp32-model
arithmetic and is not re-derived here. What is measured above is the
like-for-like comparison this
PR has to answer — same corpus, same horizon, same serving path, one
implementation swapped.

**History:** commits 1–3 restore the autocast design as it was split
out; commits 4–6 replace it.
Happy to squash before review.

<details>
<summary>Alternatives measured and rejected, so they do not get
re-proposed</summary>

- **Swapping `p.data` to the master and calling `super().step()`**
(reuses all of AdamW, ~20 lines
instead of ~50): bit-identical on ordinary parameters over 25 steps, but
silently wrong under
FSDP2 — assigning `.data` on a `DTensor` parameter updates the wrapper's
reported dtype while the
local shard keeps the model's, so `p.dtype` reads fp32, `p.data.dtype`
reads bf16, and
`zeros_like(p)` allocates the moments in bf16 anyway. CPU tests pass
either way.
- **Narrowing the autocast from `__call__` to `forward`** (while it
still existed): turns 10
Domino/DSpark tests red — the variants apply their heads in their own
`forward` overrides,
  outside `DFlashModule.forward`.
- **Building the rotary buffer on meta and letting the loader
materialise it**: makes RoPE
correctness depend on transformers selecting a branch by class-name
substring
(`"RotaryEmbedding" in module.__class__.__name__`), and the `if not
hasattr` guard is then
permanently satisfied, so a later `to_empty()` leaves garbage forever —
measured `4.56e-41`, i.e.
  cos=1 / sin=0, no positional encoding at all.
- **Building it eagerly in `DFlashModule.__init__`**: lands before the
dtype cast, so `Module.to`
rounds the RoPE frequencies to bf16 on the default path — measured
`0.8659643530845642` →
  `0.8671875`, loss `3.47230935097` → `3.47114777565`.

</details>

🤖 Generated with [Claude Code](https://claude.com/claude-code)




<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

- **New Features**
- DFlash now uses FP32 optimizer master weights and Adam moments by
default while keeping draft parameters in the base model’s dtype.
- Master-weight training preserves optimizer precision when restoring
checkpoints.
  - The feature can be disabled to reduce optimizer memory usage.
  - Draft models consistently follow the base model’s dtype and device.

- **Bug Fixes**
  - DFlash workflows now support operation without autocast.
- Added validation for compatible AdamW-family optimizers and
master-weight precision, including resumed training runs.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-09-23 15:39:33 +08:00
..
2026-06-20 02:14:31 +00:00
2026-04-08 02:54:18 -07:00
2025-07-31 22:48:09 +05:30

Speculative Decoding

Documentation

Speculative decoding accelerates auto-regressive generation in large language models (LLMs) by leveraging a lightweight draft model to predict the next γ tokens. The main LLM then verifies these candidate tokens in a single forward pass. If the draft model correctly predicts α tokens, the LLM can accept and generate α+1 tokens per verification step, significantly improving generation speed.

This folder contains an end-to-end runnable speculative decoding fine‑tuning pipeline in which Llama‑3.2‑1B (Hugging Face) is trained on the Daring-Anteater dataset.

This example focuses on training with Hugging Face. To train with Megatron‑LM, see the Megatron‑LM example.

Contents

Section Description Jump To
Pre-Requisites Required & optional dependencies [Link]
Simplified Workflow Train, evaluate, and export EAGLE model with one-line command [Link]
Online Training Train draft model alongside base model in GPU memory [Link]
Offline Training Train draft model using pre-computed hidden states [Link]
Streaming Training Train draft on hidden states streamed from a live vLLM serve (no disk dump) [Link]
After Training Evaluation, export and deployment [Link]
Advanced Usage Data synthesis, vocab compression, and configuration [Link]
Support Matrix Supported models for speculative decoding training [Link]
Speculation Module Checkpoints View pre-trained speculation modules ready to deploy! [Link]
Resources Extra links to relevant resources [Link]

Pre-Requisites

Docker

Please use the PyTorch docker image (e.g., nvcr.io/nvidia/pytorch:25.08-py3) or visit our installation docs for more information.

Also follow the installation steps below to upgrade to the latest version of Model Optimizer and install dataset and example-specific dependencies.

Local Installation

Install Modelopt with hf dependencies and other requirements for this example:

pip install -U nvidia-modelopt[hf]
pip install -r requirements.txt

Data Preparation

We support a range of input datasets. In this example, we will use the Daring-Anteater dataset.

python ../dataset/make_dataset.py -f ../dataset/example_data_config.yaml --full-conversations

See other-datasets section for other dataset options and instruction for user-provided data.

Omit --full-conversations if you plan to run synthetic data generation (see data-synthesis).

For large-scale training with NVIDIA's Nemotron datasets, use the dedicated scripts described in Nemotron Datasets.

Getting Started: Simplified Workflow

bash train_eagle3_and_export.sh --base_model meta-llama/Llama-3.2-1B-Instruct

This one-line command runs a minimal example workflow of training and exporting an EAGLE draft model in Modelopt. Specifically, it

  • Initializes the draft model with default settings
  • Fine-tunes the model on the dataset
  • Evaluates the acceptance rate on MT-Bench
  • Exports a checkpoint ready for deployment

Training Draft Model with Online Base Model

For small base models that fit in GPU memory, we can collocate them with draft models and train with the following command:

./launch_train.sh \
    --config ../../modelopt_recipes/general/speculative_decoding/eagle3.yaml \
    model.model_name_or_path=meta-llama/Llama-3.2-1B-Instruct \
    data.data_path=input_conversations/train.jsonl \
    training.output_dir=ckpts/llama-3.2-1b-online

All default training settings are in eagle3.yaml. You can adjust them by editing the YAML file or by specifying command-line overrides with OmegaConf dotlist arguments.

To enable context parallelism for long-context training, add training.cp_size=<N>. The saved modelopt checkpoint is similar in architecture to HF models. It can be further optimized through ModelOpt, e.g., PTQ and QAT.

Training Draft Model with Offline Base Model

For large models, you can export intermediate hidden states to disk and train only the draft model. This significantly reduces GPU memory requirements, but requires several to tens of terabytes of disk storage depending on dataset size.

Dumpping Hidden States to Disk

We support two backends for generating base model hidden states. For better effciency, it is recommended to use TRT-LLM:

python collect_hidden_states/compute_hidden_states_trtllm.py \
            --model $BASE_MODEL \
            --input-file input_conversations/train.jsonl \
            --output-dir $HIDDEN_STATES_DIR

NOTE: TRT-LLM installation needed for the above command.

Alternatively, you can generate the same hidden states with HF:

python collect_hidden_states/compute_hidden_states_hf.py \
            --model $BASE_MODEL \
            --input-file input_conversations/train.jsonl  \
            --output-dir $HIDDEN_STATES_DIR

NOTE: See run_hf_compute_hiddens_dp.sh and run_trtllm_compute_hiddens_dp.sh for a simple example using data parallelism (DP) to accelerate hidden state generation.

Train Draft Model with Dumped Hidden States

Once we finish dumping hidden states, launch offline training pointing to the hidden states directory:

./launch_train.sh \
    --config ../../modelopt_recipes/general/speculative_decoding/eagle3.yaml \
    model.model_name_or_path=meta-llama/Llama-3.2-1B-Instruct \
    data.offline_data_path=$HIDDEN_STATES_DIR \
    training.output_dir=ckpts/llama-3.2-1b-offline

Training Draft Model with Streaming Base Model

For large base models, you can stream hidden states from a live vllm serve instead of dumping them to disk: a co-located server produces the base-model hidden states on the fly and sends them to the trainer over NIXL RDMA, scaling to multiple nodes (dedicated serve replicas + DDP trainers). See the launcher examples, e.g. Kimi-K2.5 streaming EAGLE3 and streaming DFlash.

Model Validation

For online training checkpoints, we can run in-framework evaluation on MT-bench:

python scripts/ar_validate.py --model_path $ONLINE_CKPT

Note: In-framework evaluation is supported only for online training. For offline training checkpoints, please export the model and evaluate it using serving frameworks.

Export

python scripts/export_hf_checkpoint.py --model_path $OUTPUT_DIR --export_path $EXPORT_PATH

This exports the model from a ModelOpt checkpoint to a deployment-compatible format.

Deployment

The exported checkpoint can be deployed on TRT-LLM or SGLang.

TRT-LLM

To serve the checkpoint with TRT-LLM, run trtllm-serve with:

trtllm-serve <base_model_checkpoint> --host 0.0.0.0 --port 8000 --backend pytorch --max_batch_size 32 --max_num_tokens 8192 --max_seq_len 8192 --extra_llm_api_options extra-llm-api-config.yml

, with extra-llm-api-config.yml being

enable_attention_dp: false
disable_overlap_scheduler: true
enable_autotuner: false

cuda_graph_config:
    max_batch_size: 1

speculative_config:
    decoding_type: Eagle
    max_draft_len: 3
    speculative_model_dir: <draft_model_checkpoint>

kv_cache_config:
    enable_block_reuse: false

Please refer to TRT-LLM Doc: Speculative Decoding for detailed usage.

vLLM

Please refer to VLLM Doc: Speculative Decoding for detailed usage.

Optionally, you can convert the exported checkpoint to contain target model information, which is accepted by vLLM to simplify depployment:

python scripts/convert_to_vllm_ckpt.py --input <exported_ckpt> --verifier <target_model> --output <output_dir>

SGLang

Please refer to SGLang Doc: Speculative Decoding for detailed usage.

SpecDec Bench

One can also use examples/specdec_bench to validate the trained Eagle3 checkpoints in a variety of frameworks (vLLM, SGLang, TRT-LLM) on a set of datasets.

Deploying Quantized model

See more details on deployment of quantized model to TRTLLM here.

Advanced Usage

Other Datasets

In addition to the default dataset, we support adding several other commonly used datasets in ../dataset/make_dataset.py:

  • MTBench (for debugging)
  • ShareGPT
  • Magpie (Full 1M, and 500k and 300k filtered)
  • Nemotron Post-Training Dataset V2

To use your own datasets, please preprocess your data into a .jsonl file with each line in the format:

{
    "messages": [{"role": "user", "content": "..."}, {"role": "assistant", "content": "..."}]
}

Nemotron Datasets

For large-scale training we provide dedicated scripts for NVIDIA's Nemotron Post-Training dataset collections. Both scripts support two modes:

  • generate (default) — strips all assistant turns, producing a conversation skeleton (system + user turns only) for synthetic data generation. The downstream pipeline feeds these to the target model turn-by-turn, appending each generated response before sending the next user turn. Optional augmentation adds language-redirect and style-hint variants to diversify prompts.
  • train — keeps all turns in clean OpenAI message format (role + content) for direct SFT training. Prompt-only rows are dropped. Tool-call context (tool_calls, tool_call_id) is preserved for agentic datasets.

Nemotron Post-Training Dataset V2 (nvidia/Nemotron-Post-Training-Dataset-v2):

# Synthetic data generation (~3.3M rows):
python ../dataset/make_nemotron_ptv2_dataset.py --output-dir /tmp/ptv2_gen

# Direct SFT training mix (~1.9M rows):
python ../dataset/make_nemotron_ptv2_dataset.py --mode train --output-dir /tmp/ptv2_train

Covers: stem, chat, math, code + 5 multilingual splits (ja/de/it/es/fr, capped at 100K each).

Nemotron Post-Training V3 collection (16 datasets):

# Synthetic data generation (~3.4M rows):
python ../dataset/make_nemotron_ptv3_dataset.py --output-dir /tmp/ptv3_gen

# Direct SFT training mix (~3.9M rows, includes agentic/tool-use datasets):
python ../dataset/make_nemotron_ptv3_dataset.py --mode train --output-dir /tmp/ptv3_train

Covers: math, code, science, instruction-following, agentic/tool-use, safety, finance, and multilingual data. The dataset mix and per-split row caps are configurable via ../dataset/nemotron_ptv3_datasets.yaml.

Augmentation (generate mode only) is controlled by ../dataset/augmentations.yaml. By default it includes 12 language-redirect variants and several style/format hints. The /no_think system-prompt variant is disabled by default (enable it for models that support it, e.g. Qwen3):

# Custom augmentation config:
python ../dataset/make_nemotron_ptv2_dataset.py \
    --augmentations-config my_augs.yaml --output-dir /tmp/ptv2_gen

Data Synthesis

To achieve higher acceptance rates during speculative decoding, it is beneficial to use conversations generated by the base model as training data. This ensures that the draft model's output distribution closely aligns with that of the base model.

First, prepare input conversation skeletons using --mode generate (default) from the Nemotron scripts above, or with make_dataset.py (omitting --full-conversations). Then launch an inference server with the base model:

pip install vllm
vllm serve meta-llama/Llama-3.2-1B-Instruct --api-key token-abc123 --port 8000  --tensor-parallel-size 1

Note: Add --quantization=modelopt flag for quantized models.

Then, we generate conversations with the base model using the prepared prompts:

python scripts/server_generate.py --data_path input_conversations/train.jsonl --output_path synthetic/train.jsonl

To add a system prompt, use the --system_prompt <system_prompt_text> argument.

For large scale data generation, please see SLURM prepare data for SLURM support.

Configuring Draft Model

For EAGLE‑1 and EAGLE‑3 we provide a default model architecture config in ModelOpt. You can override default settings via eagle.eagle_architecture_config in the YAML. E.g. to use a 2-layer EAGLE head with 8192 intermediate size:

eagle:
  eagle_architecture_config:
    num_hidden_layers: 2
    intermediate_size: 8192

Draft Vocabulary Compression

We can optionally use smaller vocab size for the draft model for faster training and inference. E.g. Llama3.2-1B has a vocab size of 128256. In this example, we construct a draft vocab mapping of size 32k by finding the most commonly appeared vocabs in our training set:

python scripts/calibrate_draft_vocab.py --model meta-llama/Llama-3.2-1B-Instruct --data input_conversations/train.jsonl --draft_vocab_size 32000 --save_dir draft_vocab_cache

This will produce a d2t.pt file in save_dir, which is the mapping from draft token to target token. During inference, draft tokens can be mapped back to target tokens by target_token = draft_token + d2t[draft_token].

Then, set eagle_architecture_config.draft_vocab_size: 32000 and data.draft_vocab_cache: <path_to_d2t.pt> in your YAML. The draft model will use this provided vocab table during training and export.

Interact with modelopt.torch.speculative

main.py provides a complete example for converting a HF base model for speculative decoding and training it. The core steps are loading the base model, converting it with an eagle config dict, and training with HF Trainer:

import modelopt.torch.speculative as mtsp

# Convert base model in-place to an EAGLE speculative decoding model
eagle_cfg = {"eagle_decoder_type": "llama", ...}  # fields from EagleConfig
mtsp.convert(model, [("eagle", eagle_cfg)])

# Train with HF Trainer as usual
trainer = transformers.Trainer(model=model, ...)
trainer.train()
trainer.save_model("<output_dir>")

See main.py for the full example including tokenizer setup, dataset loading, and checkpoint handling.

Support Matrix

Model Medusa EAGLE1/2 EAGLE3
LLAMA 2 ✅ ✅ ✅
LLAMA 3, 3.1 ✅ ✅ ✅
Mistral ✅ ✅ ✅
Phi 3 ✅ ✅ ✅
QWen 1.5,2,2.5,3 ✅ ✅ ✅
Kimi-K2.5, K2.6 ✅

Speculation Module Checkpoints

Ready-to-deploy speculation module checkpoints [🤗 Hugging Face - NVIDIA Speculative Decoding Modules Collection] Deployable on TensorRT-LLM and SGLang!
More models coming soon!

Resources

DFlash (Block Diffusion for Speculative Decoding)

DFlash is a parallel speculative decoding method based on Block Diffusion. Unlike autoregressive draft models (EAGLE3), DFlash predicts an entire block of tokens in a single forward pass using masked parallel prediction with KV injection from the target model's hidden states.

Quick Start

For a complete end-to-end example (training + evaluation), see the launcher example:

uv run launch.py --yaml examples/Qwen/Qwen3-8B/hf_online_dflash.yaml --yes

Key Configuration (dflash.yaml)

Field Default Description
dflash.dflash_block_size 8 Block size for parallel prediction
dflash.dflash_num_anchors 512 Number of anchor positions per sample
dflash.dflash_loss_decay_factor 4.0 Exponential decay gamma (0 disables)
dflash.dflash_self_logit_distillation true Use logit distillation from target
dflash.dflash_architecture_config.num_hidden_layers 5 Draft decoder layers
dflash.dflash_architecture_config.mask_token_id auto Token ID for masked positions
training.answer_only_loss false Mask loss on non-assistant tokens

Qwen3 sliding window attention is automatically supported — draft layers inherit layer_types and sliding_window from the config, matching the target model's attention pattern.

Export

python scripts/export_hf_checkpoint.py \
    --model_path /path/to/training/output \
    --export_path /path/to/exported/model

Results

See doc/dflash.md for design details, benchmark results, and open items.