From c2aaa44f6040658a21a2f5d2213c10ecc6542512 Mon Sep 17 00:00:00 2001 From: h-guo18 <67671475+h-guo18@users.noreply.github.com> Date: Tue, 29 Sep 2026 11:43:02 +0800 Subject: [PATCH] [Speculative Decoding] DFlash2 draft variant (grouped sublayer convolution + candidate selector) (#2216) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ### 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=.jsonl \ training.output_dir= ``` ```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 image **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. ## 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. --------- Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> Co-authored-by: Claude Opus 5 (1M context) --- .pre-commit-config.yaml | 2 + CHANGELOG.rst | 4 + .../torch/export/plugins/hf_spec_export.py | 43 ++ modelopt/torch/speculative/config.py | 46 ++ .../torch/speculative/dflash/conversion.py | 10 +- .../torch/speculative/plugins/__init__.py | 1 + .../torch/speculative/plugins/hf_dflash.py | 31 +- .../torch/speculative/plugins/hf_dflash2.py | 279 +++++++++ .../speculative/plugins/modeling_dflash2.py | 317 ++++++++++ .../speculative/plugins/modeling_fakebase.py | 11 +- .../speculative/plugins/modeling_lilicorr.py | 18 +- .../general/speculative_decoding/dflash2.yaml | 104 ++++ .../speculative_decoding/lilicorr_conv.yaml | 13 +- .../torch/export/test_hf_spec_rope_export.py | 69 +- .../speculative/plugins/test_fakebase.py | 81 ++- .../speculative/plugins/test_hf_dflash2.py | 588 ++++++++++++++++++ .../speculative/plugins/test_hf_lilicorr.py | 93 +++ .../Qwen/Qwen3-8B/hf_online_dflash2.yaml | 107 ++++ .../Qwen/Qwen3-8B/hf_streaming_dflash2.yaml | 124 ++++ 19 files changed, 1922 insertions(+), 19 deletions(-) create mode 100644 modelopt/torch/speculative/plugins/hf_dflash2.py create mode 100644 modelopt/torch/speculative/plugins/modeling_dflash2.py create mode 100644 modelopt_recipes/general/speculative_decoding/dflash2.yaml create mode 100644 tests/unit/torch/speculative/plugins/test_hf_dflash2.py create mode 100644 tools/launcher/examples/Qwen/Qwen3-8B/hf_online_dflash2.yaml create mode 100644 tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 7654acab5..0078076d0 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -117,6 +117,8 @@ repos: modelopt/torch/speculative/plugins/modeling_domino.py| modelopt/torch/speculative/plugins/hf_dflash.py| modelopt/torch/speculative/plugins/modeling_dflash.py| + modelopt/torch/speculative/plugins/hf_dflash2.py| + modelopt/torch/speculative/plugins/modeling_dflash2.py| modelopt/torch/speculative/plugins/hf_dspark.py| modelopt/torch/speculative/plugins/modeling_dspark.py| modelopt/torch/speculative/plugins/hf_lilicorr.py| diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 6c76f6f6e..041f1fe16 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -26,6 +26,10 @@ Changelog - Add an end-to-end BEVFormer ONNX PTQ example with temporal calibration data generation, INT8 and FP8 quantization, TensorRT engine building, and nuScenes accuracy evaluation. See `examples/onnx_ptq/bevformer/README.md `_ for details. - Add a reusable local-Hessian NVFP4 PTQ recipe and the quantization recipe used for ``nvidia/Qwen3.8-27B-NVFP4``. +*Speculative Decoding* + +- Add the DFlash2 draft variant, selected with ``dflash_architecture_config.projector_type="dflash2"``: DFlash's one-pass parallel backbone plus a grouped dynamic convolution around every attention/MLP sublayer (``conv_kernel_size`` / ``conv_group_size``) and a low-rank candidate selector (``selector_rank`` / ``selector_top_k``, weighted by ``dflash_selector_loss_alpha``). Exported checkpoints declare ``DFlash2DraftModel`` and load in the SGLang/vLLM DFlash2 serving path. + *Megatron Framework (M-LM / M-Bridge)* - Add an end-to-end W4A4 NVFP4 PTQ and QAD tutorial for Qwen3.6-35B-A3B also covering evaluation and vLLM throughput benchmarking. See `examples/megatron_bridge/tutorials/Qwen3.6-35B-A3B/README.md `_ for details. diff --git a/modelopt/torch/export/plugins/hf_spec_export.py b/modelopt/torch/export/plugins/hf_spec_export.py index 06be11b8a..38ee13c7a 100644 --- a/modelopt/torch/export/plugins/hf_spec_export.py +++ b/modelopt/torch/export/plugins/hf_spec_export.py @@ -605,3 +605,46 @@ class DSparkExporter(DFlashExporter): } ) return config + + +class DFlash2Exporter(DFlashExporter): + """Draft model exporter for DFlash2 (DFlash backbone + convolutions + selector). + + Same z-lab-compatible format as DFlash, plus the DFlash2 weights + (``layers.*.attention_conv.*`` / ``layers.*.mlp_conv.*`` / + ``candidate_selector.*``, already captured by the inherited ``dflash_module.`` + stripping) and the config fields the SGLang/vLLM ``DFlash2DraftModel`` loader + needs to rebuild them (``conv_kernel_size``, ``conv_group_size``, + ``selector_rank``, ``selector_top_k``). + + The architecture name is what selects the DFlash2 serving path: a checkpoint + declaring ``DFlashDraftModel`` loads as a plain DFlash draft and would silently + ignore the convolutions and the selector. + """ + + def _export_config(self): + """Extend the DFlash config with the DFlash2 architecture fields.""" + config = super()._export_config() + draft_config = self.model.dflash_config + + config["architectures"] = ["DFlash2DraftModel"] + # Present because HFDFlash2Model.modify validates them at convert time. + config["dflash_config"].update( + { + "projector_type": getattr(draft_config, "projector_type", "dflash2"), + "conv_kernel_size": draft_config.conv_kernel_size, + "conv_group_size": draft_config.conv_group_size, + "selector_rank": draft_config.selector_rank, + "selector_top_k": draft_config.selector_top_k, + # The published DFlash2 checkpoints carry block_size inside + # dflash_config; the DFlash loader reads it from the top level. + # Emit both so either contract resolves to the same value. + "block_size": config["block_size"], + } + ) + # vLLM reads is_causal from the TOP level, and the published DFlash2 checkpoints + # state it explicitly rather than leaving it to be inferred from layer_types. + # Mirror dflash_config.causal, which DFlashExporter writes unconditionally from + # dflash_draft_attention; no parent writes a top-level is_causal. + config["is_causal"] = config["dflash_config"]["causal"] + return config diff --git a/modelopt/torch/speculative/config.py b/modelopt/torch/speculative/config.py index cc90f5680..b8ecf0aa9 100644 --- a/modelopt/torch/speculative/config.py +++ b/modelopt/torch/speculative/config.py @@ -375,6 +375,52 @@ class DFlashConfig(ModeloptBaseConfig): ), ) + dflash_lk_loss_type: Literal["ce", "tv", "lambda"] = ModeloptField( + default="ce", + description=( + "DFlash2 only: which divergence the block objective minimizes against the " + "hard target. 'ce' is -log q(gold), today's behavior. 'tv' is 1 - q(gold), the " + "total variation to the one-hot target, which is also the per-position expected " + "acceptance loss. 'lambda' anneals between them: the CE share is " + "dflash_lk_ce_scale * exp(-dflash_lk_ce_decay * a), where a is the mean q(gold) " + "over supervised positions, so the objective moves from fitting the " + "distribution to maximizing acceptance as acceptance improves. " + "'lambda' and 'tv' require dflash_self_logit_distillation=false: both read " + "q(gold) from the per-position cross-entropy, which the KD path does not " + "produce. Ignored unless dflash_architecture_config.projector_type == 'dflash2'." + ), + ) + + dflash_lk_ce_scale: float = ModeloptField( + default=1.0, + ge=0.0, + description=( + "DFlash2 only: scale of the CE share in the dflash_lk_loss_type='lambda' " + "blend. 1.0 starts the run as pure CE. Ignored for other loss types." + ), + ) + + dflash_lk_ce_decay: float = ModeloptField( + default=1.0, + ge=0.0, + description=( + "DFlash2 only: how fast the CE share decays as acceptance rises in the " + "dflash_lk_loss_type='lambda' blend. 0 pins the blend at dflash_lk_ce_scale. " + "Ignored for other loss types." + ), + ) + + dflash_selector_loss_alpha: float = ModeloptField( + default=1.0, + ge=0.0, + description=( + "DFlash2 only: weight of the candidate-selector cross-entropy term, added to " + "the backbone loss. The selector re-ranks the backbone's top-k candidates per " + "block position; 0 trains the backbone and convolutions only. " + "Ignored unless dflash_architecture_config.projector_type == 'dflash2'." + ), + ) + @model_validator(mode="after") def _check_dpace_alpha(self) -> "DFlashConfig": # Validate at construction regardless of the active objective, so a bad alpha diff --git a/modelopt/torch/speculative/dflash/conversion.py b/modelopt/torch/speculative/dflash/conversion.py index 3a1645382..2f7152ab5 100644 --- a/modelopt/torch/speculative/dflash/conversion.py +++ b/modelopt/torch/speculative/dflash/conversion.py @@ -35,6 +35,12 @@ DominoDMRegistry = _DMRegistryCls(prefix="Domino") # ``dflash_architecture_config.projector_type == "dspark"`` and kept in its own # registry so its wrapper (HFDSparkModel) does not overwrite HFDFlashModel. DSparkDMRegistry = _DMRegistryCls(prefix="DSpark") +# DFlash2 also reuses the dflash mode/config/recipe, converting the base model to a +# DFlash backbone whose sublayers are wrapped in grouped dynamic convolutions, plus a +# low-rank candidate selector. Selected via +# ``dflash_architecture_config.projector_type == "dflash2"`` and kept in its own +# registry so its wrapper (HFDFlash2Model) does not overwrite HFDFlashModel. +DFlash2DMRegistry = _DMRegistryCls(prefix="DFlash2") # LiLiCorr also reuses the dflash mode/config/recipe, converting the base model to a # DFlash backbone augmented with a reranker over the candidate lattice the backbone # already produces. Selected via @@ -59,6 +65,8 @@ def convert_to_dflash_model(model: nn.Module, config: DFlashConfig) -> ConvertRe registry = DominoDMRegistry elif projector_type == "dspark": registry = DSparkDMRegistry + elif projector_type == "dflash2": + registry = DFlash2DMRegistry elif projector_type == "lilicorr": registry = LiLiCorrDMRegistry elif projector_type in (None, "dflash"): @@ -66,7 +74,7 @@ def convert_to_dflash_model(model: nn.Module, config: DFlashConfig) -> ConvertRe else: raise ValueError( f"Unsupported dflash_architecture_config.projector_type: {projector_type!r}. " - "Expected 'dflash' (default), 'domino', 'dspark' or 'lilicorr'." + "Expected 'dflash' (default), 'domino', 'dspark', 'dflash2' or 'lilicorr'." ) original_cls = type(model) diff --git a/modelopt/torch/speculative/plugins/__init__.py b/modelopt/torch/speculative/plugins/__init__.py index 36cde677d..bbdcac755 100644 --- a/modelopt/torch/speculative/plugins/__init__.py +++ b/modelopt/torch/speculative/plugins/__init__.py @@ -31,6 +31,7 @@ with import_plugin("megatron_medusa"): with import_plugin("transformers"): from .hf_dflash import * + from .hf_dflash2 import * from .hf_domino import * from .hf_dspark import * from .hf_eagle import * diff --git a/modelopt/torch/speculative/plugins/hf_dflash.py b/modelopt/torch/speculative/plugins/hf_dflash.py index da9180223..084040a5e 100644 --- a/modelopt/torch/speculative/plugins/hf_dflash.py +++ b/modelopt/torch/speculative/plugins/hf_dflash.py @@ -73,7 +73,7 @@ Draft model components: import logging from pathlib import Path -from typing import Any +from typing import Any, NamedTuple import torch import torch.nn.functional as F @@ -200,6 +200,21 @@ def _dpace_position_weights( return weights.to(dtype=confidences.dtype) +class _DFlashLossTerms(NamedTuple): + """The unreduced pieces behind the block loss, for variants that re-derive it. + + ``ce_per_token`` is ``None`` on the KD path, which never forms a per-position + cross-entropy. ``weights`` carries the position weighting (decay or D-PACE) and + normalizes by ``weight_sum``; ``supervised_mask`` is the unweighted mask the + reported accuracy uses. + """ + + ce_per_token: torch.Tensor | None + weights: torch.Tensor + weight_sum: torch.Tensor + supervised_mask: torch.Tensor + + @DFlashDMRegistry.register({PreTrainedModel: "hf.PreTrainedModel"}) class HFDFlashModel(DFlashModel): """DFlash Model for HuggingFace transformers.""" @@ -823,6 +838,7 @@ class HFDFlashModel(DFlashModel): base_logits=None, draft_hidden=None, base_outputs=None, + return_terms=False, ): """Compute weighted cross-entropy (or KD) loss and accuracy. @@ -835,6 +851,11 @@ class HFDFlashModel(DFlashModel): base_logits: Base model logits for KD loss [B, seq_len, vocab], or None for CE. draft_hidden: Draft hidden states [B, N*block_size, H] behind ``logits``. Unused here; passed for variants whose head consumes them. + return_terms: Also return the unreduced pieces behind the loss, so a variant + can recompose the block objective from a different divergence without + rebuilding the target alignment and position weighting. + TODO: promote this into a shared divergence seam when the DFlash-family + loss code is refactored; DFlash2 is the only consumer today. Returns: (loss, accuracy) tuple. @@ -929,6 +950,14 @@ class HFDFlashModel(DFlashModel): loss = flat_logits.sum() * 0.0 accuracy = 0.0 + if return_terms: + terms = _DFlashLossTerms( + ce_per_token=loss_per_token, + weights=flat_weights, + weight_sum=valid_count, + supervised_mask=binary_eval_mask, + ) + return loss, accuracy, terms return loss, accuracy def forward( diff --git a/modelopt/torch/speculative/plugins/hf_dflash2.py b/modelopt/torch/speculative/plugins/hf_dflash2.py new file mode 100644 index 000000000..625bc6226 --- /dev/null +++ b/modelopt/torch/speculative/plugins/hf_dflash2.py @@ -0,0 +1,279 @@ +# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""HF DFlash2 model wrapper — DFlash training plus the candidate-selector objective. + +DFlash2 differs from DFlash only in the draft module (grouped dynamic convolutions +around every sublayer, plus a candidate selector) and in one extra loss term, so +this wrapper reuses ``HFDFlashModel``'s forward wholesale and overrides just +:meth:`_compute_loss`. + +The convolutions need no supervision of their own: they sit inside the backbone +and are trained by the backbone loss. The selector does, because at serving time +it — not an independent argmax — picks the drafted token at each block position. + +Selector supervision (following the SGLang/SpecForge reference): + +- Take the backbone's top-k candidates per block position. +- Score each candidate against its *teacher-forced* predecessor token, so the + positions train in parallel exactly as the backbone does. +- When the gold token is missing from the top-k, substitute it into the last + candidate slot. Without this the selector sees no positive class on the hard + positions and never learns those edges. +""" + +import torch +import torch.nn.functional as F +from transformers import PreTrainedModel + +from ..dflash.conversion import DFlash2DMRegistry +from .hf_dflash import HFDFlashModel +from .modeling_dflash2 import DFlash2Module + +__all__ = ["HFDFlash2Model"] + + +@DFlash2DMRegistry.register({PreTrainedModel: "hf.PreTrainedModel"}) +class HFDFlash2Model(HFDFlashModel): + """DFlash model with DFlash2's sublayer convolutions and candidate selector. + + Registered in ``DFlash2DMRegistry`` so that ``convert_to_dflash_model`` routes + to it when ``dflash_architecture_config.projector_type == "dflash2"``. + """ + + def _build_draft_module(self, dflash_config): + """Build the DFlash2 draft module (DFlash backbone + convolutions + selector).""" + return DFlash2Module(dflash_config) + + def modify(self, config): + """Initialize the DFlash2 draft module and read the selector loss weight.""" + arch_config = config.dflash_architecture_config + missing = [ + name + for name in ("conv_kernel_size", "conv_group_size", "selector_rank", "selector_top_k") + if arch_config.get(name) is None + ] + if missing: + raise ValueError( + f"DFlash2 (projector_type='dflash2') requires {missing} in " + "dflash_architecture_config (convolution taps/group size and the " + "candidate selector's rank/top-k)." + ) + super().modify(config) + self.dflash_selector_loss_alpha = getattr(config, "dflash_selector_loss_alpha", 1.0) + self.dflash_lk_loss_type = getattr(config, "dflash_lk_loss_type", "ce") + self.dflash_lk_ce_scale = getattr(config, "dflash_lk_ce_scale", 1.0) + self.dflash_lk_ce_decay = getattr(config, "dflash_lk_ce_decay", 1.0) + if self.dflash_lk_loss_type != "ce" and self.dflash_self_logit_distillation: + raise ValueError( + f"dflash_lk_loss_type={self.dflash_lk_loss_type!r} needs the draft's " + "probability of the gold token, which the KD path never forms -- it " + "optimizes a soft target instead. Set dflash_self_logit_distillation=false, " + "or dflash_lk_loss_type='ce' to keep distillation." + ) + self._selector_metrics = None + + def forward(self, *args, **kwargs): + """Run the DFlash training forward and attach the candidate-selector metrics. + + The variant is the objective, so the pipeline above it is inherited verbatim; + this only carries out what ``_compute_loss`` produced. Without it + ``selector_coverage`` -- the only signal distinguishing a selector that is + choosing from one being handed the gold token -- never reaches the logs. + """ + self._selector_metrics = None + outputs = super().forward(*args, **kwargs) + if self._selector_metrics is not None: + outputs["selector_metrics"] = self._selector_metrics + self._selector_metrics = None + return outputs + + def get_exporter(self): + """Get the exporter for the DFlash2 draft model.""" + from modelopt.torch.export.plugins.hf_spec_export import DFlash2Exporter + + return DFlash2Exporter(self) + + def _selector_loss(self, logits, target_ids, hidden, predecessor_ids, weight_mask): + """Cross-entropy over the selector's candidate set, and its top-1 accuracy. + + Args: + logits: Backbone logits per block position ``[B, N, block_size, V]``. + target_ids: Gold token ids ``[B, N, block_size]``. + hidden: Backbone hidden states ``[B, N, block_size, H]``. + predecessor_ids: Teacher-forced predecessor ids ``[B, N, block_size]``. + weight_mask: Per-position loss weights ``[B, N, block_size]``. + + Returns: + ``(loss, accuracy, coverage)`` — coverage is the fraction of supervised + positions whose gold token was already in the backbone's top-k, i.e. how + often the selector is choosing rather than being handed the answer. + """ + selector = self.dflash_module.candidate_selector + top_k = selector.top_k + + unary_logits, candidate_ids = logits.topk(top_k, dim=-1) + + # Train on the strict top-k, the candidate set serving actually builds. A gold + # token the backbone did not propose is a backbone recall failure, not a + # selector classification example: substituting it in would teach the selector + # to override the unary ranking on a set it will never be shown. Those + # positions carry no selector gradient and leave the denominator instead. + gold_matches = candidate_ids == target_ids.unsqueeze(-1) + gold_in_topk = gold_matches.any(dim=-1) + gold_slot = gold_matches.long().argmax(dim=-1) + + selector_logits = selector.score_candidates( + candidate_ids, unary_logits, hidden, predecessor_ids + ) + + covered = weight_mask * gold_in_topk.to(weight_mask.dtype) + flat_weights = covered.reshape(-1) + denominator = flat_weights.sum() + 1e-6 + per_token = F.cross_entropy( + selector_logits.float().reshape(-1, top_k), + gold_slot.reshape(-1), + reduction="none", + ) + loss = (per_token * flat_weights).sum() / denominator + + with torch.no_grad(): + chosen = selector_logits.argmax(dim=-1).reshape(-1) + accuracy = ( + (chosen == gold_slot.reshape(-1)).float() * flat_weights + ).sum() / denominator + # Coverage keeps the full supervised mask as its denominator: it measures + # how often the selector was given a solvable problem at all. + # clamp, not an epsilon: a fully covered batch must read exactly 1.0. + supervised = weight_mask.reshape(-1).sum() + coverage = flat_weights.sum() / supervised.clamp(min=1.0) + # Detached tensors, not Python scalars: .item() would force a CPU-GPU sync on + # every training step. The trainer converts them at the logging boundary. + return loss, accuracy.detach(), coverage.detach() + + def _lk_loss(self, terms): + """Re-weight the block objective between cross-entropy and acceptance. + + Both terms are read off the same per-position cross-entropy the backbone loss + already produced, so the target alignment and position weighting are shared:: + + q = exp(-ce) draft probability of the gold token + L_ce = _w today's objective + L_tv = <1 - q>_w total variation to the one-hot target, + i.e. the per-position acceptance loss + a = _mask mean acceptance over supervised positions + L = s*exp(-d*a) * L_ce + (1 - s*exp(-d*a)) * L_tv + + ``<.>_w`` averages under the position weighting, ``<.>_mask`` under the + unweighted supervised mask, matching how the reported accuracy is normalized. + The blend weight is detached, so it reshapes the objective without adding a + gradient path of its own. + """ + ce, weights, weight_sum = terms.ce_per_token, terms.weights, terms.weight_sum + assert ce is not None, ( + "the KD path produced no per-position cross-entropy, so q(gold) is " + "unavailable and the blend would silently fall back to it; modify() is " + "supposed to have rejected this combination at convert time" + ) + gold_probability = torch.exp(-ce) + tv_loss = ((1.0 - gold_probability) * weights).sum() / weight_sum + if self.dflash_lk_loss_type == "tv": + return tv_loss + + ce_loss = (ce * weights).sum() / weight_sum + mask = terms.supervised_mask + acceptance = (gold_probability.detach() * mask).sum() / (mask.sum() + 1e-6) + ce_share = self.dflash_lk_ce_scale * torch.exp(-self.dflash_lk_ce_decay * acceptance) + return ce_share * ce_loss + (1.0 - ce_share) * tv_loss + + def _compute_loss( + self, + logits, + input_ids, + anchor_positions, + block_keep_mask, + loss_mask, + base_logits=None, + draft_hidden=None, + base_outputs=None, + ): + """Backbone DFlash loss plus the candidate-selector cross-entropy. + + Reuses ``HFDFlashModel._compute_loss`` for the backbone term, then rebuilds + the same target/weight alignment for the selector term. Reported accuracy + stays the backbone's top-1, so DFlash and DFlash2 runs remain comparable; + the selector's own accuracy is logged separately. + """ + loss, accuracy, terms = super()._compute_loss( + logits, + input_ids, + anchor_positions, + block_keep_mask, + loss_mask, + base_logits, + draft_hidden=draft_hidden, + base_outputs=base_outputs, + return_terms=True, + ) + if self.dflash_lk_loss_type != "ce": + loss = self._lk_loss(terms) + if self.dflash_selector_loss_alpha <= 0 or draft_hidden is None: + return loss, accuracy + + bsz, seq_len = input_ids.shape + block_size = self.dflash_block_size + n_blocks = anchor_positions.shape[1] + device = input_ids.device + + offsets = torch.arange(block_size, device=device).view(1, 1, -1) + label_indices = anchor_positions.unsqueeze(-1) + offsets + valid_label = label_indices < seq_len + safe_label_indices = label_indices.clamp(max=seq_len - 1) + expanded_ids = input_ids.unsqueeze(1).expand(-1, n_blocks, -1) + target_ids = torch.gather(expanded_ids, 2, safe_label_indices) + + # Same supervision mask as the backbone loss: valid block, in bounds, not the + # anchor slot, and inside the answer span. Position weighting (decay/D-PACE) is + # deliberately not applied — it shapes *where* the backbone spends capacity, + # while the selector should learn every position's transition equally. + weight_mask = block_keep_mask.unsqueeze(-1).expand(-1, -1, block_size).float() + weight_mask = weight_mask * valid_label.float() + weight_mask = weight_mask * (offsets > 0).float() + weight_mask = weight_mask * torch.gather( + loss_mask.unsqueeze(1).expand(-1, n_blocks, -1), 2, safe_label_indices + ) + + # Teacher-forced predecessor of block position k is the real token at anchor+k-1. + # Only offsets 1..block_size-1 are supervised (weight_mask zeroes slot 0), so the + # first supervised position, k=1, has the anchor's own token as its predecessor. + # Slot 0's entry here resolves to anchor-1 and never contributes. + # + # This is the train/serve contract: CandidateSelector.greedy_path seeds its walk + # from the anchor token, so ITS position 0 corresponds to block offset 1. A caller + # that hands greedy_path the full 0..block_size-1 candidate set is off by one. + predecessor_ids = torch.gather(expanded_ids, 2, (safe_label_indices - 1).clamp(min=0)) + + selector_loss, selector_accuracy, selector_coverage = self._selector_loss( + logits.reshape(bsz, n_blocks, block_size, -1), + target_ids, + draft_hidden.reshape(bsz, n_blocks, block_size, -1), + predecessor_ids, + weight_mask, + ) + self._selector_metrics = { + "selector_accuracy": selector_accuracy, + "selector_coverage": selector_coverage, + } + return loss + self.dflash_selector_loss_alpha * selector_loss, accuracy diff --git a/modelopt/torch/speculative/plugins/modeling_dflash2.py b/modelopt/torch/speculative/plugins/modeling_dflash2.py new file mode 100644 index 000000000..e54a9b0f9 --- /dev/null +++ b/modelopt/torch/speculative/plugins/modeling_dflash2.py @@ -0,0 +1,317 @@ +# Adapted from https://github.com/sgl-project/SpecForge/pull/772 +# Copyright (c) 2025 sgl-project +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. + +# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 AND MIT +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""DFlash2 draft module — DFlash backbone plus local convolution and candidate selection. + +DFlash2 (Inco AI / Z Lab, https://inco.ai/blog/dflash2/) keeps DFlash's one-pass +parallel backbone and adds two small components that address the two ways a +purely parallel draft loses acceptance: + +- :class:`DFlashGroupedConv` — a grouped *dynamic* depthwise convolution wrapped + around every attention and MLP sublayer. Each block position mixes in its + predecessors inside the block, which injects the intra-block sequential + dependency the parallel backbone lacks (mitigating suffix acceptance decay) + without a second backbone pass. Taps do not cross the block boundary. + +- :class:`CandidateSelector` — a low-rank transition scorer. Instead of an + independent argmax per block position, the drafter keeps the target head's + top-k candidates per position and scores adjacent transitions, so serving can + walk one coherent path through the block. + +Where Domino uses a GRU and DSpark a Markov transition bias, DFlash2 spends its +extra capacity on these two pieces: both are cheap (a few percent of draft +parameters, ~1% of serving step latency in the reference measurements). + +This module owns the parameters only; the training wrapper (``HFDFlash2Model`` +in ``hf_dflash2.py``) orchestrates the forward and the selector loss. Module and +parameter names (``attention_conv`` / ``mlp_conv`` / ``base_kernel`` / +``kernel_projection`` / ``candidate_selector`` / ``predecessor_codebook`` / +``successor_codebook`` / ``hidden_projection``) match the SGLang and vLLM +``DFlash2DraftModel`` loaders so an exported checkpoint is served directly. +""" + +import torch +import torch.nn.functional as F +from torch import nn + +from .modeling_dflash import DFlashModule + +__all__ = ["CandidateSelector", "DFlash2Module", "DFlashGroupedConv"] + + +class DFlashGroupedConv(nn.Module): + """Grouped dynamic depthwise convolution over positions within a proposal block. + + Wraps one sublayer: :meth:`prepare` convolves the sublayer input and emits the + dynamic kernel for the output side, :meth:`finish` convolves the sublayer + output. One projection of the sublayer input produces both sides' kernel + deltas. + + Both halves start as a no-op: ``base_kernel`` is an identity (tap 0 weight 1, + later taps 0) and ``kernel_projection`` is zero-initialized, so ``delta`` is + zero and a freshly built DFlash2 draft computes exactly what its DFlash + backbone would. That makes the convolution a stable extension rather than a + perturbation. Matches the reference implementation (SpecForge #772) and the + way ``modeling_lilicorr`` installs this same class. + """ + + def __init__(self, hidden_size: int, block_size: int, taps: int, group_size: int): + """Build the identity-initialized base kernel and the dynamic-kernel projection.""" + super().__init__() + if taps < 1: + raise ValueError(f"DFlash2 conv_kernel_size must be >= 1, got {taps}.") + if taps > block_size: + raise ValueError( + f"DFlash2 conv_kernel_size ({taps}) must not exceed " + f"dflash_block_size ({block_size})." + ) + if group_size < 1 or hidden_size % group_size: + raise ValueError( + f"DFlash2 conv_group_size ({group_size}) must be >= 1 and divide " + f"hidden_size ({hidden_size})." + ) + + self.block_size = int(block_size) + self.taps = int(taps) + self.group_size = int(group_size) + self.num_groups = int(hidden_size) // self.group_size + + # [input/output side, tap, channel]; identity at tap 0. Layout matches the + # SGLang/vLLM DFlash2 weight loader. + base_kernel = torch.zeros(2, self.taps, int(hidden_size)) + base_kernel[:, 0] = 1.0 + self.base_kernel = nn.Parameter(base_kernel) + self.kernel_projection = nn.Linear( + int(hidden_size), 2 * self.taps * self.num_groups, bias=False + ) + # Zero here, not only in DFlash2Module._init_head_weights, so the wrapper is an + # exact identity however it is built -- modeling_lilicorr constructs it directly. + nn.init.zeros_(self.kernel_projection.weight) + + def _convolve(self, hidden_states, delta, side: int): + """Apply the depthwise convolution for one side, with taps clipped at block starts.""" + bsz, seq_len, hidden_size = hidden_states.shape + if seq_len % self.block_size: + raise ValueError( + f"DFlash2 convolution needs a sequence length divisible by " + f"block_size ({self.block_size}), got {seq_len}." + ) + + n_blocks = seq_len // self.block_size + blocks = hidden_states.reshape( + bsz, n_blocks, self.block_size, self.num_groups, self.group_size + ) + dynamic = delta.reshape(bsz, n_blocks, self.block_size, self.taps, self.num_groups) + base = self.base_kernel[side].reshape(self.taps, self.num_groups, self.group_size) + + # (base + delta) * x is expanded as base * x + delta * x rather than formed as a + # dense coefficient tensor. The summed form would be [.., taps, groups, group_size] + # -- taps * hidden floats per position -- and, being a multiplicand, autograd would + # hold it until backward; the delta it is built from is group_size times smaller. + output = base[0] * blocks + dynamic[:, :, :, 0].unsqueeze(-1) * blocks + for tap in range(1, self.taps): + # Shift within the block only: position k reads k-tap, and the first + # `tap` positions of each block read zeros rather than the previous block. + shifted = F.pad(blocks[:, :, : self.block_size - tap], (0, 0, 0, 0, tap, 0)) + output = output + base[tap] * shifted + dynamic[:, :, :, tap].unsqueeze(-1) * shifted + return output.reshape(bsz, seq_len, hidden_size) + + def prepare(self, hidden_states): + """Convolve the sublayer input; return it with the output side's dynamic kernel.""" + coefficients = self.kernel_projection(hidden_states).reshape( + *hidden_states.shape[:-1], 2, self.taps, self.num_groups + ) + return self._convolve(hidden_states, coefficients[..., 0, :, :], side=0), coefficients[ + ..., 1, :, : + ] + + def finish(self, hidden_states, state): + """Convolve the sublayer output using the kernel produced by :meth:`prepare`.""" + return self._convolve(hidden_states, state, side=1) + + +class CandidateSelector(nn.Module): + """Low-rank scorer for transitions between adjacent block positions' candidates. + + Scores an edge from a predecessor token ``p`` to a candidate token ``c`` at a + block position with hidden state ``h`` as:: + + edge(p -> c) = + unary_logit[c] + + i.e. a bilinear form between the two token codebooks, gated by the context. + ``successor_codebook`` is zero-initialized, so ``transition`` is zero and a + fresh selector reproduces the backbone's unary ranking exactly, as in the + reference implementation (SpecForge #772). + + Training scores each position's candidate set independently under teacher + forcing (:meth:`score_candidates`); serving walks the resulting lattice. + """ + + def __init__(self, hidden_size: int, vocab_size: int, rank: int, top_k: int, std: float): + """Build the predecessor/successor codebooks and the context projection.""" + super().__init__() + if rank < 1: + raise ValueError(f"DFlash2 selector_rank must be >= 1, got {rank}.") + if not 1 <= top_k <= vocab_size: + raise ValueError( + f"DFlash2 selector_top_k must be in [1, vocab_size={vocab_size}], got {top_k}." + ) + self.top_k = int(top_k) + self.rank = int(rank) + self.predecessor_codebook = nn.Parameter(torch.empty(int(vocab_size), int(rank))) + self.successor_codebook = nn.Parameter(torch.empty(int(vocab_size), int(rank))) + self.hidden_projection = nn.Linear(int(hidden_size), int(rank), bias=False) + nn.init.normal_(self.predecessor_codebook, std=std) + # The transition term starts as a no-op, so a fresh selector reproduces the + # backbone's unary proposal exactly and only learns to deviate from it. + nn.init.zeros_(self.successor_codebook) + + def score_candidates(self, candidate_ids, unary_logits, hidden_states, predecessor_ids): + """Add the predecessor transition score to a candidate set's unary logits. + + Args: + candidate_ids: Candidate token ids ``[..., K]``. + unary_logits: Backbone logits for those candidates ``[..., K]``. + hidden_states: Backbone hidden at this position ``[..., H]``. + predecessor_ids: Teacher-forced predecessor token ids ``[...]``. + + Returns: + Selector logits over the candidate set ``[..., K]``. + """ + predecessor = self.predecessor_codebook[predecessor_ids] + successor = self.successor_codebook[candidate_ids] + context = predecessor * self.hidden_projection(hidden_states) + transition = torch.einsum("...r,...kr->...k", context.to(successor.dtype), successor) + return unary_logits + transition + + @torch.no_grad() + def greedy_path(self, candidate_ids, unary_logits, hidden_states, anchor_token_ids): + """Walk the candidate lattice greedily, mirroring the serving-side path walk. + + Position 0 here is block offset 1, not 0: the walk is seeded with the anchor + token, which is the predecessor the training objective pairs with offset 1 + (slot 0 is the given anchor and is never supervised). Passing the full + ``0..block_size-1`` candidate set therefore shifts the whole path by one. + + Args: + candidate_ids: ``[B, L, K]`` candidate ids for block offsets ``1..L``. + unary_logits: ``[B, L, K]`` backbone logits for those candidates. + hidden_states: ``[B, L, H]`` backbone hidden at those offsets. + anchor_token_ids: ``[B]`` the verified token at the anchor, i.e. offset 0. + + Returns: + Selected token ids ``[B, L]``. + """ + predecessor_ids = anchor_token_ids + path = [] + for position in range(candidate_ids.shape[1]): + scores = self.score_candidates( + candidate_ids[:, position], + unary_logits[:, position], + hidden_states[:, position], + predecessor_ids, + ) + selected = scores.argmax(dim=-1, keepdim=True) + predecessor_ids = candidate_ids[:, position].gather(1, selected)[:, 0] + path.append(predecessor_ids) + return torch.stack(path, dim=1) + + +class DFlash2Module(DFlashModule): + """DFlash draft backbone with per-sublayer convolutions and a candidate selector.""" + + def __init__(self, config): + """Initialize the DFlash backbone, then attach the convolutions and the selector.""" + super().__init__(config) + + self.projector_type = getattr(config, "projector_type", "dflash2") + + def required_int(name: str) -> int: + """Read an int architecture field, rejecting missing values and bools.""" + value = getattr(config, name, None) + if not isinstance(value, int) or isinstance(value, bool): + raise ValueError( + f"DFlash2 (projector_type='dflash2') requires an integer " + f"'{name}' in dflash_architecture_config, got {value!r}." + ) + return value + + taps = required_int("conv_kernel_size") + group_size = required_int("conv_group_size") + rank = required_int("selector_rank") + top_k = required_int("selector_top_k") + + std = getattr(config, "initializer_range", 0.02) + + # Replace each layer's no-op sublayer wrappers with real convolutions. The + # backbone layer forward already calls prepare()/finish() around attention + # and the MLP, so nothing else in the layer changes. + for layer in self.layers: + for wrapper_name in ("attention_conv", "mlp_conv"): + setattr( + layer, + wrapper_name, + DFlashGroupedConv( + hidden_size=config.hidden_size, + block_size=self.block_size, + taps=taps, + group_size=group_size, + ), + ) + + self.candidate_selector = CandidateSelector( + hidden_size=config.hidden_size, + vocab_size=config.vocab_size, + rank=rank, + top_k=top_k, + std=std, + ) + + # DFlashModule.__init__ already ran _init_weights before these modules + # existed, so initialize the new Linear layers explicitly. base_kernel and + # the codebooks keep the init set in their own constructors. + self._init_head_weights(std) + + def _init_head_weights(self, std: float): + """Initialize the convolution and selector Linear layers (matching HF _init_weights).""" + nn.init.normal_(self.candidate_selector.hidden_projection.weight, mean=0.0, std=std) + # The dynamic kernel starts at zero so every convolution is an exact identity at + # init and the draft begins as its DFlash backbone. The projection still trains: + # delta multiplies the sublayer activation, so dL/dW = dL/d(delta) . x^T is nonzero. + for layer in self.layers: + for wrapper in (layer.attention_conv, layer.mlp_conv): + nn.init.zeros_(wrapper.kernel_projection.weight) diff --git a/modelopt/torch/speculative/plugins/modeling_fakebase.py b/modelopt/torch/speculative/plugins/modeling_fakebase.py index 2b5fe989c..9ef6a1d26 100644 --- a/modelopt/torch/speculative/plugins/modeling_fakebase.py +++ b/modelopt/torch/speculative/plugins/modeling_fakebase.py @@ -124,7 +124,12 @@ class FakeBaseConfig(PretrainedConfig): ) self.intermediate_size = intermediate_size # For some drafter algo (e.g. DFlash) rope theta must match target model. Extract here. + # Published in both shapes: Transformers 5 consumers read the dict first, and it has + # to be well-formed -- the EAGLE draft config is built from this class too, and + # rotary embeddings index rope_parameters["rope_type"] unconditionally. self.rope_theta = rope_theta + if rope_theta is not None: + self.rope_parameters = {"rope_theta": rope_theta, "rope_type": "default"} if isinstance(dtype, str): dtype = getattr(torch, dtype) self.dtype = dtype @@ -179,6 +184,8 @@ class FakeBaseModel(PreTrainedModel): local checkpoint; otherwise it is treated as a Hub repo ID and the required files are downloaded via ``huggingface_hub``. """ + from modelopt.torch.export.plugins.hf_spec_export import _get_rope_theta + orig_config = transformers.AutoConfig.from_pretrained( source, trust_remote_code=trust_remote_code ) @@ -203,7 +210,9 @@ class FakeBaseModel(PreTrainedModel): num_key_value_heads=getattr(base_cfg, "num_key_value_heads", None), intermediate_size=getattr(base_cfg, "intermediate_size", None), rms_norm_eps=getattr(base_cfg, "rms_norm_eps", 1e-6), - rope_theta=getattr(base_cfg, "rope_theta", None), + # Shared with the exporter: where a config keeps rope_theta depends on the + # transformers version, and reading it wrong is silent until serve time. + rope_theta=_get_rope_theta(base_cfg), final_norm_type=_select_final_norm_type( getattr(base_cfg, "model_type", None), base_cfg ), diff --git a/modelopt/torch/speculative/plugins/modeling_lilicorr.py b/modelopt/torch/speculative/plugins/modeling_lilicorr.py index 6644f97e5..7d8924d54 100644 --- a/modelopt/torch/speculative/plugins/modeling_lilicorr.py +++ b/modelopt/torch/speculative/plugins/modeling_lilicorr.py @@ -527,14 +527,16 @@ class LiLiCorrModule(DFlashModule): ``DFlashDecoderLayer`` already exposes, so the convolution itself is shared code rather than a second implementation of the same arithmetic. - The initialization is the one deliberate difference. DFlash2 draws - ``kernel_projection`` from ``normal_(0, initializer_range)``, so its convolution - is not the identity at step 0. Here it is zero by default, and since - ``base_kernel`` is identity at tap 0 the whole wrapper is then an *exact* - identity at init: ``prepare`` emits a zero dynamic kernel, ``coefficients == - base``, and the convolution returns its input unchanged. That makes the - difference between a conv and a non-conv run attributable to the convolutions - rather than to a perturbed starting point. + The initialization is written here rather than inherited. + ``conv_projection_init_std`` is zero by default, and since ``base_kernel`` is + identity at tap 0 the whole wrapper is then an *exact* identity at init: + ``prepare`` emits a zero dynamic kernel, ``coefficients == base``, and the + convolution returns its input unchanged. That makes the difference between a + conv and a non-conv run attributable to the convolutions rather than to a + perturbed starting point. Assigning the weight explicitly is what keeps that + property LiLiCorr's own: it holds whatever ``DFlashGroupedConv`` does for + DFlash2, which since the DFlash2 merge also zeroes the projection in its own + constructor. ``conv_projection_init_std`` is a separate key from ``initializer_range`` on purpose: the latter also seeds the reranker, so overloading it would couple two diff --git a/modelopt_recipes/general/speculative_decoding/dflash2.yaml b/modelopt_recipes/general/speculative_decoding/dflash2.yaml new file mode 100644 index 000000000..1f02e44b1 --- /dev/null +++ b/modelopt_recipes/general/speculative_decoding/dflash2.yaml @@ -0,0 +1,104 @@ +# DFlash2 speculative-decoding training recipe. +# +# DFlash2 (https://inco.ai/blog/dflash2/) reuses the DFlash mode/pipeline and adds +# two components, selected via dflash_architecture_config.projector_type=dflash2: +# - a grouped dynamic depthwise convolution around every attention/MLP sublayer, +# giving each block position a view of its predecessors inside the block; +# - a low-rank candidate selector that scores transitions between adjacent block +# positions' top-k candidates, so serving walks one coherent path. +# The selector is trained by an extra cross-entropy term weighted by +# dflash_selector_loss_alpha. Online training is the default path (data.mode=online). +# Override fields via an OmegaConf dotlist. + +# modelopt-schema: modelopt.recipe.config.ModelOptDFlashRecipe +metadata: + description: DFlash2 training recipe (DFlash backbone + sublayer conv + candidate selector). + +# maps to ModelArguments (main.py) +model: + model_name_or_path: + trust_remote_code: false + use_fake_base_for_offline: false + +# maps to DataArguments (main.py) +data: + mode: online + data_path: + offline_data_path: + # Jinja chat template with {% generation %} tags for answer_only_loss. + chat_template: + +# maps to TrainingArguments (main.py) +training: + # --- commonly modified --- + output_dir: + num_train_epochs: 6 + per_device_train_batch_size: 1 + learning_rate: 6.0e-4 + warmup_ratio: 0.04 + training_seq_len: 3072 + logging_steps: 50 + save_steps: 2000 + cp_size: 1 + dp_shard_size: 1 + disable_tqdm: true + # Keep off: eval takes a plain per-position argmax, so the candidate selector is + # not applied. The convolutions are -- they run inside the backbone layers -- so AR + # here measures backbone + convolutions and understates the trained model. Compare + # via export + the offline acceptance-length harness instead. + estimate_ar: false + ar_validate_steps: 0 + answer_only_loss: true + + # --- rarely modified --- + do_eval: false + lr_scheduler_type: linear + save_strategy: steps + weight_decay: 0.0 + max_grad_norm: 1.0 + dataloader_drop_last: true + bf16: true + tf32: true + remove_unused_columns: false + # Safe default: the selector params are unused when + # dflash_selector_loss_alpha == 0, which would otherwise trip DDP. + ddp_find_unused_parameters: true + ddp_timeout: 1800 + report_to: tensorboard + +# maps to DFlashConfig (modelopt/torch/speculative/config.py). +dflash: + dflash_block_size: 16 + dflash_num_anchors: 256 + dflash_use_torch_compile: false + dflash_self_logit_distillation: false + # gamma for exponential loss decay (block_size=16 -> 7). + dflash_loss_decay_factor: 7.0 + # Qwen3 has no native mask token; 151669 is an unused id used by the reference. + dflash_mask_token_id: 151669 + # Weight of the candidate-selector cross-entropy term (0 disables it and trains + # the backbone + convolutions only). + dflash_selector_loss_alpha: 1.0 + # Anneal the block objective from cross-entropy toward acceptance as acceptance + # rises; see dflash_lk_loss_type. Requires dflash_self_logit_distillation: false. + dflash_lk_loss_type: lambda + dflash_lk_ce_scale: 1.0 + dflash_lk_ce_decay: 1.0 + dflash_architecture_config: + num_hidden_layers: 5 + # Draft attention/MLP dims — set explicitly (the draft is an independent + # Qwen3 model and does NOT inherit these from the base). GQA: 8 KV heads. + num_attention_heads: 32 + num_key_value_heads: 8 + head_dim: 128 + intermediate_size: 12288 + projector_type: dflash2 + # Grouped dynamic depthwise convolution. conv_kernel_size is the number of taps + # (2 = each position also sees its predecessor); it must not exceed the block + # size. conv_group_size must divide hidden_size. + conv_kernel_size: 2 + conv_group_size: 16 + # Candidate selector: rank of the transition codebooks, and how many of the + # backbone top-k candidates per position it re-ranks. + selector_rank: 256 + selector_top_k: 16 diff --git a/modelopt_recipes/general/speculative_decoding/lilicorr_conv.yaml b/modelopt_recipes/general/speculative_decoding/lilicorr_conv.yaml index 0fdea64b2..7a7a06b9c 100644 --- a/modelopt_recipes/general/speculative_decoding/lilicorr_conv.yaml +++ b/modelopt_recipes/general/speculative_decoding/lilicorr_conv.yaml @@ -13,12 +13,13 @@ # either. It is a standalone recipe because recipes do not compose; keep the two files in # sync when editing shared fields. # -# THE INIT IS THE ONE DELIBERATE DIFFERENCE FROM DFlash2. `conv_projection_init_std: 0.0` -# zeroes `kernel_projection`, and `base_kernel` is identity at tap 0, so the wrapper is an -# EXACT identity at step 0: a run from this recipe begins as the plain reranker and any -# difference is attributable to the convolutions rather than to a perturbed start. DFlash2 -# draws the same projection from `normal_(0, initializer_range)` and is therefore not the -# identity at init. Do not "align" the two -- they answer different questions. +# THE INIT IS THIS RECIPE'S OWN. `conv_projection_init_std: 0.0` zeroes +# `kernel_projection`, and `base_kernel` is identity at tap 0, so the wrapper is an EXACT +# identity at step 0: a run from this recipe begins as the plain reranker and any +# difference is attributable to the convolutions rather than to a perturbed start. +# `_install_sublayer_convs` assigns the weight explicitly, so this holds whatever +# `DFlashGroupedConv` does for DFlash2 -- which since the DFlash2 merge also zeroes it. +# Raise this key if you want a perturbed start instead. # # MEMORY. The convolutions add ~42M trainable parameters (20 tensors for a 5-layer draft), # and the activations they hold are the binding constraint. At an 8B target, combined with diff --git a/tests/unit/torch/export/test_hf_spec_rope_export.py b/tests/unit/torch/export/test_hf_spec_rope_export.py index fbeb21879..7cb0a317d 100644 --- a/tests/unit/torch/export/test_hf_spec_rope_export.py +++ b/tests/unit/torch/export/test_hf_spec_rope_export.py @@ -18,9 +18,14 @@ from types import SimpleNamespace from unittest.mock import MagicMock +import pytest import torch -from modelopt.torch.export.plugins.hf_spec_export import DFlashExporter, EagleExporter +from modelopt.torch.export.plugins.hf_spec_export import ( + DFlashExporter, + EagleExporter, + _get_rope_theta, +) DEFAULT_ROPE_SCALING = { "rope_type": "yarn", @@ -152,3 +157,65 @@ def test_dflash_rope_theta_inherits_base_rope_parameters(): config = exporter._export_config() assert config["rope_theta"] == 5000000.0 + + +class TestGetRopeTheta: + """Where a config keeps rope_theta depends on the transformers version. + + Every consumer of a base config -- the exporter, the draft builder, the fake base -- + has to agree on this, so they share this one reader. Reading it wrong is silent: the + draft trains and exports without complaint against a RoPE base the target never used, + and only misbehaves at serve time. + """ + + def test_reads_the_rope_parameters_dict(self): + """The transformers 5.12+ layout: the value lives only in the dict.""" + assert _get_rope_theta(SimpleNamespace(rope_parameters={"rope_theta": 1000000.0})) == ( + 1000000.0 + ) + + def test_reads_the_legacy_rope_scaling_dict(self): + """Older transformers spell the same dict rope_scaling.""" + assert _get_rope_theta(SimpleNamespace(rope_scaling={"rope_theta": 1000000.0})) == ( + 1000000.0 + ) + + def test_prefers_the_dict_over_a_disagreeing_flat_field(self): + """Both present and disagreeing: the dict wins. + + This is the regression guard. The precedence was the other way round on main from + 2026-07-30 to 2026-09-09, and no test noticed -- a config can carry the real base + in the dict while the class default (10000.0 for Qwen3) stays visible as a flat + rope_theta, so reading flat first exports a drafter whose RoPE base is 100x off. + """ + config = SimpleNamespace(rope_theta=10000.0, rope_parameters={"rope_theta": 1000000.0}) + assert _get_rope_theta(config) == 1000000.0 + + def test_falls_back_to_a_flat_attribute(self): + """The transformers 4.x layout: only the flat field exists.""" + assert _get_rope_theta(SimpleNamespace(rope_theta=12345.0)) == 12345.0 + + def test_missing_everywhere_returns_the_default(self): + """An absent base must stay absent rather than become a wrong number.""" + assert _get_rope_theta(SimpleNamespace()) is None + assert _get_rope_theta(SimpleNamespace(), 7.0) == 7.0 + + def test_reads_a_real_config_whichever_layout_it_uses(self): + """A real config resolves on every supported transformers version. + + Asserts the outcome, not the layout: 5.12 keeps the value only in the dict while + the minimum supported version (4.57) has only the flat field and no dict at all. + The layouts themselves are pinned above, built explicitly, so they stay covered + where transformers is absent -- this is the only test here that needs it. + """ + transformers = pytest.importorskip("transformers") + config = transformers.Qwen3Config( + hidden_size=32, + num_hidden_layers=2, + num_attention_heads=4, + num_key_value_heads=2, + intermediate_size=64, + vocab_size=64, + rope_theta=1000000.0, + ) + assert _get_rope_theta(config) == 1000000.0 diff --git a/tests/unit/torch/speculative/plugins/test_fakebase.py b/tests/unit/torch/speculative/plugins/test_fakebase.py index cf6dfe1a6..be959b22e 100644 --- a/tests/unit/torch/speculative/plugins/test_fakebase.py +++ b/tests/unit/torch/speculative/plugins/test_fakebase.py @@ -24,7 +24,7 @@ import torch pytest.importorskip("transformers") import transformers -from modelopt.torch.speculative.plugins.modeling_fakebase import FakeBaseModel +from modelopt.torch.speculative.plugins.modeling_fakebase import FakeBaseConfig, FakeBaseModel from modelopt.torch.speculative.utils import load_vlm_or_llm _HIDDEN_SIZE = 16 @@ -162,3 +162,82 @@ def test_load_vlm_or_llm_uses_transformers5_vlm_auto_class(monkeypatch): assert load_vlm_or_llm("qwen3-vl", dtype="auto") is not None assert captured["args"] == ("qwen3-vl",) assert captured["kwargs"]["torch_dtype"] == "auto" + + +class TestFakeBaseRopeTheta: + """The RoPE base has to survive the fake base, whichever shape the target stores it in. + + The draft injects the target's KV, so a mismatched base trains and exports without + complaint and only misbehaves at serve time. + """ + + def test_from_source_carries_a_transformers_5_base_theta(self, tmp_path, monkeypatch): + """The seam, not the reader. + + Reading rope_theta correctly is the exporter's ``_get_rope_theta`` and is tested + there. What is pinned here is that this call site uses it: a plain + ``getattr(base_cfg, "rope_theta")`` reads None from a transformers 5 config and + silently builds a draft with no RoPE base at all. + """ + base_cfg = transformers.PretrainedConfig( + model_type="llama", + hidden_size=_HIDDEN_SIZE, + vocab_size=_VOCAB_SIZE, + num_hidden_layers=2, + max_position_embeddings=128, + tie_word_embeddings=False, + rope_parameters={"rope_theta": 1000000.0}, + ) + monkeypatch.setattr(transformers.AutoConfig, "from_pretrained", lambda *a, **kw: base_cfg) + tensors = { + "lm_head.weight": torch.zeros(_VOCAB_SIZE, _HIDDEN_SIZE), + "embed_tokens.weight": torch.zeros(_VOCAB_SIZE, _HIDDEN_SIZE), + "norm.weight": torch.ones(_HIDDEN_SIZE), + } + shard = tmp_path / "model-00001-of-00001.safetensors" + safetensors.torch.save_file(tensors, shard) + (tmp_path / "model.safetensors.index.json").write_text( + json.dumps({"weight_map": dict.fromkeys(tensors, shard.name)}) + ) + + model = FakeBaseModel.from_source(str(tmp_path)) + + assert model.config.rope_theta == 1000000.0 + assert model.config.rope_parameters == {"rope_theta": 1000000.0, "rope_type": "default"} + + def test_config_publishes_both_shapes(self): + """Consumers that prefer the dict must find it on a fake base too.""" + config = FakeBaseConfig(num_hidden_layers=2, hidden_size=32, rope_theta=1000000.0) + assert config.rope_theta == 1000000.0 + assert config.rope_parameters == {"rope_theta": 1000000.0, "rope_type": "default"} + + def test_config_drives_a_transformers_rotary_embedding(self): + """The published dict has to satisfy transformers, not just carry the number. + + This class is also the class the EAGLE draft config is built from, so the dict + reaches `LlamaRotaryEmbedding`, which indexes ``rope_parameters["rope_type"]`` + unconditionally. + """ + from transformers.models.llama.modeling_llama import LlamaRotaryEmbedding + + config = FakeBaseConfig( + num_hidden_layers=2, + hidden_size=64, + num_attention_heads=4, + num_key_value_heads=2, + max_position_embeddings=128, + rope_theta=1000000.0, + ) + config.head_dim = 16 + + rotary = LlamaRotaryEmbedding(config=config) + cos, sin = rotary(torch.zeros(1, 4, 64), torch.arange(4).unsqueeze(0)) + + assert rotary.rope_type == "default" + assert cos.shape == (1, 4, 16) + + def test_unknown_theta_publishes_no_dict(self): + """An absent base must stay absent rather than become a wrong default.""" + config = FakeBaseConfig(num_hidden_layers=2, hidden_size=32, rope_theta=None) + assert config.rope_theta is None + assert not getattr(config, "rope_parameters", None) diff --git a/tests/unit/torch/speculative/plugins/test_hf_dflash2.py b/tests/unit/torch/speculative/plugins/test_hf_dflash2.py new file mode 100644 index 000000000..931a4a3ee --- /dev/null +++ b/tests/unit/torch/speculative/plugins/test_hf_dflash2.py @@ -0,0 +1,588 @@ +# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""CPU unit tests for the DFlash2 speculative decoding plugin. + +DFlash2 reuses the DFlash mode/pipeline and adds grouped dynamic convolutions +around every attention/MLP sublayer plus a low-rank candidate selector. These +tests cover conversion routing, the convolution's two structural invariants +(identity at initialization, no leakage across the block boundary), the selector +training objective, and the export format against the SGLang/vLLM +``DFlash2DraftModel`` layout (``attention_conv.*`` / ``mlp_conv.*`` / +``candidate_selector.*``). +""" + +import json +from copy import deepcopy + +import pytest +import torch +from _test_utils.torch.transformers_models import get_tiny_llama +from safetensors.torch import load_file + +import modelopt.torch.speculative as mtsp +from modelopt.torch.speculative.config import DFLASH_DEFAULT_CFG +from modelopt.torch.speculative.plugins.hf_dflash import HFDFlashModel +from modelopt.torch.speculative.plugins.hf_dflash2 import HFDFlash2Model +from modelopt.torch.speculative.plugins.modeling_dflash import ( + DFlashModule, + _IdentitySublayerWrapper, +) +from modelopt.torch.speculative.plugins.modeling_dflash2 import ( + CandidateSelector, + DFlash2Module, + DFlashGroupedConv, +) + +BLOCK_SIZE = 4 +NUM_DRAFT_LAYERS = 2 +SEQ_LEN = 16 # must be a multiple of BLOCK_SIZE +CONV_KERNEL_SIZE = 2 +CONV_GROUP_SIZE = 4 +SELECTOR_RANK = 8 +SELECTOR_TOP_K = 5 + +ARCH_FIELDS = ["conv_kernel_size", "conv_group_size", "selector_rank", "selector_top_k"] + + +def _get_dflash2_config(selector_loss_alpha=1.0, block_size=BLOCK_SIZE, **arch_overrides): + """Create a DFlash2 config for testing (dflash mode + projector_type=dflash2).""" + config = deepcopy(DFLASH_DEFAULT_CFG["config"]) + config["dflash_block_size"] = block_size + config["dflash_use_torch_compile"] = False + config["dflash_mask_token_id"] = 0 # token 0 as mask for the tiny model + config["dflash_self_logit_distillation"] = False + config["dflash_selector_loss_alpha"] = selector_loss_alpha + config["dflash_architecture_config"] = { + "num_hidden_layers": NUM_DRAFT_LAYERS, + "projector_type": "dflash2", + "conv_kernel_size": CONV_KERNEL_SIZE, + "conv_group_size": CONV_GROUP_SIZE, + "selector_rank": SELECTOR_RANK, + "selector_top_k": SELECTOR_TOP_K, + **arch_overrides, + } + return config + + +def _make_batch(vocab_size): + torch.manual_seed(0) + input_ids = torch.randint(1, vocab_size, (2, SEQ_LEN)) + return input_ids, torch.ones_like(input_ids), input_ids.clone() + + +class TestDFlash2Convert: + """Test DFlash2 conversion routing and module construction.""" + + def test_convert_creates_dflash2_model(self): + """projector_type=dflash2 routes to HFDFlash2Model (a HFDFlashModel subclass).""" + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", _get_dflash2_config())]) + assert isinstance(model, HFDFlash2Model) + assert isinstance(model, HFDFlashModel) + assert isinstance(model.dflash_module, DFlash2Module) + + def test_every_sublayer_wrapped_in_a_convolution(self): + """Both sublayer wrappers on every draft layer become real convolutions.""" + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", _get_dflash2_config())]) + layers = model.dflash_module.layers + assert len(layers) == NUM_DRAFT_LAYERS + for layer in layers: + for conv in (layer.attention_conv, layer.mlp_conv): + assert isinstance(conv, DFlashGroupedConv) + assert conv.taps == CONV_KERNEL_SIZE + assert conv.group_size == CONV_GROUP_SIZE + + def test_selector_shapes(self): + """The candidate selector's codebooks and projection are sized from the config.""" + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", _get_dflash2_config())]) + selector = model.dflash_module.candidate_selector + vocab = model.dflash_config.vocab_size + assert selector.top_k == SELECTOR_TOP_K + assert selector.predecessor_codebook.shape == (vocab, SELECTOR_RANK) + assert selector.successor_codebook.shape == (vocab, SELECTOR_RANK) + assert selector.hidden_projection.out_features == SELECTOR_RANK + assert selector.hidden_projection.bias is None + + def test_new_params_trainable(self): + """The convolution and selector parameters are trainable.""" + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", _get_dflash2_config())]) + new = [ + (n, p) + for n, p in model.named_parameters() + if "_conv." in n or "candidate_selector" in n + ] + assert len(new) >= 2 * NUM_DRAFT_LAYERS * 2 + 3 + assert all(p.requires_grad for _, p in new) + + @pytest.mark.parametrize("field", ARCH_FIELDS) + def test_missing_architecture_field_raises(self, field): + """projector_type=dflash2 without a required architecture field is an error.""" + config = _get_dflash2_config() + del config["dflash_architecture_config"][field] + model = get_tiny_llama(num_hidden_layers=4) + with pytest.raises(ValueError, match=field): + mtsp.convert(model, [("dflash", config)]) + + def test_conv_kernel_larger_than_block_raises(self): + """A convolution tap count exceeding the block size is an error.""" + model = get_tiny_llama(num_hidden_layers=4) + config = _get_dflash2_config(conv_kernel_size=BLOCK_SIZE + 1) + with pytest.raises(ValueError, match="conv_kernel_size"): + mtsp.convert(model, [("dflash", config)]) + + def test_conv_group_size_must_divide_hidden(self): + """A conv_group_size that does not divide hidden_size is an error.""" + model = get_tiny_llama(num_hidden_layers=4) + config = _get_dflash2_config(conv_group_size=model.config.hidden_size - 1) + with pytest.raises(ValueError, match="conv_group_size"): + mtsp.convert(model, [("dflash", config)]) + + def test_dflash_mode_still_creates_plain_dflash(self): + """Without projector_type=dflash2, conversion still yields a plain DFlash model.""" + config = deepcopy(DFLASH_DEFAULT_CFG["config"]) + config["dflash_mask_token_id"] = 0 + config["dflash_architecture_config"] = {"num_hidden_layers": NUM_DRAFT_LAYERS} + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", config)]) + assert isinstance(model, HFDFlashModel) + assert not isinstance(model, HFDFlash2Model) + assert type(model.dflash_module) is DFlashModule + # The sublayer seam stays a parameterless no-op for a plain DFlash draft. + for layer in model.dflash_module.layers: + assert isinstance(layer.attention_conv, _IdentitySublayerWrapper) + assert isinstance(layer.mlp_conv, _IdentitySublayerWrapper) + assert not any("_conv." in n for n, _ in model.dflash_module.named_parameters()) + + +class TestDFlashGroupedConv: + """Test the convolution's structural invariants directly.""" + + def _conv(self, hidden_size=32, taps=2): + torch.manual_seed(0) + return DFlashGroupedConv( + hidden_size=hidden_size, block_size=BLOCK_SIZE, taps=taps, group_size=CONV_GROUP_SIZE + ).double() + + def _trained_conv(self, **kwargs): + """A conv whose dynamic kernel is non-zero, i.e. what training produces. + + The default construction is an exact identity, so the structural assertions + below would pass vacuously on it. + """ + conv = self._conv(**kwargs) + with torch.no_grad(): + torch.nn.init.normal_(conv.kernel_projection.weight, std=0.5) + return conv + + def test_identity_at_initialization(self): + """A freshly built conv is an exact identity -- no manual zeroing needed. + + Both halves start as a no-op: the base kernel is an identity and + ``kernel_projection`` is zero-initialized. This is what makes enabling + DFlash2 a stable extension of a DFlash backbone rather than a perturbation + of it, and it is asserted on the DEFAULT construction so that a change to + either half fails here. + """ + conv = self._conv() + assert conv.kernel_projection.weight.abs().max() == 0.0 + x = torch.randn(2, SEQ_LEN, 32, dtype=torch.double) + out = conv.finish(*conv.prepare(x)) + assert torch.equal(out, x) + + def test_taps_do_not_cross_the_block_boundary(self): + """Perturbing the last position of a block leaves later blocks untouched.""" + conv = self._trained_conv() + x = torch.randn(2, SEQ_LEN, 32, dtype=torch.double) + baseline = conv.finish(*conv.prepare(x)) + + perturbed_input = x.clone() + perturbed_input[:, BLOCK_SIZE - 1] += 5.0 + perturbed = conv.finish(*conv.prepare(perturbed_input)) + + assert torch.allclose(baseline[:, BLOCK_SIZE:], perturbed[:, BLOCK_SIZE:], atol=1e-12) + assert not torch.allclose(baseline[:, :BLOCK_SIZE], perturbed[:, :BLOCK_SIZE]) + + def test_intra_block_dependency_is_backward_only(self): + """A position influences its successors inside the block, never its predecessors. + + This is the point of the convolution: it injects the sequential dependency the + parallel backbone lacks, without letting a position see the future. + """ + conv = self._trained_conv() + x = torch.randn(2, SEQ_LEN, 32, dtype=torch.double) + baseline = conv.finish(*conv.prepare(x)) + + perturbed_input = x.clone() + perturbed_input[:, 1] += 5.0 + perturbed = conv.finish(*conv.prepare(perturbed_input)) + + assert torch.allclose(baseline[:, 0], perturbed[:, 0], atol=1e-12) + assert not torch.allclose(baseline[:, 2], perturbed[:, 2]) + + def test_sequence_length_must_be_block_aligned(self): + """A sequence length not divisible by the block size is an error.""" + conv = self._conv() + with pytest.raises(ValueError, match="block_size"): + conv.prepare(torch.randn(1, BLOCK_SIZE + 1, 32, dtype=torch.double)) + + +class TestDFlash2Forward: + """Test the DFlash2 training forward (online path on CPU).""" + + def test_forward_grads_reach_conv_and_selector(self): + """Backward fills gradients on the convolutions, the selector and the backbone.""" + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", _get_dflash2_config())]) + model.train() + + input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size) + out = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels) + assert out.loss.requires_grad + assert out.loss.dim() == 0 + out.loss.backward() + + module = model.dflash_module + selector = module.candidate_selector + for grad in ( + selector.successor_codebook.grad, + module.layers[0].attention_conv.base_kernel.grad, + module.layers[0].mlp_conv.kernel_projection.weight.grad, + module.fc.weight.grad, + ): + assert grad is not None and torch.isfinite(grad).all() + assert grad.abs().sum() > 0 + + # The zero-initialized ``successor_codebook`` is one factor of a bilinear form, + # so on the FIRST step the other two factors get exactly zero gradient. This is + # a one-step delay, not a dead branch: the assertions below show they train as + # soon as ``successor_codebook`` moves off zero. The convolution has no such + # delay -- its delta is added to a non-zero base kernel, hence the grad above. + for grad in (selector.predecessor_codebook.grad, selector.hidden_projection.weight.grad): + assert grad is not None and torch.isfinite(grad).all() + assert grad.abs().sum() == 0 + + def test_selector_factors_train_once_the_successor_codebook_moves(self): + """After the first step the whole bilinear selector receives gradient. + + Guards the zero-init of ``successor_codebook`` against becoming a permanently + dead branch rather than the intended one-step warm start. + """ + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", _get_dflash2_config())]) + model.train() + + selector = model.dflash_module.candidate_selector + with torch.no_grad(): # stand in for the first optimizer step + torch.nn.init.normal_(selector.successor_codebook, std=0.02) + + input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size) + model(input_ids=input_ids, attention_mask=attention_mask, labels=labels).loss.backward() + + for grad in ( + selector.predecessor_codebook.grad, + selector.successor_codebook.grad, + selector.hidden_projection.weight.grad, + ): + assert grad is not None and torch.isfinite(grad).all() + assert grad.abs().sum() > 0 + + def test_selector_metrics_reported(self): + """The forward records selector accuracy and top-k coverage.""" + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", _get_dflash2_config())]) + model.train() + + input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size) + out = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels) + # Carried out on the forward output, not left on a private attribute, so the + # trainer can log them; and kept as tensors to avoid a per-step CPU-GPU sync. + metrics = out["selector_metrics"] + for key in ("selector_accuracy", "selector_coverage"): + assert torch.is_tensor(metrics[key]) + assert 0.0 <= metrics[key].item() <= 1.0 + + def test_selector_alpha_zero_disables_the_term(self): + """alpha=0 trains the backbone and convolutions only; the selector gets no grad.""" + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", _get_dflash2_config(selector_loss_alpha=0.0))]) + model.train() + + input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size) + out = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels) + out.loss.backward() + + module = model.dflash_module + codebook_grad = module.candidate_selector.predecessor_codebook.grad + assert codebook_grad is None or codebook_grad.abs().sum() == 0 + # The convolutions still train: they live inside the backbone. + conv_grad = module.layers[0].attention_conv.kernel_projection.weight.grad + assert conv_grad is not None and conv_grad.abs().sum() > 0 + + def test_selector_loss_increases_total_loss(self): + """The selector term adds to the backbone loss rather than replacing it.""" + input_ids, attention_mask, labels = _make_batch(32) + + losses = {} + for alpha in (0.0, 1.0): + torch.manual_seed(0) + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", _get_dflash2_config(selector_loss_alpha=alpha))]) + model.train() + torch.manual_seed(0) + out = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels) + losses[alpha] = float(out.loss.detach()) + assert losses[1.0] > losses[0.0] + + def test_overfits_a_single_batch(self): + """A few steps on one batch drive backbone and selector accuracy up. + + Guards the target/predecessor alignment: a misaligned selector objective still + produces a finite decreasing loss, but its accuracy does not reach 1. + """ + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", _get_dflash2_config())]) + model.train() + + input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size) + optimizer = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad], lr=5e-3) + for _ in range(60): + out = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels) + optimizer.zero_grad() + out.loss.backward() + optimizer.step() + + assert out.train_acc[0][0] > 0.9 + assert out["selector_metrics"]["selector_accuracy"].item() > 0.9 + + +class TestCandidateSelectorAlignment: + """Pin the offset convention shared by the training objective and the lattice walk.""" + + def _selector(self, vocab=16, rank=4, top_k=3, hidden=8): + torch.manual_seed(0) + selector = CandidateSelector( + hidden_size=hidden, vocab_size=vocab, rank=rank, top_k=top_k, std=0.5 + ).double() + with torch.no_grad(): # a fresh selector is a no-op; give it a real transition + torch.nn.init.normal_(selector.successor_codebook, std=0.5) + return selector + + def _selector_terms(self, selector, target_ids, candidate_ids): + """Drive HFDFlash2Model._selector_loss with a hand-built candidate set.""" + b, length, k = candidate_ids.shape + model = HFDFlash2Model.__new__(HFDFlash2Model) + model.dflash_module = type("_M", (), {"candidate_selector": selector})() + logits = torch.zeros(b, length, 32, dtype=torch.double) + # Put the chosen candidates on top so logits.topk reproduces candidate_ids. + for slot in range(k): + logits.scatter_(-1, candidate_ids[..., slot : slot + 1], float(k - slot)) + return HFDFlash2Model._selector_loss( + model, + logits, + target_ids, + torch.randn(b, length, 8, dtype=torch.double), + torch.zeros(b, length, dtype=torch.long), + torch.ones(b, length, dtype=torch.double), + ) + + def test_a_miss_carries_no_selector_gradient(self): + """A gold token outside the backbone's top-k is excluded, not substituted in. + + Serving only ever shows the strict top-k, so supervising a set with the gold + forced into it would teach the selector to override the unary ranking on a + candidate set it will never see. Contrasted against a covered set below, which + must still train -- masking everything would pass a one-sided assertion. + """ + selector = self._selector(vocab=32, top_k=2) + candidates = torch.tensor([[[5, 6], [5, 6], [5, 6]]] * 2) + + _, _, covered = self._selector_terms( + selector, torch.full((2, 3), 5, dtype=torch.long), candidates + ) + assert float(covered) == 1.0 + + loss, _, coverage = self._selector_terms( + selector, torch.full((2, 3), 9, dtype=torch.long), candidates + ) + assert float(coverage) == 0.0 + assert float(loss.detach()) == 0.0 + selector.zero_grad() + loss.backward() + assert selector.successor_codebook.grad.abs().sum() == 0.0 + + def test_a_covered_position_does_train_the_selector(self): + """The mask must not be so aggressive that nothing trains.""" + selector = self._selector(vocab=32, top_k=2) + candidates = torch.tensor([[[5, 6], [5, 6], [5, 6]]] * 2) + loss, _, coverage = self._selector_terms( + selector, torch.full((2, 3), 5, dtype=torch.long), candidates + ) + assert float(coverage) == 1.0 + assert float(loss.detach()) > 0.0 + selector.zero_grad() + loss.backward() + assert selector.successor_codebook.grad.abs().sum() > 0.0 + + def test_greedy_path_position_zero_is_seeded_by_the_anchor(self): + """``greedy_path`` position 0 scores against the anchor, i.e. block offset 1. + + ``HFDFlash2Model._compute_loss`` supervises offsets 1..block_size-1 and pairs + offset 1 with the anchor's own token. If either side changes its base offset + the objective and the walk silently disagree by one position, which a finite + decreasing loss does not catch. + """ + selector = self._selector() + b, length, k, hidden = 2, 3, 3, 8 + candidate_ids = torch.randint(0, 16, (b, length, k)) + unary = torch.randn(b, length, k, dtype=torch.double) + hiddens = torch.randn(b, length, hidden, dtype=torch.double) + anchor = torch.randint(0, 16, (b,)) + + path = selector.greedy_path(candidate_ids, unary, hiddens, anchor) + + expected_first = candidate_ids[:, 0].gather( + 1, + selector.score_candidates( + candidate_ids[:, 0], unary[:, 0], hiddens[:, 0], anchor + ).argmax(dim=-1, keepdim=True), + )[:, 0] + assert torch.equal(path[:, 0], expected_first) + + def test_greedy_path_feeds_each_choice_forward(self): + """Position n+1 is scored against the token position n actually selected.""" + selector = self._selector() + b, length, k, hidden = 2, 3, 3, 8 + candidate_ids = torch.randint(0, 16, (b, length, k)) + unary = torch.randn(b, length, k, dtype=torch.double) + hiddens = torch.randn(b, length, hidden, dtype=torch.double) + anchor = torch.randint(0, 16, (b,)) + + path = selector.greedy_path(candidate_ids, unary, hiddens, anchor) + + expected_second = candidate_ids[:, 1].gather( + 1, + selector.score_candidates( + candidate_ids[:, 1], unary[:, 1], hiddens[:, 1], path[:, 0] + ).argmax(dim=-1, keepdim=True), + )[:, 0] + assert torch.equal(path[:, 1], expected_second) + + +class TestDFlash2BlockObjective: + """The cross-entropy/acceptance blend selected by ``dflash_lk_loss_type``.""" + + def _loss(self, **overrides): + torch.manual_seed(0) + config = _get_dflash2_config() + config.update(overrides) + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", config)]) + model.train() + torch.manual_seed(1) + input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size) + out = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels) + return float(out.loss.detach()) + + def test_zero_decay_and_unit_scale_is_exactly_cross_entropy(self): + """A constant CE share of 1 must leave the objective bit-identical to 'ce'. + + This is the degenerate case that catches a blend wired up backwards: a wrong + sign or a swapped term still produces a finite decreasing loss. + """ + assert self._loss( + dflash_lk_loss_type="lambda", dflash_lk_ce_scale=1.0, dflash_lk_ce_decay=0.0 + ) == self._loss(dflash_lk_loss_type="ce") + + def test_zero_scale_is_exactly_the_acceptance_term(self): + """A CE share of 0 must leave the objective bit-identical to 'tv'.""" + assert self._loss(dflash_lk_loss_type="lambda", dflash_lk_ce_scale=0.0) == self._loss( + dflash_lk_loss_type="tv" + ) + + def test_blend_lies_between_its_two_terms(self): + ce = self._loss(dflash_lk_loss_type="ce") + tv = self._loss(dflash_lk_loss_type="tv") + blended = self._loss(dflash_lk_loss_type="lambda") + assert min(ce, tv) <= blended <= max(ce, tv) + + def test_acceptance_term_is_a_probability(self): + """1 - q(gold) is a weighted mean of probabilities, so it cannot leave [0, 1].""" + loss = self._loss(dflash_lk_loss_type="tv", dflash_selector_loss_alpha=0.0) + assert 0.0 <= loss <= 1.0 + + @pytest.mark.parametrize("loss_type", ["tv", "lambda"]) + def test_distillation_conflict_is_rejected(self, loss_type): + """Both terms read q(gold), which the KD path never forms.""" + config = _get_dflash2_config() + config["dflash_lk_loss_type"] = loss_type + config["dflash_self_logit_distillation"] = True + model = get_tiny_llama(num_hidden_layers=4) + with pytest.raises(ValueError, match="dflash_self_logit_distillation"): + mtsp.convert(model, [("dflash", config)]) + + +class TestDFlash2Export: + """Test the DFlash2 export format (weights + config).""" + + def _export(self, tmp_path): + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", _get_dflash2_config())]) + export_dir = tmp_path / "exported" + model.get_exporter().export(export_dir) + return export_dir + + def test_export_weight_keys_match_reference(self, tmp_path): + """Exported weights carry the DFlash2 tensors under reference names, no prefix.""" + sd = load_file(str(self._export(tmp_path) / "model.safetensors")) + for key in sd: + assert "dflash_module." not in key + assert "rotary_emb" not in key + + assert "candidate_selector.predecessor_codebook" in sd + assert "candidate_selector.successor_codebook" in sd + assert "candidate_selector.hidden_projection.weight" in sd + for layer_idx in range(NUM_DRAFT_LAYERS): + for wrapper in ("attention_conv", "mlp_conv"): + assert f"layers.{layer_idx}.{wrapper}.base_kernel" in sd + assert f"layers.{layer_idx}.{wrapper}.kernel_projection.weight" in sd + + def test_export_config_declares_dflash2_architecture(self, tmp_path): + """config.json selects the DFlash2 serving path and carries its fields. + + The architecture name matters: a checkpoint declaring ``DFlashDraftModel`` + loads as a plain DFlash draft and silently ignores these weights. + """ + with open(self._export(tmp_path) / "config.json") as f: + cfg = json.load(f) + + assert cfg["architectures"] == ["DFlash2DraftModel"] + # Emitted by the published-contract commit and read by the vLLM loader. Assert the + # top-level block_size too, not just the nested one: DFlash2Exporter derives the + # nested value FROM the top-level, so a wrong top-level passes the nested check. + assert cfg["is_causal"] is False + assert cfg["block_size"] == BLOCK_SIZE + dflash_config = cfg["dflash_config"] + assert dflash_config["block_size"] == BLOCK_SIZE + assert dflash_config["projector_type"] == "dflash2" + assert dflash_config["conv_kernel_size"] == CONV_KERNEL_SIZE + assert dflash_config["conv_group_size"] == CONV_GROUP_SIZE + assert dflash_config["selector_rank"] == SELECTOR_RANK + assert dflash_config["selector_top_k"] == SELECTOR_TOP_K + assert "mask_token_id" in dflash_config + assert "target_layer_ids" in dflash_config diff --git a/tests/unit/torch/speculative/plugins/test_hf_lilicorr.py b/tests/unit/torch/speculative/plugins/test_hf_lilicorr.py index f9bdbbca3..a21a8bdec 100644 --- a/tests/unit/torch/speculative/plugins/test_hf_lilicorr.py +++ b/tests/unit/torch/speculative/plugins/test_hf_lilicorr.py @@ -420,6 +420,99 @@ class TestLiLiCorrForward: assert draft_tokens.shape == (1, 3) +class TestLiLiCorrSublayerConvs: + """The optional grouped-convolution path, shared with DFlash2 via DFlashGroupedConv. + + This is the composition `lilicorr_conv.yaml` ships and the only place + `_install_sublayer_convs` runs, so it is also what guards LiLiCorr's init from + changes made on the DFlash2 side of the shared class. + """ + + CONV_KWARGS = {"conv_kernel_size": 2, "conv_group_size": 8} + + def _conv_converted(self, **arch_overrides): + config = _get_lilicorr_config() + config["dflash_architecture_config"].update(self.CONV_KWARGS) + config["dflash_architecture_config"].update(arch_overrides) + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", config)]) + return model + + def test_convs_replace_the_no_op_wrappers(self): + """Both sublayer wrappers on every draft layer become a real convolution.""" + from modelopt.torch.speculative.plugins.modeling_dflash2 import DFlashGroupedConv + + module = self._conv_converted().dflash_module + assert len(module.layers) == NUM_DRAFT_LAYERS + for layer in module.layers: + for wrapper_name in ("attention_conv", "mlp_conv"): + conv = getattr(layer, wrapper_name) + assert isinstance(conv, DFlashGroupedConv) + assert conv.taps == self.CONV_KWARGS["conv_kernel_size"] + assert conv.group_size == self.CONV_KWARGS["conv_group_size"] + + def test_plain_lilicorr_keeps_the_no_op_wrappers(self): + """Without the two geometry keys the draft is the plain reranker.""" + from modelopt.torch.speculative.plugins.modeling_dflash import _IdentitySublayerWrapper + + module = _converted().dflash_module + for layer in module.layers: + assert isinstance(layer.attention_conv, _IdentitySublayerWrapper) + assert isinstance(layer.mlp_conv, _IdentitySublayerWrapper) + + def test_default_init_is_an_exact_identity(self): + """`conv_projection_init_std` defaults to 0, so a conv run starts as the plain reranker. + + This is LiLiCorr's own choice, written by `_install_sublayer_convs` after the + conv is constructed. It must not depend on how `DFlashGroupedConv` happens to + initialize itself for DFlash2. + """ + module = self._conv_converted().dflash_module + for layer in module.layers: + for wrapper_name in ("attention_conv", "mlp_conv"): + conv = getattr(layer, wrapper_name) + assert conv.kernel_projection.weight.abs().max() == 0.0 + + conv = module.layers[0].attention_conv.double() + x = torch.randn(2, SEQ_LEN, conv.base_kernel.shape[-1], dtype=torch.double) + assert torch.equal(conv.finish(*conv.prepare(x)), x) + + def test_non_zero_init_std_perturbs_the_start(self): + """A non-zero `conv_projection_init_std` is still honoured, and only LiLiCorr sets it.""" + module = self._conv_converted(conv_projection_init_std=0.5).dflash_module + for layer in module.layers: + for wrapper_name in ("attention_conv", "mlp_conv"): + conv = getattr(layer, wrapper_name) + assert conv.kernel_projection.weight.abs().max() > 0.0 + + conv = module.layers[0].attention_conv.double() + x = torch.randn(2, SEQ_LEN, conv.base_kernel.shape[-1], dtype=torch.double) + assert not torch.allclose(conv.finish(*conv.prepare(x)), x) + + def test_forward_trains_the_convs(self): + """The conv path produces a finite loss and gradients reach the convolutions.""" + model = self._conv_converted() + model.train() + out = model(**_make_batch(model.dflash_config.vocab_size)) + assert torch.isfinite(out.loss).all() + out.loss.backward() + + for layer in model.dflash_module.layers: + for wrapper_name in ("attention_conv", "mlp_conv"): + conv = getattr(layer, wrapper_name) + for grad in (conv.base_kernel.grad, conv.kernel_projection.weight.grad): + assert grad is not None and torch.isfinite(grad).all() + assert grad.abs().sum() > 0 + + def test_one_geometry_key_alone_is_rejected(self): + """Half the geometry would silently build a draft with no convolutions.""" + config = _get_lilicorr_config() + config["dflash_architecture_config"]["conv_kernel_size"] = 2 + model = get_tiny_llama(num_hidden_layers=4) + with pytest.raises(ValueError, match="conv_kernel_size"): + mtsp.convert(model, [("dflash", config)]) + + class TestLiLiCorrOptimization: """The objective is trainable: a fixed batch is driven down. diff --git a/tools/launcher/examples/Qwen/Qwen3-8B/hf_online_dflash2.yaml b/tools/launcher/examples/Qwen/Qwen3-8B/hf_online_dflash2.yaml new file mode 100644 index 000000000..e300e0665 --- /dev/null +++ b/tools/launcher/examples/Qwen/Qwen3-8B/hf_online_dflash2.yaml @@ -0,0 +1,107 @@ +# DFlash2 online speculative decoding training for Qwen3-8B. +# +# DFlash2 = the DFlash draft backbone plus two additions: +# * a grouped dynamic depthwise convolution wrapped around every attention and +# MLP sublayer (taps clipped at block boundaries, identity-initialized so a +# fresh DFlash2 draft computes exactly what its DFlash backbone would), and +# * a low-rank candidate selector that scores transitions between adjacent +# block positions' top-k candidates, so serving walks one coherent path. +# See the dflash2.yaml recipe and +# modelopt/torch/speculative/plugins/{modeling,hf}_dflash2.py. +# +# 2-step pipeline: +# task_0: Build training conversations (Daring-Anteater multi-turn SFT, 50K) +# task_1: Online DFlash2 training + export of the drafter checkpoint +# +# As configured this is a short convergence check (max_steps=2000), matching the +# other Qwen3-8B online examples so it finishes on one node. To reproduce the +# published Qwen3-8B DFlash2 curve instead, see "Full run" below. +# +# Reference: inco.ai/blog/dflash2 | vLLM PR #52816 (serving support) +# +# Usage: +# uv run launch.py --yaml examples/Qwen/Qwen3-8B/hf_online_dflash2.yaml --yes +# uv run slurm.py --yaml modules/Model-Optimizer/tools/launcher/examples/Qwen/Qwen3-8B/hf_online_dflash2.yaml --yes +# +# Full run (the strict A/B against DFlash / DSpark / Domino, 3 epochs = ~92K +# steps on the 1.96M-conversation Spec-Decoding-Dataset-v1, 8 nodes x 8 H100): +# - point data.data_path at that corpus instead of task_0's output +# - training.num_train_epochs=3 and drop training.max_steps +# - training.save_steps=4000 +# - slurm_config.nodes=8 (global batch stays 64 = nodes x gpus x bs x accum) +# Every other knob below is already the A/B setting. + +job_name: Qwen3-8B_DFlash2_online +pipeline: + global_vars: + hf_model: /hf-local/Qwen/Qwen3-8B + + # Step 1: Build input conversations. example_data_config.yaml enables only the + # daring-anteater source (train: 50000) — multi-turn SFT with real assistant + # completions. --full-conversations keeps those completions so answer_only_loss + # has assistant spans to mask. make_dataset.sh writes /scratchspace/data/train.jsonl. + 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.0rc10 + + # Step 2: Online DFlash2 training (the script exports the drafter at the end). + # Consumes the conversations built in task_0 (shared via /scratchspace). + task_1: + script: common/specdec/dflash_online_training.sh + args: + - --config modules/Model-Optimizer/modelopt_recipes/general/speculative_decoding/dflash2.yaml + - model.model_name_or_path=<> + - data.data_path=/scratchspace/data/train.jsonl + - data.chat_template=examples/Qwen/Qwen3-8B/chat_template_train.jinja + - training.output_dir=/scratchspace/dflash2_bs16 + - training.per_device_train_batch_size=1 + - training.num_train_epochs=1 + - training.max_steps=2000 + - training.training_seq_len=4096 + - training.learning_rate=6.0e-4 + - training.warmup_ratio=0.04 + - training.warmup_steps=0 + - training.lr_scheduler_type=linear + - training.save_steps=5000 + - training.logging_steps=100 + - training.disable_tqdm=true + - training.answer_only_loss=true + # Draft backbone — identical to the DFlash / DSpark / Domino arms so the + # only difference between them is the correction head. + - dflash.dflash_block_size=16 + - dflash.dflash_num_anchors=512 + - dflash.dflash_loss_decay_factor=7 + - dflash.dflash_mask_token_id=151669 + - dflash.dflash_self_logit_distillation=false + - dflash.dflash_architecture_config.num_hidden_layers=5 + - dflash.dflash_architecture_config.num_attention_heads=32 + - dflash.dflash_architecture_config.num_key_value_heads=8 + - dflash.dflash_architecture_config.head_dim=128 + - dflash.dflash_architecture_config.intermediate_size=12288 + # DFlash2 knobs (also set in the recipe; repeated here for visibility). + # A draft dim NOT set explicitly falls back to the Qwen3Config default, not + # to the base model's — hence the five dims above are always spelled out. + - dflash.dflash_architecture_config.projector_type=dflash2 + - dflash.dflash_architecture_config.conv_kernel_size=2 + - dflash.dflash_architecture_config.conv_group_size=16 + - dflash.dflash_architecture_config.selector_rank=256 + - dflash.dflash_architecture_config.selector_top_k=16 + - dflash.dflash_selector_loss_alpha=1.0 + # Sliding-window draft attention, matching the published DFlash2 drafters. + - dflash.dflash_swa_window_size=2048 + environment: + - MAX_FINAL_LOSS: "5.0" + - MIN_FINAL_ACC: "0.15" + slurm_config: + _factory_: "slurm_factory" + nodes: 1 + ntasks_per_node: 1 + gpus_per_node: 8 diff --git a/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml b/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml new file mode 100644 index 000000000..89c06eece --- /dev/null +++ b/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml @@ -0,0 +1,124 @@ +# 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=<> + # 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: <> + # 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: <> + - 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