mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
### What does this PR do? Type of change: new feature Adds **DFlash2** ([blog](https://inco.ai/blog/dflash2/)) as a draft variant of the existing DFlash mode, selected with `dflash_architecture_config.projector_type="dflash2"` alongside `domino`, `dspark` and `lilicorr`. DFlash2 keeps DFlash's one-pass parallel backbone and adds two components that recover the acceptance a purely parallel draft loses: - **Grouped dynamic depthwise convolution** around every attention and MLP sublayer, giving each block position a view of its predecessors *inside* the block. Taps do not cross the block boundary, so the draft stays one forward pass. - **Low-rank candidate selector** scoring transitions between adjacent positions' top-k candidates, so serving walks one coherent path instead of taking an independent argmax per position. Both start as exact no-ops — the convolution's `base_kernel` is an identity and `kernel_projection` is zeroed; the selector's `successor_codebook` is zeroed — so a freshly built DFlash2 draft *is* its DFlash backbone, and enabling the variant is an extension rather than a perturbation. This matches the reference implementation ([SpecForge#772](https://github.com/sgl-project/SpecForge/pull/772), merged) and the way `modeling_lilicorr` already installs this same convolution class. **This unblocks a recipe already shipped on `main`.** `modeling_lilicorr._install_sublayer_convs` imports `DFlashGroupedConv` from `modeling_dflash2`, so `modelopt_recipes/general/speculative_decoding/lilicorr_conv.yaml` raises at model build today and its CHANGELOG entry documents a feature that cannot run. Landing this makes it runnable. Module and parameter names match the SGLang/vLLM `DFlash2DraftModel` loaders. Verified against the released `z-lab/Qwen3.8-27B-DFlash2` checkpoint: **81 tensors, 21 name patterns, zero difference in either direction**. The serving side, [vllm-project/vllm#52816](https://github.com/vllm-project/vllm/pull/52816), has since merged (`b389ac294`) with no change to the checkpoint contract. ### Usage ```bash python examples/speculative_decoding/main.py \ --config modelopt_recipes/general/speculative_decoding/dflash2.yaml \ model.model_name_or_path=Qwen/Qwen3-8B \ data.data_path=<corpus>.jsonl \ training.output_dir=<out> ``` ```yaml # modelopt_recipes/general/speculative_decoding/dflash2.yaml dflash: dflash_selector_loss_alpha: 1.0 # weight of the candidate-selector CE term dflash_architecture_config: projector_type: dflash2 conv_kernel_size: 2 # taps; must not exceed the block size conv_group_size: 16 # must divide hidden_size selector_rank: 256 selector_top_k: 16 ``` ### Testing <img width="2000" height="1320" alt="image" src="https://github.com/user-attachments/assets/869004d1-92c6-41a4-a03d-fd824a06255c" /> **Unit** — 25 CPU tests in `tests/unit/torch/speculative/plugins/test_hf_dflash2.py`; the full `tests/unit/torch/speculative/` suite passes with no regressions. The ones worth keeping pin invariants that a decreasing loss does not catch: - the convolution is an exact identity on the **default** construction, and its taps stay inside the block while a position still sees its predecessors; - the block-offset contract shared by the training objective and `CandidateSelector.greedy_path` — a misaligned objective still converges; - the export fields the vLLM loader requires, including the top-level `block_size` that `DFlash2Exporter` derives the nested copy from; - which selector factors receive gradient on the first step. `successor_codebook` starts at zero, so `predecessor_codebook` and `hidden_projection` take one step to begin moving. That is a warm start, not a dead branch, and both sides are asserted. **End-to-end** — trained on Qwen3-8B against a plain DFlash control with every other argument identical (plot above). Monotonic convergence, no NaN/divergence, no DDP unused-parameter issues. Note the losses are **not comparable** across arms: DFlash2's includes the selector CE term. **Serving (vLLM)** — the exported drafter loads and drafts under the merged DFlash2 path (`RESOLVED draft architectures: ['DFlash2DraftModel']`). Two notes for anyone reproducing: vLLM sizes the convolution from `1 + num_speculative_tokens` at runtime rather than from the checkpoint, so a `block_size=16` drafter is only correct at `num_speculative_tokens=15`; and at that value the upstream path currently hits an illegal memory access in `_cache_draft_logits` ([vllm#55279](https://github.com/vllm-project/vllm/issues/55279)), independent of which checkpoint is used. ### Before your PR is "*Ready for review*" - Is this change backward compatible?: ✅ — additive. New `projector_type`, its own registry and exporter, one new config field; DFlash / Domino / DSpark / LiLiCorr numerics and `state_dict` contents are untouched. - If you copied code from any other sources or added a new PIP dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ — `modeling_dflash2.py` is adapted from [SpecForge#772](https://github.com/sgl-project/SpecForge/pull/772) and carries its MIT notice. No new dependencies. - Did you write any new necessary tests?: ✅ - Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?: ✅ — under `0.48.0`. - Did you get Claude approval on this PR?: ✅ — run on 2026-08-20; all review threads addressed and resolved. ### Additional Information Rebased onto current `main`. Two commits from the original branch were dropped because [#2342](https://github.com/NVIDIA/Model-Optimizer/pull/2342) landed them first, with authorship preserved: the no-op sublayer seam in `modeling_dflash.py`, and the `rope_theta`/`rope_parameters` fix — `main`'s version of the latter is stricter, so this PR no longer touches `hf_dflash.py` at all. <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Added DFlash2 speculative decoding with grouped dynamic convolutions and low-rank candidate selection. * Added configurable selector-loss weighting, including an option to disable it. * Added DFlash2 model conversion and export support. * Added checkpoints compatible with SGLang and vLLM DFlash2 serving. * **Documentation** * Added training recipes and a Qwen3-8B online DFlash2 training configuration. * **Tests** * Added coverage for conversion, training, metrics, gradients, and export compatibility. <!-- 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>
125 lines
5.6 KiB
YAML
125 lines
5.6 KiB
YAML
# DFlash2 streaming speculative decoding pipeline for Qwen3-8B.
|
|
#
|
|
# Same streaming transport as hf_streaming_dflash.yaml: a live `vllm serve` captures the
|
|
# target model's hidden states and moves them to the trainer over NIXL RDMA (no disk
|
|
# round-trip). DFlash2 trains the same block-diffusion backbone plus the grouped sublayer
|
|
# convolutions and the candidate selector. See common/eagle3/train_eagle_streaming.sh for
|
|
# the dispatch and rendezvous.
|
|
#
|
|
# Two nodes, all four GPUs of each: node 0 is one vLLM replica at TP=4, node 1 is a
|
|
# 4-rank DDP trainer. Scale down by lowering gpus_per_node and SERVE_TP together; scale up
|
|
# by raising nodes and SERVE_NODES together, as hf_streaming_dflash_multi_node.yaml does.
|
|
#
|
|
# 3-step pipeline:
|
|
# task_0: Build input conversations (jsonl)
|
|
# task_1: Streaming train — 1 serve node (TP=4) + 1 trainer node (4-rank DDP)
|
|
# task_2: vLLM smoke test with the exported drafter
|
|
#
|
|
# Site-specific settings left unset on purpose: an aarch64 cluster needs its own image
|
|
# (the default is x86), slurm_config.time where the queue's default is short, the NIXL and
|
|
# NCCL HCA pinning an InfiniBand fabric needs, and container-visible cache paths (srun
|
|
# runs with --no-container-mount-home, and TRITON_CACHE_DIR must be node-local).
|
|
#
|
|
# Usage:
|
|
# uv run launch.py --yaml examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml --yes
|
|
|
|
job_name: Qwen3-8B_DFlash2_streaming
|
|
pipeline:
|
|
allow_to_fail: false
|
|
skip: false
|
|
note:
|
|
|
|
global_vars:
|
|
hf_model: /hf-local/Qwen/Qwen3-8B
|
|
|
|
# Step 1: Build input conversations
|
|
task_0:
|
|
script: common/eagle3/make_dataset.sh
|
|
args:
|
|
- -f modules/Model-Optimizer/examples/dataset/example_data_config.yaml
|
|
- --full-conversations
|
|
slurm_config:
|
|
_factory_: "slurm_factory"
|
|
nodes: 1
|
|
ntasks_per_node: 1
|
|
gpus_per_node: 1
|
|
container: nvcr.io/nvidia/tensorrt-llm/release:1.3.0rc20
|
|
|
|
# Step 2: Streaming DFlash2 training — node 0 vllm serve (TP=4), node 1 trainer.
|
|
# DFlash2 extracts 5 target layers (build_target_layer_ids(36,5)=[1,9,17,25,33], the
|
|
# draft's fc input); vLLM's capture ids are those +1 -> [2,10,18,26,34]. No final layer
|
|
# (36): this recipe trains against the hard target, so the base hidden is never read.
|
|
# Dropping it must be paired with data.final_aux_is_base_hidden=true below. To train
|
|
# with distillation instead, set dflash_lk_loss_type=ce and add 36 back without it.
|
|
task_1:
|
|
script: common/eagle3/train_eagle_streaming.sh
|
|
args:
|
|
- --config modules/Model-Optimizer/modelopt_recipes/general/speculative_decoding/dflash2.yaml
|
|
- model.model_name_or_path=<<global_vars.hf_model>>
|
|
# Streaming never runs the base's transformer layers on the trainer; load only the
|
|
# embeddings, final norm and lm_head rather than all 36 layers on every rank.
|
|
- model.use_fake_base_for_offline=true
|
|
- data.mode=streaming
|
|
- data.data_path=/scratchspace/data/train.jsonl
|
|
# All five captured planes are aux features; without this the streaming dataset
|
|
# peels the last one off as the (unused) KD target and the draft's fc gets
|
|
# 4x4096 where it wants 5x4096.
|
|
- data.final_aux_is_base_hidden=true
|
|
- training.output_dir=/scratchspace/dflash2
|
|
- training.training_seq_len=4096
|
|
- training.disable_tqdm=true
|
|
# The serve prefills the corpus conversation rather than generating, so the
|
|
# assistant turn is present; Qwen3-8B's stock template just has no
|
|
# {% generation %} tags to locate it. Pass chat_template_train.jinja to mask to it.
|
|
- training.answer_only_loss=false
|
|
# 4 per device x 4 ranks = global batch 16 sequences = 65,536 tokens/step.
|
|
- training.per_device_train_batch_size=4
|
|
- training.num_train_epochs=1
|
|
- training.max_steps=5000
|
|
# dflash2.yaml sets report_to=tensorboard, which hard-fails if tensorboard
|
|
# isn't in the serve container; the streaming trainer doesn't need it.
|
|
- training.report_to=none
|
|
environment:
|
|
- HF_MODEL_CKPT: <<global_vars.hf_model>>
|
|
# No spaces: nemo_run emits unquoted `export FOO=value`, so spaces would split.
|
|
- EAGLE_CAPTURE_IDS: "[2,10,18,26,34]"
|
|
# Serve replica nodes (Slurm nodes 0..SERVE_NODES-1); the rest are trainers.
|
|
- SERVE_NODES: "1"
|
|
- SERVE_TP: "4"
|
|
# training_seq_len plus headroom for the single decode step.
|
|
- SERVE_MAX_MODEL_LEN: "4608"
|
|
- SERVE_MAX_NUM_SEQS: "32"
|
|
- SERVE_GPU_MEM_UTIL: "0.9"
|
|
# RDMA pool slot capacity in tokens. Must be >= training_seq_len or long prompts
|
|
# overflow the slot and the producer silently skips capture. 32 slots is 5.4 GiB of
|
|
# pinned host memory against a peak in-flight of 4 ranks x 4 workers.
|
|
- HS_MAX_TOKENS: "4608"
|
|
- HS_POOL_SLOTS: "32"
|
|
- STREAMING_NUM_WORKERS: "4"
|
|
# DFlash2 uses a custom modeling file; export must trust remote code.
|
|
- EXPORT_EXTRA_ARGS: "--trust_remote_code"
|
|
slurm_config:
|
|
_factory_: "slurm_factory"
|
|
nodes: 2
|
|
ntasks_per_node: 1
|
|
gpus_per_node: 4
|
|
container: vllm/vllm-openai:latest
|
|
|
|
# Step 3: vLLM smoke test (uses the exported checkpoint from training).
|
|
# The method stays "dflash": vLLM has no separate dflash2 method and selects the
|
|
# DFlash2 path from the checkpoint's architectures: ["DFlash2DraftModel"].
|
|
task_2:
|
|
script: common/specdec/vllm_smoke_test.sh
|
|
environment:
|
|
- HF_MODEL_CKPT: <<global_vars.hf_model>>
|
|
- DRAFT_MODEL: /scratchspace/export
|
|
- SPEC_METHOD: "dflash"
|
|
- NUM_SPEC_TOKENS: "7"
|
|
- MIN_ACCEPTANCE_LENGTH: "1.2"
|
|
slurm_config:
|
|
_factory_: "slurm_factory"
|
|
container: vllm/vllm-openai:nightly
|
|
nodes: 1
|
|
ntasks_per_node: 1
|
|
gpus_per_node: 1
|